feat: EOD pull, EMA52 screening, DuckDB transient connections, config refactor

This commit is contained in:
jackyu66git
2026-07-03 12:12:48 +08:00
parent 21378c4f6d
commit cd602bf3ac
15 changed files with 1068 additions and 294 deletions
+275 -49
View File
@@ -4,7 +4,9 @@ from __future__ import annotations
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from fastapi import FastAPI from typing import Optional
from fastapi import FastAPI, HTTPException, Query
from fastapi.responses import HTMLResponse from fastapi.responses import HTMLResponse
from loguru import logger from loguru import logger
@@ -82,6 +84,10 @@ DOCS_HTML = r"""
<div class="stat"><div class="num">全天</div><div class="lbl">自动盘后拉取</div></div> <div class="stat"><div class="num">全天</div><div class="lbl">自动盘后拉取</div></div>
</div> </div>
<div class="tip" style="border-left-color: var(--green);">
📊 <strong><a href="screening/ema52">EMA52 筛选</a></strong> — 扫描日线/周线价格在 EMA52 ±3% 附近的股票,按成交额排序
</div>
<div class="toc"> <div class="toc">
<a href="#freqs">K线周期</a> <a href="#freqs">K线周期</a>
<a href="#rest-stocks">股票查询</a> <a href="#rest-stocks">股票查询</a>
@@ -279,17 +285,142 @@ ashare-dp query stats</code></pre>
""" """
EMA52_HTML = r"""
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>EMA52 筛选 — A-Share Data Platform</title>
<style>
:root {
--bg: #0d1117; --card: #161b22; --border: #30363d;
--text: #c9d1d9; --muted: #8b949e; --accent: #58a6ff;
--green: #3fb950; --red: #f85149; --orange: #d2991d;
}
* { margin: 0; padding: 0; box-sizing: border-box; }
body { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif; background: var(--bg); color: var(--text); }
.container { max-width: 1100px; margin: 0 auto; padding: 20px; }
h1 { font-size: 1.6em; margin-bottom: 4px; }
h1 span { color: var(--accent); }
.subtitle { color: var(--muted); margin-bottom: 20px; font-size: 0.9em; }
.controls { display: flex; gap: 12px; align-items: center; margin-bottom: 16px; flex-wrap: wrap; }
.controls button, .controls select {
background: var(--card); border: 1px solid var(--border); color: var(--text);
padding: 8px 16px; border-radius: 6px; cursor: pointer; font-size: 0.9em;
}
.controls button:hover { border-color: var(--accent); }
.controls button.active { background: var(--accent); color: #fff; border-color: var(--accent); }
.stats { color: var(--muted); font-size: 0.85em; margin-bottom: 12px; }
table { width: 100%; border-collapse: collapse; font-size: 0.9em; }
th { text-align: left; padding: 10px 12px; border-bottom: 2px solid var(--border); color: var(--muted); font-weight: 600; }
td { padding: 8px 12px; border-bottom: 1px solid var(--border); }
tr:hover { background: #1c2129; }
.code { font-family: 'SF Mono', 'Fira Code', monospace; color: var(--accent); }
.name { max-width: 100px; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; }
.pos { color: var(--green); }
.neg { color: var(--red); }
.near { color: var(--orange); }
.bar { display: inline-block; height: 6px; border-radius: 3px; vertical-align: middle; margin-right: 4px; }
.loading { text-align: center; padding: 40px; color: var(--muted); }
.empty { text-align: center; padding: 40px; color: var(--muted); }
a { color: var(--accent); text-decoration: none; }
a:hover { text-decoration: underline; }
.legend { display: flex; gap: 16px; margin-bottom: 12px; font-size: 0.8em; }
.legend span { display: flex; align-items: center; gap: 4px; }
.legend .dot { width: 8px; height: 8px; border-radius: 50%; }
.footer { text-align: center; color: var(--muted); font-size: 0.8em; margin-top: 30px; padding-bottom: 20px; }
</style>
</head>
<body>
<div class="container">
<h1><span>EMA52</span> 筛选</h1>
<p class="subtitle">每日收盘后扫描 1d / 1w 价格在 EMA52 附近的股票,按成交额降序排列</p>
<div class="controls">
<button id="btn-1d" class="active" onclick="switchFreq('1d')">日线 1d</button>
<button id="btn-1w" onclick="switchFreq('1w')">周线 1w</button>
<span style="flex:1"></span>
<span class="stats" id="stats">加载中…</span>
</div>
<div class="legend">
<span><span class="dot" style="background:var(--green)"></span> 价格 > EMA52</span>
<span><span class="dot" style="background:var(--red)"></span> 价格 < EMA52</span>
<span>阈值: ±3% | 按成交额↓</span>
</div>
<div id="content"><div class="loading">加载中…</div></div>
<div class="footer">
<a href="/">← 返回首页</a>
</div>
</div>
<script>
let currentFreq = '1d';
function switchFreq(freq) {
currentFreq = freq;
document.getElementById('btn-1d').className = freq === '1d' ? 'active' : '';
document.getElementById('btn-1w').className = freq === '1w' ? 'active' : '';
loadData();
}
function loadData() {
const content = document.getElementById('content');
content.innerHTML = '<div class="loading">加载中…</div>';
fetch(`/api/v1/screening/ema52?freq=${currentFreq}&limit=200`)
.then(r => r.json())
.then(data => {
document.getElementById('stats').textContent =
`共 ${data.total} 只股票在 EMA52 附近`;
if (!data.items || data.items.length === 0) {
content.innerHTML = '<div class="empty">暂无数据。请等待今日盘后扫描完成。</div>';
return;
}
let html = '<table><thead><tr><th>股票</th><th>价格</th><th>EMA52</th><th>成交额(亿)</th><th>偏离</th></tr></thead><tbody>';
data.items.forEach(item => {
const dist = item.distance_pct;
let cls = dist > 0 ? 'pos' : 'neg';
let barColor = dist > 0 ? 'var(--green)' : 'var(--red)';
let barW = Math.min(Math.abs(dist) * 20, 60);
const amt = item.amount != null ? item.amount.toFixed(1) : '';
html += `<tr>
<td><span class="code">${item.ts_code}</span> <span class="name" title="${item.name}">${item.name}</span></td>
<td>${item.close_price.toFixed(2)}</td>
<td>${item.ema52.toFixed(2)}</td>
<td>${amt}</td>
<td>
<span class="${cls}">
<span class="bar" style="width:${barW}px; background:${barColor}"></span>
${dist > 0 ? '+' : ''}${dist.toFixed(2)}%
</span>
</td>
</tr>`;
});
html += '</tbody></table>';
content.innerHTML = html;
})
.catch(() => {
content.innerHTML = '<div class="empty">加载失败,请确认服务在运行且盘后扫描已完成。</div>';
});
}
loadData();
</script>
</body>
</html>
"""
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
"""Startup and shutdown lifecycle.""" """Startup and shutdown lifecycle."""
logger.info("Starting A-Share Data Platform...") logger.info("Starting A-Share Data Platform...")
db = get_db() # Run DDL with write access (short-lived)
db.connect() with get_db(read_only=False) as db:
for ddl in DDL_STATEMENTS: for ddl in DDL_STATEMENTS:
try: try:
db.execute(ddl) db.execute(ddl)
except Exception as e: except Exception as e:
logger.warning(f"DDL warning: {e}") logger.warning(f"DDL warning: {e}")
logger.info("Database initialized") logger.info("Database initialized")
# Start realtime poller # Start realtime poller
@@ -302,18 +433,6 @@ async def lifespan(app: FastAPI):
logger.warning(f"Realtime poller not started: {e}") logger.warning(f"Realtime poller not started: {e}")
app.state.poller = None app.state.poller = None
# Start scheduler
try:
from ashare_dp.scheduler.scheduler import Scheduler
from ashare_dp.data.akshare_client import AKShareClient
scheduler = Scheduler(client=AKShareClient())
scheduler.start()
app.state.scheduler = scheduler
logger.info("Scheduler started")
except Exception as e:
logger.warning(f"Scheduler not started: {e}")
app.state.scheduler = None
yield yield
# Shutdown # Shutdown
@@ -322,12 +441,6 @@ async def lifespan(app: FastAPI):
await app.state.poller.stop() await app.state.poller.stop()
except Exception: except Exception:
pass pass
if app.state.scheduler:
try:
app.state.scheduler.shutdown()
except Exception:
pass
db.close()
logger.info("A-Share Data Platform stopped") logger.info("A-Share Data Platform stopped")
@@ -347,37 +460,150 @@ def create_app() -> FastAPI:
@app.get("/health") @app.get("/health")
async def health(): async def health():
"""Basic health check.""" """Health check in data_provider format (compatible with Chan system)."""
db = get_db()
repo = KLineRepository() repo = KLineRepository()
from ashare_dp.core.models import BACKFILLABLE_FREQS, DERIVED_FREQS, INTRADAY_FREQS, Freq
all_freqs = BACKFILLABLE_FREQS + INTRADAY_FREQS + DERIVED_FREQS
base_freqs = [f.value for f in BACKFILLABLE_FREQS + INTRADAY_FREQS]
derived_freqs = [f.value for f in DERIVED_FREQS]
try:
with get_db(read_only=True) as db:
stocks = db.query("SELECT ts_code FROM stock_info")
symbols = [s[0] for s in stocks]
except Exception:
symbols = []
return { return {
"status": "ok", "status": "ok",
"db_path": db.db_path, "exchange": "ashare",
"stocks": db.query("SELECT count(*) FROM stock_info")[0][0], "symbols": symbols,
"trading_days": db.query("SELECT count(*) FROM trading_calendar")[0][0], "base_timeframes": base_freqs,
"daily_records": repo.count_records(Freq.d1), "derived_timeframes": derived_freqs,
"timeframes": [f.value for f in all_freqs],
"ready": True,
} }
@app.get("/stats") @app.get("/timeframes")
async def stats(): async def timeframes():
"""Full DB statistics.""" """List available timeframes in data_provider format."""
db = get_db() from ashare_dp.core.models import BACKFILLABLE_FREQS, DERIVED_FREQS, INTRADAY_FREQS
repo = KLineRepository()
from ashare_dp.core.models import Freq all_freqs = BACKFILLABLE_FREQS + INTRADAY_FREQS + DERIVED_FREQS
freq_stats = {}
for f in [Freq.d1, Freq.w1, Freq.M1, Freq.m1, Freq.m5, Freq.m15, Freq.m30, Freq.h1]:
dr = repo.get_date_range(f)
freq_stats[f.value] = {
"records": repo.count_records(f),
"start_date": dr[0].isoformat() if dr[0] else None,
"end_date": dr[1].isoformat() if dr[1] else None,
}
return { return {
"stocks": db.query("SELECT count(*) FROM stock_info")[0][0], "base_timeframes": [f.value for f in BACKFILLABLE_FREQS + INTRADAY_FREQS],
"trading_days": db.query("SELECT count(*) FROM trading_calendar")[0][0], "derived_timeframes": [f.value for f in DERIVED_FREQS],
"frequencies": freq_stats, "timeframes": [f.value for f in all_freqs],
} }
@app.get("/api/candles")
async def api_candles(
symbol: str = Query(..., description="Stock code, e.g. 000001.SZ"),
tf: str = Query("1d", description="Timeframe: 1d, 1w, 1M, 1h, 5m, 1m, etc."),
start: Optional[int] = Query(None, description="Start timestamp in ms"),
end: Optional[int] = Query(None, description="End timestamp in ms"),
limit: Optional[int] = Query(None, description="Max candles to return"),
):
"""Fetch K-line candles in data_provider-compatible format.
Returns [{"timestamp": ..., "open": ..., "high": ..., "low": ..., "close": ..., "volume": ...}]
"""
from datetime import datetime, timezone
from ashare_dp.core.models import BACKFILLABLE_FREQS, DERIVED_FREQS, INTRADAY_FREQS, Freq
all_freqs = BACKFILLABLE_FREQS + INTRADAY_FREQS + DERIVED_FREQS
freq_values = [f.value for f in all_freqs]
if tf not in freq_values:
raise HTTPException(status_code=400, detail=f"Invalid timeframe: {tf}")
freq = Freq(tf)
repo = KLineRepository()
start_date = None
end_date = None
if start is not None:
start_date = datetime.fromtimestamp(start / 1000, tz=timezone.utc).date()
if end is not None:
end_date = datetime.fromtimestamp(end / 1000, tz=timezone.utc).date()
df = repo.read_klines(
freq=freq, ts_code=symbol,
start_date=start_date, end_date=end_date,
limit=limit or 10000, offset=0,
)
if df is None or df.empty:
return []
candles = []
for _, row in df.iterrows():
ts = row.get("trade_time")
if hasattr(ts, "timestamp"):
ts_ms = int(ts.timestamp() * 1000)
else:
ts_ms = 0
candles.append({
"timestamp": ts_ms,
"open": float(row.get("open", 0)),
"high": float(row.get("high", 0)),
"low": float(row.get("low", 0)),
"close": float(row.get("close", 0)),
"volume": float(row.get("volume", 0)),
})
if limit and len(candles) > limit:
candles = candles[-limit:]
return candles
# ── EMA52 Screening ──
@app.get("/screening/ema52", response_class=HTMLResponse)
async def ema52_screening_page():
"""EMA52 screening results page."""
return EMA52_HTML
@app.get("/api/v1/screening/ema52")
async def api_ema52_screening(
freq: str = Query("1d", description="Frequency: 1d or 1w"),
limit: int = Query(200, ge=1, le=1000),
offset: int = Query(0, ge=0),
):
"""Get EMA52 screening results from the database."""
if freq not in ("1d", "1w"):
raise HTTPException(status_code=400, detail="freq must be 1d or 1w")
with get_db(read_only=True) as db:
row = db.conn.execute(
"SELECT COUNT(*) FROM ema52_screening WHERE freq = ?",
[freq],
).fetchone()
total = row[0] if row else 0
results = db.query(
"""SELECT ts_code, name, trade_date, close_price, ema52, distance_pct, amount
FROM ema52_screening
WHERE freq = ?
ORDER BY amount DESC NULLS LAST
LIMIT ? OFFSET ?""",
[freq, limit, offset],
)
items = [
{
"ts_code": r[0],
"name": r[1],
"trade_date": str(r[2]) if r[2] else None,
"close_price": r[3],
"ema52": r[4],
"distance_pct": r[5],
"amount": r[6],
}
for r in results
]
return {"total": total, "limit": limit, "offset": offset, "items": items}
# Register routers # Register routers
app.include_router(stocks.router, prefix="/api/v1") app.include_router(stocks.router, prefix="/api/v1")
app.include_router(kline.router, prefix="/api/v1") app.include_router(kline.router, prefix="/api/v1")
-1
View File
@@ -1,6 +1,5 @@
"""FastAPI dependency injection.""" """FastAPI dependency injection."""
from ashare_dp.storage.database import get_db
from ashare_dp.storage.repository import KLineRepository from ashare_dp.storage.repository import KLineRepository
+22 -22
View File
@@ -19,7 +19,6 @@ async def list_stocks(
offset: int = Query(0, ge=0), offset: int = Query(0, ge=0),
): ):
"""List stocks with optional filters.""" """List stocks with optional filters."""
db = get_db()
conditions = ["1=1"] conditions = ["1=1"]
params = [] params = []
@@ -31,14 +30,15 @@ async def list_stocks(
params.append(market) params.append(market)
where = " AND ".join(conditions) where = " AND ".join(conditions)
rows = db.query( with get_db(read_only=True) as db:
f"SELECT * FROM stock_info WHERE {where} ORDER BY ts_code LIMIT ? OFFSET ?", rows = db.query(
tuple(params) + (limit, offset), f"SELECT * FROM stock_info WHERE {where} ORDER BY ts_code LIMIT ? OFFSET ?",
) tuple(params) + (limit, offset),
total = db.query( )
f"SELECT count(*) FROM stock_info WHERE {where}", total = db.query(
tuple(params), f"SELECT count(*) FROM stock_info WHERE {where}",
)[0][0] tuple(params),
)[0][0]
cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"] cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"]
items = [] items = []
@@ -57,11 +57,11 @@ async def list_stocks(
@router.get("/search") @router.get("/search")
async def search_stocks(q: str = Query(..., min_length=1, description="Search query")): async def search_stocks(q: str = Query(..., min_length=1, description="Search query")):
"""Fuzzy search stocks by name or code.""" """Fuzzy search stocks by name or code."""
db = get_db() with get_db(read_only=True) as db:
rows = db.query( rows = db.query(
"SELECT * FROM stock_info WHERE name LIKE ? OR symbol LIKE ? OR ts_code LIKE ? LIMIT 50", "SELECT * FROM stock_info WHERE name LIKE ? OR symbol LIKE ? OR ts_code LIKE ? LIMIT 50",
(f"%{q}%", f"%{q}%", f"%{q}%"), (f"%{q}%", f"%{q}%", f"%{q}%"),
) )
cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"] cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"]
items = [] items = []
for row in rows: for row in rows:
@@ -78,14 +78,14 @@ async def search_stocks(q: str = Query(..., min_length=1, description="Search qu
@router.get("/{ts_code}") @router.get("/{ts_code}")
async def get_stock(ts_code: str): async def get_stock(ts_code: str):
"""Get single stock info by ts_code (e.g. '000001.SZ').""" """Get single stock info by ts_code (e.g. '000001.SZ')."""
db = get_db() with get_db(read_only=True) as db:
row = db.query( row = db.query(
"SELECT * FROM stock_info WHERE ts_code = ?", "SELECT * FROM stock_info WHERE ts_code = ?",
(ts_code,), (ts_code,),
) )
if not row: if not row:
raise HTTPException(status_code=404, detail=f"Stock not found: {ts_code}") raise HTTPException(status_code=404, detail=f"Stock not found: {ts_code}")
row = row[0] row = row[0]
cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"] cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"]
item = {} item = {}
for i, col in enumerate(cols): for i, col in enumerate(cols):
-7
View File
@@ -11,7 +11,6 @@ from loguru import logger
from ashare_dp.core.models import BACKFILLABLE_FREQS, INTRADAY_FREQS from ashare_dp.core.models import BACKFILLABLE_FREQS, INTRADAY_FREQS
from ashare_dp.data.akshare_client import AKShareClient from ashare_dp.data.akshare_client import AKShareClient
from ashare_dp.data.backfill import BackfillPipeline from ashare_dp.data.backfill import BackfillPipeline
from ashare_dp.storage.database import get_db
backfill_app = typer.Typer() backfill_app = typer.Typer()
@@ -19,8 +18,6 @@ backfill_app = typer.Typer()
@backfill_app.command("init") @backfill_app.command("init")
def init_db(): def init_db():
"""Initialize database schema.""" """Initialize database schema."""
db = get_db()
db.connect()
pipeline = BackfillPipeline() pipeline = BackfillPipeline()
pipeline.init_db() pipeline.init_db()
pipeline.load_stock_list() pipeline.load_stock_list()
@@ -36,8 +33,6 @@ def backfill_daily(
symbols: str = typer.Option(None, help="Comma-separated stock symbols (default: all)"), symbols: str = typer.Option(None, help="Comma-separated stock symbols (default: all)"),
): ):
"""Backfill daily/weekly/monthly K-line data.""" """Backfill daily/weekly/monthly K-line data."""
db = get_db()
db.connect()
start_date = datetime.strptime(start, "%Y%m%d").date() start_date = datetime.strptime(start, "%Y%m%d").date()
end_date = datetime.strptime(end, "%Y%m%d").date() if end else date.today() end_date = datetime.strptime(end, "%Y%m%d").date() if end else date.today()
@@ -70,8 +65,6 @@ def backfill_minute(
symbols: str = typer.Option(None, help="Comma-separated stock symbols (default: all)"), symbols: str = typer.Option(None, help="Comma-separated stock symbols (default: all)"),
): ):
"""Backfill recent minute K-line data (limited API history).""" """Backfill recent minute K-line data (limited API history)."""
db = get_db()
db.connect()
sym_list = [s.strip() for s in symbols.split(",")] if symbols else None sym_list = [s.strip() for s in symbols.split(",")] if symbols else None
+58
View File
@@ -0,0 +1,58 @@
"""EOD CLI subcommand: pull end-of-day data."""
from __future__ import annotations
from datetime import date, datetime
import typer
from loguru import logger
from ashare_dp.core.calendar import BEIJING_TZ
from ashare_dp.data.eod import EODPipeline
eod_app = typer.Typer()
@eod_app.command("pull")
def eod_pull(
date_str: str = typer.Option(
None, "--date", "-d",
help="Trade date in YYYY-MM-DD format (default: today in Beijing time)",
),
daily_only: bool = typer.Option(
False, "--daily-only",
help="Only pull daily data, skip minute/weekly/monthly",
),
):
"""Pull end-of-day K-line data for all active stocks."""
now = datetime.now(BEIJING_TZ)
today = now.date()
if date_str:
trade_date = date.fromisoformat(date_str)
else:
trade_date = today
# Weekend check
if now.weekday() >= 5 and date_str is None:
logger.info("Weekend, skipping EOD pull (use --date to force)")
return
logger.info(f"EOD pull: {trade_date.isoformat()}")
# EOD writes to Parquet files directly (via pipeline), no DuckDB lock needed.
pipeline = EODPipeline()
symbols = pipeline.get_active_symbols()
logger.info(f"EOD: {len(symbols)} active stocks")
# Pull daily
pipeline.pull_daily(symbols, trade_date)
if not daily_only:
# Pull minute
pipeline.pull_minute(symbols, trade_date)
# Derive weekly/monthly
pipeline.pull_weekly_monthly(symbols, trade_date)
logger.info(f"EOD pull complete for {trade_date.isoformat()}")
+4
View File
@@ -7,6 +7,8 @@ import typer
from ashare_dp.cli.backfill_cmd import backfill_app from ashare_dp.cli.backfill_cmd import backfill_app
from ashare_dp.cli.serve_cmd import serve_app from ashare_dp.cli.serve_cmd import serve_app
from ashare_dp.cli.query_cmd import query_app from ashare_dp.cli.query_cmd import query_app
from ashare_dp.cli.eod_cmd import eod_app
from ashare_dp.cli.screening_cmd import screening_app
app = typer.Typer( app = typer.Typer(
name="ashare-dp", name="ashare-dp",
@@ -16,6 +18,8 @@ app = typer.Typer(
app.add_typer(backfill_app, name="backfill", help="Historical data backfill") app.add_typer(backfill_app, name="backfill", help="Historical data backfill")
app.add_typer(serve_app, name="serve", help="Start API server") app.add_typer(serve_app, name="serve", help="Start API server")
app.add_typer(query_app, name="query", help="Ad-hoc data queries") app.add_typer(query_app, name="query", help="Ad-hoc data queries")
app.add_typer(eod_app, name="eod", help="End-of-day data pull")
app.add_typer(screening_app, name="screening", help="Stock screening (EMA52, etc.)")
@app.command() @app.command()
+13 -22
View File
@@ -22,9 +22,6 @@ def query_kline(
limit: int = typer.Option(100, help="Max records"), limit: int = typer.Option(100, help="Max records"),
): ):
"""Query K-line data from the command line.""" """Query K-line data from the command line."""
db = get_db()
db.connect()
freq_enum = Freq(freq) freq_enum = Freq(freq)
repo = KLineRepository() repo = KLineRepository()
@@ -48,9 +45,6 @@ def query_latest(
ts_code: str = typer.Option(None, help="Stock code (optional)"), ts_code: str = typer.Option(None, help="Stock code (optional)"),
): ):
"""Show latest K-line data.""" """Show latest K-line data."""
db = get_db()
db.connect()
freq_enum = Freq(freq) freq_enum = Freq(freq)
repo = KLineRepository() repo = KLineRepository()
df = repo.get_latest(freq=freq_enum, ts_code=ts_code) df = repo.get_latest(freq=freq_enum, ts_code=ts_code)
@@ -70,19 +64,17 @@ def query_stocks(
limit: int = typer.Option(50, help="Max records"), limit: int = typer.Option(50, help="Max records"),
): ):
"""List stocks.""" """List stocks."""
db = get_db() with get_db(read_only=True) as db:
db.connect() if exchange:
rows = db.query(
if exchange: "SELECT ts_code, symbol, name, exchange, market, list_date FROM stock_info WHERE exchange = ? LIMIT ?",
rows = db.query( (exchange.upper(), limit),
"SELECT ts_code, symbol, name, exchange, market, list_date FROM stock_info WHERE exchange = ? LIMIT ?", )
(exchange.upper(), limit), else:
) rows = db.query(
else: "SELECT ts_code, symbol, name, exchange, market, list_date FROM stock_info LIMIT ?",
rows = db.query( (limit,),
"SELECT ts_code, symbol, name, exchange, market, list_date FROM stock_info LIMIT ?", )
(limit,),
)
typer.echo(f"\n{'ts_code':<12} {'symbol':<8} {'name':<12} {'exchange':<8} {'market':<10} {'list_date'}") typer.echo(f"\n{'ts_code':<12} {'symbol':<8} {'name':<12} {'exchange':<8} {'market':<10} {'list_date'}")
typer.echo("-" * 60) typer.echo("-" * 60)
@@ -95,15 +87,14 @@ def query_stocks(
@query_app.command("stats") @query_app.command("stats")
def query_stats(): def query_stats():
"""Show database statistics.""" """Show database statistics."""
db = get_db()
db.connect()
repo = KLineRepository() repo = KLineRepository()
typer.echo("\nDatabase Statistics:") typer.echo("\nDatabase Statistics:")
typer.echo("-" * 40) typer.echo("-" * 40)
# Stock count # Stock count
n = db.query("SELECT count(*) FROM stock_info")[0][0] with get_db(read_only=True) as db:
n = db.query("SELECT count(*) FROM stock_info")[0][0]
typer.echo(f" Stocks: {n}") typer.echo(f" Stocks: {n}")
# Record count and date range per frequency # Record count and date range per frequency
+168
View File
@@ -0,0 +1,168 @@
"""Screening CLI subcommand: EMA52 scanning and other screens."""
from __future__ import annotations
import typer
from loguru import logger
from ashare_dp.storage.database import get_db
screening_app = typer.Typer()
@screening_app.command("ema52")
def ema52(
threshold: float = typer.Option(
None, "--threshold", "-t",
help="EMA52 proximity threshold (e.g. 0.03 = ±3%%). Default from .env",
),
):
"""Run EMA52 screening for 1d and 1w.
Scans all stocks, computes EMA52 from close prices,
and saves results near EMA52 to the ema52_screening table.
"""
import numpy as np
import pandas as pd
from ashare_dp.config import Settings
from ashare_dp.core.models import Freq
settings = Settings()
if threshold is None:
threshold = settings.ema52_threshold
with get_db(read_only=False) as db:
# Get all ts_codes and names
stocks = db.query(
"SELECT ts_code, name FROM stock_info ORDER BY ts_code"
)
symbols = [row[0] for row in stocks]
symbol_names = {row[0]: row[1] for row in stocks}
logger.info(
f"EMA52 screening: {len(symbols)} stocks, "
f"threshold=±{threshold * 100:.1f}%"
)
for freq, freq_label, bars_needed in [
(Freq.d1, "1d", 100),
(Freq.w1, "1w", 60),
]:
try:
from ashare_dp.storage.partitioning import partition_glob
glob = partition_glob(freq)
sql = f"""
WITH ranked AS (
SELECT *, ROW_NUMBER() OVER (
PARTITION BY ts_code ORDER BY trade_time DESC
) as rn
FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true)
)
SELECT * FROM ranked WHERE rn <= {bars_needed}
ORDER BY ts_code, trade_time ASC
"""
df = _query_parquet(sql)
if df.empty:
logger.warning(f"EMA52: no data for {freq_label}")
continue
except Exception as e:
logger.error(f"EMA52: failed to read {freq_label}: {e}")
continue
df["trade_time"] = pd.to_datetime(df["trade_time"])
df = df.sort_values(["ts_code", "trade_time"])
results = []
for ts_code, group in df.groupby("ts_code"):
if ts_code not in symbol_names:
continue
group = group.tail(bars_needed)
if len(group) < 26:
continue
closes = group["close"].astype(float).values
ema_values = _compute_ema(closes, 52)
if len(ema_values) == 0:
continue
latest_close = closes[-1]
latest_ema = ema_values[-1]
if latest_ema <= 0:
continue
distance = (latest_close - latest_ema) / latest_ema
# Compute daily amount in 亿 CNY
latest_amount = float(group["amount"].iloc[-1])
amount_yi = latest_amount / 100_000_000 # 元 → 亿
if abs(distance) <= threshold:
results.append({
"trade_date": group["trade_date"].iloc[-1],
"ts_code": ts_code,
"name": symbol_names[ts_code],
"freq": freq_label,
"close_price": round(latest_close, 2),
"ema52": round(latest_ema, 2),
"distance_pct": round(distance * 100, 2),
"amount": round(amount_yi, 2),
})
if results:
# Sort by amount DESC (largest turnover first)
results.sort(key=lambda r: r.get("amount", 0) or 0, reverse=True)
try:
with get_db(read_only=False) as wdb:
wdb.conn.execute(
"DELETE FROM ema52_screening WHERE freq = ?",
[freq_label],
)
wdb.conn.executemany(
"""INSERT INTO ema52_screening
(trade_date, ts_code, name, freq, close_price, ema52, distance_pct, amount)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT (trade_date, ts_code, freq) DO UPDATE SET
close_price=excluded.close_price,
ema52=excluded.ema52,
distance_pct=excluded.distance_pct,
amount=excluded.amount,
name=excluded.name,
updated_at=now()""",
[(r["trade_date"], r["ts_code"], r["name"],
r["freq"], r["close_price"], r["ema52"], r["distance_pct"], r["amount"])
for r in results],
)
except Exception as e:
logger.error(f"EMA52: failed to save {freq_label}: {e}")
continue
logger.info(
f"EMA52 {freq_label}: {len(results)} near "
f"{threshold * 100:.1f}%)"
)
logger.info("EMA52 screening complete")
def _query_parquet(sql: str) -> "pd.DataFrame":
"""Execute a SQL query against Parquet files using a transient in-memory DuckDB."""
import duckdb
import pandas as pd
conn = duckdb.connect()
try:
return conn.execute(sql).fetchdf()
finally:
conn.close()
def _compute_ema(values: "np.ndarray", period: int) -> "np.ndarray":
"""Compute EMA for a 1D array of values."""
import numpy as np
if len(values) < period:
return np.array([])
alpha = 2.0 / (period + 1)
ema = np.zeros(len(values))
ema[period - 1] = np.mean(values[:period])
for i in range(period, len(values)):
ema[i] = alpha * values[i] + (1 - alpha) * ema[i - 1]
return ema
+5
View File
@@ -28,6 +28,11 @@ class Settings(BaseSettings):
# Realtime # Realtime
realtime_poll_interval: int = 5 realtime_poll_interval: int = 5
# EMA52 screening
ema52_threshold: float = 0.03
ema52_amount_min: float = 25.0 # 最低日成交额(亿元)
ema52_amount_max: float = 35.0 # 最高日成交额(亿元)
# Logging # Logging
log_level: str = "INFO" log_level: str = "INFO"
+216 -90
View File
@@ -9,9 +9,12 @@ from __future__ import annotations
import time import time
from concurrent.futures import ThreadPoolExecutor, as_completed from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import date, datetime, timedelta from datetime import date, datetime, timedelta
from pathlib import Path
from typing import Optional from typing import Optional
import pandas as pd import pandas as pd
import pyarrow as pa
import pyarrow.parquet as pq
from loguru import logger from loguru import logger
from ashare_dp.config import Settings from ashare_dp.config import Settings
@@ -23,7 +26,8 @@ from ashare_dp.core.models import (
) )
from ashare_dp.data.akshare_client import AKShareClient from ashare_dp.data.akshare_client import AKShareClient
from ashare_dp.storage.database import get_db from ashare_dp.storage.database import get_db
from ashare_dp.storage.repository import KLineRepository from ashare_dp.storage.repository import KLineRepository, _standardize_df
from ashare_dp.storage.partitioning import partition_glob
settings = Settings() settings = Settings()
@@ -74,6 +78,16 @@ LEGACY_MIN_COLUMN_MAPPING = {
REQUIRED_COLS = ["ts_code", "trade_time", "open", "high", "low", "close", "volume", "amount"] REQUIRED_COLS = ["ts_code", "trade_time", "open", "high", "low", "close", "volume", "amount"]
def _to_ts_code(symbol: str) -> str:
"""Convert 6-digit symbol to ts_code with exchange suffix."""
code = str(symbol).zfill(6)
if code.startswith(("4", "8")) or code.startswith("92"):
return f"{code}.BJ"
elif code.startswith("6") or code.startswith("9"):
return f"{code}.SH"
return f"{code}.SZ"
def _normalize_hist_df(df: pd.DataFrame, symbol: str, freq: Freq) -> pd.DataFrame: def _normalize_hist_df(df: pd.DataFrame, symbol: str, freq: Freq) -> pd.DataFrame:
"""Normalize daily K-line output to standard K-line schema.""" """Normalize daily K-line output to standard K-line schema."""
df = df.copy() df = df.copy()
@@ -89,7 +103,7 @@ def _normalize_hist_df(df: pd.DataFrame, symbol: str, freq: Freq) -> pd.DataFram
break break
if "ts_code" not in df.columns: if "ts_code" not in df.columns:
df["ts_code"] = symbol df["ts_code"] = _to_ts_code(symbol)
if "trade_time" in df.columns: if "trade_time" in df.columns:
df["trade_time"] = pd.to_datetime(df["trade_time"]) df["trade_time"] = pd.to_datetime(df["trade_time"])
@@ -121,7 +135,7 @@ def _normalize_min_df(df: pd.DataFrame, symbol: str) -> pd.DataFrame:
break break
if "ts_code" not in df.columns: if "ts_code" not in df.columns:
df["ts_code"] = symbol df["ts_code"] = _to_ts_code(symbol)
if "trade_time" in df.columns: if "trade_time" in df.columns:
df["trade_time"] = pd.to_datetime(df["trade_time"]) df["trade_time"] = pd.to_datetime(df["trade_time"])
@@ -175,6 +189,19 @@ def _resample_daily_to_period(df: pd.DataFrame, freq: Freq) -> pd.DataFrame:
return result return result
def _code_to_ts_code(code: str) -> str:
"""Convert bare 6-digit code to ts_code format: '000001' -> '000001.SZ'."""
code = str(code).zfill(6)
if code.startswith(("4", "8")) or code.startswith("92"):
return f"{code}.BJ"
elif code.startswith("6") or code.startswith("9"):
return f"{code}.SH"
return f"{code}.SZ"
WRITE_BATCH = 500 # Write to disk every N stocks to limit memory
class BackfillPipeline: class BackfillPipeline:
"""Orchestrates historical data backfill.""" """Orchestrates historical data backfill."""
@@ -187,18 +214,18 @@ class BackfillPipeline:
self.client = client or AKShareClient() self.client = client or AKShareClient()
self.repo = repo or KLineRepository() self.repo = repo or KLineRepository()
self.max_workers = max_workers or settings.backfill_workers self.max_workers = max_workers or settings.backfill_workers
self._db = get_db()
def init_db(self): def init_db(self):
"""Initialize database schema (tables and views).""" """Initialize database schema (tables and views)."""
from ashare_dp.storage.schema import DDL_STATEMENTS from ashare_dp.storage.schema import DDL_STATEMENTS
logger.info("Initializing database schema...") logger.info("Initializing database schema...")
for ddl in DDL_STATEMENTS: with get_db(read_only=False) as db:
try: for ddl in DDL_STATEMENTS:
self._db.execute(ddl) try:
except Exception as e: db.execute(ddl)
logger.warning(f"DDL warning: {e}") except Exception as e:
logger.warning(f"DDL warning: {e}")
logger.info("Database schema initialized") logger.info("Database schema initialized")
def load_stock_list(self) -> pd.DataFrame: def load_stock_list(self) -> pd.DataFrame:
@@ -230,21 +257,40 @@ class BackfillPipeline:
df["ts_code"] = df["symbol"].apply(_make_ts_code) df["ts_code"] = df["symbol"].apply(_make_ts_code)
df["exchange"] = df["ts_code"].str[-2:] df["exchange"] = df["ts_code"].str[-2:]
# Upsert into DuckDB # Upsert into DuckDB (batch insert, not row-by-row)
if "ts_code" in df.columns and "symbol" in df.columns: if "ts_code" in df.columns and "symbol" in df.columns:
for _, row in df.iterrows(): import duckdb
try: try:
self._db.execute(""" with get_db(read_only=False) as db:
INSERT OR REPLACE INTO stock_info (ts_code, symbol, name, exchange, updated_at) # Build column arrays for bulk insert
VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP) ts_codes = df["ts_code"].tolist()
""", ( symbols = df["symbol"].astype(str).tolist()
row.get("ts_code"), names = df["name"].astype(str).tolist()
str(row.get("symbol", "")), exchanges = df["exchange"].tolist()
str(row.get("name", "")), db.conn.execute("""
row.get("exchange", ""), INSERT OR REPLACE INTO stock_info (ts_code, symbol, name, exchange, updated_at)
)) SELECT ts_code, symbol, name, exchange, now()
except Exception: FROM (SELECT UNNEST($1::VARCHAR[]) AS ts_code,
pass UNNEST($2::VARCHAR[]) AS symbol,
UNNEST($3::VARCHAR[]) AS name,
UNNEST($4::VARCHAR[]) AS exchange)
""", [ts_codes, symbols, names, exchanges])
except Exception as e:
logger.warning(f"Batch stock_info insert failed ({e}), falling back to row-by-row")
with get_db(read_only=False) as db2:
for _, row in df.iterrows():
try:
db2.execute("""
INSERT OR REPLACE INTO stock_info (ts_code, symbol, name, exchange, updated_at)
VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)
""", (
row.get("ts_code"),
str(row.get("symbol", "")),
str(row.get("name", "")),
row.get("exchange", ""),
))
except Exception:
pass
logger.info(f"Loaded {len(df)} stocks into stock_info") logger.info(f"Loaded {len(df)} stocks into stock_info")
return df return df
@@ -260,20 +306,33 @@ class BackfillPipeline:
df["trade_date"] = pd.to_datetime(df["trade_date"]).dt.date df["trade_date"] = pd.to_datetime(df["trade_date"]).dt.date
for _, row in df.iterrows(): # Batch insert into DuckDB
d = row["trade_date"] dates = [d for d in df["trade_date"]]
try: weekdays = [d.weekday() for d in dates]
self._db.execute(""" years = [d.year for d in dates]
months = [d.month for d in dates]
try:
with get_db(read_only=False) as db:
db.conn.execute("""
INSERT OR REPLACE INTO trading_calendar (trade_date, is_trading_day, week_day, year, month) INSERT OR REPLACE INTO trading_calendar (trade_date, is_trading_day, week_day, year, month)
VALUES (?, 1, ?, ?, ?) SELECT trade_date, 1, week_day, year, month
""", ( FROM (SELECT UNNEST($1::DATE[]) AS trade_date,
d, UNNEST($2::INTEGER[]) AS week_day,
d.weekday(), UNNEST($3::INTEGER[]) AS year,
d.year, UNNEST($4::INTEGER[]) AS month)
d.month, """, [dates, weekdays, years, months])
)) except Exception as e:
except Exception: logger.warning(f"Batch trading_calendar insert failed ({e}), falling back to row-by-row")
pass with get_db(read_only=False) as db2:
for _, row in df.iterrows():
d = row["trade_date"]
try:
db2.execute("""
INSERT OR REPLACE INTO trading_calendar (trade_date, is_trading_day, week_day, year, month)
VALUES (?, 1, ?, ?, ?)
""", (d, d.weekday(), d.year, d.month))
except Exception:
pass
logger.info(f"Loaded {len(df)} trading days into trading_calendar") logger.info(f"Loaded {len(df)} trading days into trading_calendar")
@@ -288,6 +347,9 @@ class BackfillPipeline:
Daily data is fetched from Sina API. Weekly and monthly are Daily data is fetched from Sina API. Weekly and monthly are
derived by resampling the daily data. derived by resampling the daily data.
Data is written in batches (every WRITE_BATCH stocks) to avoid
holding all 5500+ stocks in memory at once.
Args: Args:
symbols: List of stock symbols (e.g. ['000001', '600000']). symbols: List of stock symbols (e.g. ['000001', '600000']).
If None, backfills all stocks from stock_info. If None, backfills all stocks from stock_info.
@@ -320,7 +382,27 @@ class BackfillPipeline:
logger.info("Backfilling 1d (daily) from Sina...") logger.info("Backfilling 1d (daily) from Sina...")
completed = 0 completed = 0
failed = 0 failed = 0
all_daily_frames = [] total_records = 0
batch_frames = []
temp_dir = Path(settings.parquet_dir) / ".tmp_daily"
temp_dir.mkdir(parents=True, exist_ok=True)
batch_idx = 0
def _flush_batch():
nonlocal total_records, batch_idx
if batch_frames:
chunk = pd.concat(batch_frames, ignore_index=True)
# Standardize: add freq, keep only standard columns
chunk["freq"] = Freq.d1.value
# Convert bare codes (000001) to ts_code format (000001.SZ)
chunk["ts_code"] = chunk["ts_code"].astype(str).apply(_code_to_ts_code)
chunk = _standardize_df(chunk)
total_records += len(chunk)
table = pa.Table.from_pandas(chunk, preserve_index=False)
pq.write_table(table, str(temp_dir / f"batch_{batch_idx:04d}.parquet"), compression="zstd", compression_level=3)
logger.debug(f" Flushed batch {batch_idx}: {len(chunk)} rows")
batch_idx += 1
batch_frames.clear()
def _backfill_daily(symbol: str): def _backfill_daily(symbol: str):
try: try:
@@ -348,22 +430,47 @@ class BackfillPipeline:
try: try:
df = future.result() df = future.result()
if df is not None and not df.empty: if df is not None and not df.empty:
all_daily_frames.append(df) batch_frames.append(df)
completed += 1 completed += 1
except Exception: except Exception:
failed += 1 failed += 1
if len(batch_frames) >= WRITE_BATCH:
_flush_batch()
if (completed + failed) % 100 == 0: if (completed + failed) % 100 == 0:
logger.info( logger.info(
f" 1d: {completed + failed}/{len(symbols)} " f" 1d: {completed + failed}/{len(symbols)} "
f"({completed} ok, {failed} fail)" f"({completed} ok, {failed} fail)"
) )
# Batch write all daily data at once _flush_batch()
total_records = 0
if all_daily_frames: # Merge all temp batches into final Hive-partitioned output
combined = pd.concat(all_daily_frames, ignore_index=True) if batch_idx > 0:
total_records = self.repo.write_klines(combined, Freq.d1) logger.info(f"Merging {batch_idx} temp batches into final parquet...")
t0 = time.monotonic()
temp_glob = str(temp_dir / "batch_*.parquet")
copy_conn = duckdb.connect()
copy_conn.execute(f"""
COPY (
SELECT *,
YEAR(trade_date) AS year,
MONTH(trade_date) AS month,
DAY(trade_date) AS day
FROM read_parquet('{temp_glob}', union_by_name=true)
ORDER BY trade_date, ts_code
) TO '{settings.parquet_dir}/kline_{Freq.d1.storage_dir}'
(FORMAT PARQUET, COMPRESSION ZSTD, COMPRESSION_LEVEL 3,
PARTITION_BY (year, month, day),
OVERWRITE_OR_IGNORE true, FILENAME_PATTERN 'data_{{i}}')
""")
copy_conn.close()
# Clean up temp files
import shutil
shutil.rmtree(temp_dir, ignore_errors=True)
elapsed = time.monotonic() - t0
logger.info(f" Merged in {elapsed:.1f}s")
results[Freq.d1.value] = { results[Freq.d1.value] = {
"records": total_records, "records": total_records,
@@ -393,56 +500,64 @@ class BackfillPipeline:
end_str: str, end_str: str,
results: dict, results: dict,
): ):
"""Derive weekly/monthly K-lines from stored daily data. """Derive weekly/monthly K-lines from daily data using DuckDB SQL.
Reads daily data for all symbols, resamples, and writes in batch Reads ALL daily parquet files in a single DuckDB scan,
to avoid file overwrite issues with the grouped-per-day write strategy. groups by ts_code + truncated date, and aggregates.
Much faster than per-stock queries — avoids 5500+ individual reads.
""" """
total_records = 0 if freq == Freq.w1:
completed = 0 trunc = "week"
failed = 0 elif freq == Freq.M1:
all_frames = [] trunc = "month"
else:
raise ValueError(f"Unsupported derive frequency: {freq}")
for symbol in symbols: parquet_glob = partition_glob(Freq.d1)
try:
df = self.repo.read_klines(
freq=Freq.d1,
ts_code=symbol,
start_date=date.fromisoformat(
pd.Timestamp(start_str).strftime("%Y-%m-%d")
) if len(start_str) == 8 else date.fromisoformat(start_str),
end_date=date.fromisoformat(
pd.Timestamp(end_str).strftime("%Y-%m-%d")
) if len(end_str) == 8 else date.fromisoformat(end_str),
limit=1_000_000,
)
if df is not None and not df.empty:
df = _resample_daily_to_period(df, freq)
all_frames.append(df)
completed += 1
except Exception as e:
logger.error(f"Failed derive {freq.value} for {symbol}: {e}")
failed += 1
if (completed + failed) % 500 == 0: sql = f"""
logger.info( SELECT
f" {freq.value}: {completed + failed}/{len(symbols)} " ts_code,
f"({completed} ok, {failed} fail)" date_trunc('{trunc}', trade_time) AS trade_time,
) CAST(date_trunc('{trunc}', trade_time) AS DATE) AS trade_date,
FIRST(open ORDER BY trade_time) AS open,
MAX(high) AS high,
MIN(low) AS low,
LAST(close ORDER BY trade_time) AS close,
SUM(volume) AS volume,
SUM(amount) AS amount
FROM read_parquet('{parquet_glob}',
hive_partitioning=true, union_by_name=true)
GROUP BY ts_code, date_trunc('{trunc}', trade_time)
ORDER BY ts_code, trade_time
"""
# Batch write all derived data at once logger.info(f"Deriving {freq.value} with single DuckDB scan...")
if all_frames: t0 = time.monotonic()
combined = pd.concat(all_frames, ignore_index=True) try:
total_records = self.repo.write_klines(combined, freq) derive_conn = duckdb.connect()
df = derive_conn.execute(sql).fetchdf()
derive_conn.close()
except Exception as e:
logger.error(f"DuckDB derive {freq.value} failed: {e}")
results[freq.value] = {"records": 0, "completed": 0, "failed": len(symbols)}
return
elapsed = time.monotonic() - t0
logger.info(f" Aggregated {len(df)} rows in {elapsed:.1f}s")
if df.empty:
results[freq.value] = {"records": 0, "completed": 0, "failed": 0}
return
total_records = self.repo.write_klines(df, freq)
results[freq.value] = { results[freq.value] = {
"records": total_records, "records": total_records,
"completed": completed, "completed": len(symbols),
"failed": failed, "failed": 0,
} }
logger.info( logger.info(
f" {freq.value} done: {total_records} records, " f" {freq.value} done: {total_records} records"
f"{completed} stocks ok, {failed} failed"
) )
def backfill_minute( def backfill_minute(
@@ -485,7 +600,15 @@ class BackfillPipeline:
completed = 0 completed = 0
failed = 0 failed = 0
all_min_frames = [] total_records = 0
batch_frames = []
def _write_min_batch():
nonlocal total_records
if batch_frames:
chunk = pd.concat(batch_frames, ignore_index=True)
total_records += self.repo.write_klines(chunk, freq)
batch_frames.clear()
def _backfill_min(symbol: str): def _backfill_min(symbol: str):
try: try:
@@ -512,22 +635,25 @@ class BackfillPipeline:
try: try:
df = future.result() df = future.result()
if df is not None and not df.empty: if df is not None and not df.empty:
all_min_frames.append(df) batch_frames.append(df)
completed += 1 completed += 1
except Exception: except Exception:
failed += 1 failed += 1
# Write batch every WRITE_BATCH stocks
if len(batch_frames) >= WRITE_BATCH:
_write_min_batch()
if (completed + failed) % 100 == 0: if (completed + failed) % 100 == 0:
logger.info( logger.info(
f" {freq.value}: {completed + failed}/{len(symbols)} " f" {freq.value}: {completed + failed}/{len(symbols)} "
f"({completed} ok, {failed} fail)" f"({completed} ok, {failed} fail)"
) )
# Batch write all data for this frequency at once # Write remaining for this frequency
total_records = 0 if batch_frames:
if all_min_frames: chunk = pd.concat(batch_frames, ignore_index=True)
combined = pd.concat(all_min_frames, ignore_index=True) total_records += self.repo.write_klines(chunk, freq)
total_records = self.repo.write_klines(combined, freq)
results[freq.value] = { results[freq.value] = {
"records": total_records, "records": total_records,
+144
View File
@@ -75,3 +75,147 @@ async def health_check_job():
# Import at bottom to avoid circular # Import at bottom to avoid circular
from ashare_dp.core.models import Freq from ashare_dp.core.models import Freq
async def ema52_screening_job():
"""EOD EMA52 screening: find stocks near EMA52 for 1d and 1w.
Triggered after EOD data pull. Scans all stocks, computes EMA52
from close prices, and saves results to ema52_screening table.
"""
from ashare_dp.config import Settings
import numpy as np
import pandas as pd
from ashare_dp.storage.database import get_db
settings = Settings()
threshold = settings.ema52_threshold
db = get_db()
# Get all ts_codes and names
stocks = db.query(
"SELECT ts_code, name FROM stock_info ORDER BY ts_code"
)
symbols = [row[0] for row in stocks]
symbol_names = {row[0]: row[1] for row in stocks}
logger.info(
f"EMA52 screening: {len(symbols)} stocks, "
f"threshold=±{threshold * 100:.1f}%"
)
for freq, freq_label, bars_needed in [
(Freq.d1, "1d", 100),
(Freq.w1, "1w", 60),
]:
try:
from ashare_dp.storage.partitioning import partition_glob
glob = partition_glob(freq)
# Get last N bars per stock using ROW_NUMBER()
sql = f"""
WITH ranked AS (
SELECT *, ROW_NUMBER() OVER (
PARTITION BY ts_code ORDER BY trade_time DESC
) as rn
FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true)
)
SELECT * FROM ranked WHERE rn <= {bars_needed}
ORDER BY ts_code, trade_time ASC
"""
df = db.conn.execute(sql).fetchdf()
if df.empty:
logger.warning(f"EMA52: no data for {freq_label}")
continue
except Exception as e:
logger.error(f"EMA52: failed to read {freq_label}: {e}")
continue
df["trade_time"] = pd.to_datetime(df["trade_time"])
df = df.sort_values(["ts_code", "trade_time"])
# Get last N bars per stock, compute EMA52
results = []
for ts_code, group in df.groupby("ts_code"):
if ts_code not in symbol_names:
continue
group = group.tail(bars_needed)
if len(group) < 26:
continue
closes = group["close"].astype(float).values
# Compute EMA52
ema_values = _compute_ema(closes, 52)
if len(ema_values) == 0:
continue
latest_close = closes[-1]
latest_ema = ema_values[-1]
if latest_ema <= 0:
continue
distance = (latest_close - latest_ema) / latest_ema
# Compute daily amount in 亿 CNY
latest_amount = float(group["amount"].iloc[-1])
amount_yi = latest_amount / 100_000_000 # 元 → 亿
if abs(distance) <= threshold:
results.append({
"trade_date": group["trade_date"].iloc[-1],
"ts_code": ts_code,
"name": symbol_names[ts_code],
"freq": freq_label,
"close_price": round(latest_close, 2),
"ema52": round(latest_ema, 2),
"distance_pct": round(distance * 100, 2),
"amount": round(amount_yi, 2),
})
if results:
# Sort by amount DESC (largest turnover first)
results.sort(key=lambda r: r.get("amount", 0) or 0, reverse=True)
# Upsert: delete old then insert
try:
db.conn.execute(
"DELETE FROM ema52_screening WHERE freq = ?",
[freq_label],
)
db.conn.executemany(
"""INSERT INTO ema52_screening
(trade_date, ts_code, name, freq, close_price, ema52, distance_pct, amount)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT (trade_date, ts_code, freq) DO UPDATE SET
close_price=excluded.close_price,
ema52=excluded.ema52,
distance_pct=excluded.distance_pct,
amount=excluded.amount,
name=excluded.name,
updated_at=now()""",
[(r["trade_date"], r["ts_code"], r["name"],
r["freq"], r["close_price"], r["ema52"], r["distance_pct"], r["amount"])
for r in results],
)
except Exception as e:
logger.error(f"EMA52: failed to save {freq_label}: {e}")
continue
near = [r for r in results if abs(r["distance_pct"]) <= threshold * 100]
logger.info(
f"EMA52 {freq_label}: {len(results)} near "
f"{threshold * 100:.1f}%), {len(near)} within threshold"
)
logger.info("EMA52 screening complete")
def _compute_ema(values: "np.ndarray", period: int) -> "np.ndarray":
"""Compute EMA for a 1D array of values."""
import numpy as np
if len(values) < period:
return np.array([])
alpha = 2.0 / (period + 1)
ema = np.zeros(len(values))
ema[period - 1] = np.mean(values[:period])
for i in range(period, len(values)):
ema[i] = alpha * values[i] + (1 - alpha) * ema[i - 1]
return ema
+16 -2
View File
@@ -19,7 +19,7 @@ class Scheduler:
def start(self): def start(self):
"""Start the scheduler and register jobs.""" """Start the scheduler and register jobs."""
from ashare_dp.scheduler.jobs import eod_pull_job, health_check_job from ashare_dp.scheduler.jobs import eod_pull_job, ema52_screening_job, health_check_job
# EOD job: 15:05 Beijing time, Mon-Fri # EOD job: 15:05 Beijing time, Mon-Fri
self._scheduler.add_job( self._scheduler.add_job(
@@ -48,8 +48,22 @@ class Scheduler:
replace_existing=True, replace_existing=True,
) )
# EMA52 screening: 15:10 Beijing time, Mon-Fri (after EOD pull)
self._scheduler.add_job(
ema52_screening_job,
trigger=CronTrigger(
day_of_week="mon-fri",
hour=15,
minute=10,
timezone=BEIJING_TZ,
),
id="ema52_screening",
name="EMA52 screening",
replace_existing=True,
)
self._scheduler.start() self._scheduler.start()
logger.info("Scheduler started with EOD (15:05 Mon-Fri) + health check (08:00 daily)") logger.info("Scheduler started with EOD (15:05) + EMA52 (15:10) + health check (08:00)")
def shutdown(self): def shutdown(self):
"""Shut down the scheduler.""" """Shut down the scheduler."""
+29 -26
View File
@@ -1,14 +1,15 @@
"""DuckDB connection management. """DuckDB connection management.
Uses a module-level connection. DuckDB connections are not thread-safe, Connections are short-lived: open, query, close. No persistent connection
so all operations are serialized through a threading lock for writes. so that the API process (read-only) and EOD process (write) can coexist
without lock contention.
""" """
from __future__ import annotations from __future__ import annotations
import threading from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Optional from typing import Iterator, Optional
import duckdb import duckdb
from loguru import logger from loguru import logger
@@ -17,16 +18,14 @@ from ashare_dp.config import Settings
settings = Settings() settings = Settings()
# Module-level lock for write serialization
_write_lock = threading.Lock()
class Database: class Database:
"""Manages a persistent DuckDB connection.""" """Short-lived DuckDB connection. Use as context manager."""
def __init__(self, db_path: str | None = None): def __init__(self, db_path: str | None = None, read_only: bool = False):
self.db_path = str(db_path or settings.duckdb_path) self.db_path = str(db_path or settings.duckdb_path)
self._conn: Optional[duckdb.DuckDBPyConnection] = None self._conn: Optional[duckdb.DuckDBPyConnection] = None
self._read_only = read_only
@property @property
def conn(self) -> duckdb.DuckDBPyConnection: def conn(self) -> duckdb.DuckDBPyConnection:
@@ -35,11 +34,14 @@ class Database:
self.connect() self.connect()
return self._conn return self._conn
@property
def read_only(self) -> bool:
return self._read_only
def connect(self) -> duckdb.DuckDBPyConnection: def connect(self) -> duckdb.DuckDBPyConnection:
"""Open a persistent connection to the DuckDB database.""" """Open a connection to the DuckDB database."""
Path(self.db_path).parent.mkdir(parents=True, exist_ok=True) Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)
self._conn = duckdb.connect(self.db_path) self._conn = duckdb.connect(self.db_path, read_only=self._read_only)
logger.info(f"Connected to DuckDB: {self.db_path}")
return self._conn return self._conn
def close(self): def close(self):
@@ -47,7 +49,6 @@ class Database:
if self._conn is not None: if self._conn is not None:
self._conn.close() self._conn.close()
self._conn = None self._conn = None
logger.info("DuckDB connection closed")
def execute(self, sql: str, params: tuple | None = None): def execute(self, sql: str, params: tuple | None = None):
"""Execute a SQL statement.""" """Execute a SQL statement."""
@@ -74,25 +75,27 @@ class Database:
).fetchone() ).fetchone()
return result[0] > 0 return result[0] > 0
def write_lock(self):
"""Acquire the write lock (use as context manager)."""
return _write_lock
def __enter__(self): def __enter__(self):
self.connect()
return self return self
def __exit__(self, *args): def __exit__(self, *args):
self.close() self.close()
# Global database instance @contextmanager
_db: Optional[Database] = None def get_db(read_only: bool = True) -> Iterator[Database]:
"""Get a short-lived DuckDB connection.
Opens, yields, and closes. No persistent singleton.
Use as context manager: `with get_db() as db: ...`
def get_db() -> Database: Default read_only=True for API queries.
"""Get or create the global database instance.""" Pass read_only=False for EOD/screening writes.
global _db """
if _db is None: db = Database(read_only=read_only)
_db = Database() try:
_db.connect() db.connect()
return _db yield db
finally:
db.close()
+93 -75
View File
@@ -6,6 +6,7 @@ from datetime import date, datetime
from pathlib import Path from pathlib import Path
from typing import Optional from typing import Optional
import duckdb
import pandas as pd import pandas as pd
import pyarrow as pa import pyarrow as pa
import pyarrow.parquet as pq import pyarrow.parquet as pq
@@ -46,7 +47,17 @@ class KLineRepository:
"""Read/write K-line data from/to Parquet files.""" """Read/write K-line data from/to Parquet files."""
def __init__(self): def __init__(self):
self._db = get_db() pass
def _query_parquet(self, sql: str, params: tuple | None = None) -> pd.DataFrame:
"""Execute a SQL query against Parquet files using a transient DuckDB."""
conn = duckdb.connect()
try:
if params:
return conn.execute(sql, params).fetchdf()
return conn.execute(sql).fetchdf()
finally:
conn.close()
# ---- Write ---- # ---- Write ----
@@ -81,31 +92,21 @@ class KLineRepository:
df = _standardize_df(df) df = _standardize_df(df)
# Determine partition date # Determine partition date
if "trade_date" not in df.columns and partition_date is None: if partition_date is None:
raise StorageError("DataFrame must have 'trade_date' column or partition_date must be provided") if "trade_date" not in df.columns:
raise StorageError("No trade_date column and no partition_date provided")
partition_date = pd.to_datetime(df["trade_date"].iloc[0]).date()
records_written = 0 # Ensure partition directory
pdir = ensure_partition_dir(freq, partition_date)
fpath = pdir / "data.parquet"
if partition_date is not None: # Write Parquet
# Single date: all records go to one file table = pa.Table.from_pandas(df)
partition_dir = ensure_partition_dir(freq, partition_date) pq.write_table(table, str(fpath), **PARQUET_WRITE_OPTIONS)
file_path = partition_dir / "data.parquet"
table = pa.Table.from_pandas(df, preserve_index=False)
pq.write_table(table, str(file_path), **PARQUET_WRITE_OPTIONS)
records_written = len(df)
else:
# Group by trade_date, one file per day
df["_pd"] = pd.to_datetime(df["trade_date"]).dt.date
for d, group in df.groupby("_pd"):
group = group.drop(columns=["_pd"])
partition_dir = ensure_partition_dir(freq, d)
file_path = partition_dir / "data.parquet"
table = pa.Table.from_pandas(group, preserve_index=False)
pq.write_table(table, str(file_path), **PARQUET_WRITE_OPTIONS)
records_written += len(group)
logger.debug(f"Wrote {records_written} records to kline_{freq.storage_dir}") logger.debug(f"Wrote {len(df)} records to {fpath}")
return records_written return len(df)
# ---- Read ---- # ---- Read ----
@@ -118,25 +119,27 @@ class KLineRepository:
limit: int = 10000, limit: int = 10000,
offset: int = 0, offset: int = 0,
) -> pd.DataFrame: ) -> pd.DataFrame:
"""Read K-line data from Parquet files.""" """Read K-line data from Parquet files.
Uses DuckDB read_parquet() with Hive partitioning.
For derived 2h freq, computes from 1h data.
"""
if freq == Freq.h2: if freq == Freq.h2:
return self._read_2h(ts_code, start_date, end_date, limit, offset) return self._read_2h(ts_code, start_date, end_date, limit, offset)
glob = partition_glob(freq) glob = partition_glob(freq)
parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}"
if not parquet_path.exists():
return pd.DataFrame(columns=STANDARD_COLS)
conditions = [] conditions = []
params = [] params = []
if ts_code: if ts_code:
conditions.append(f"ts_code = ${len(params) + 1}") conditions.append("ts_code = ?")
params.append(ts_code) params.append(ts_code)
if start_date: if start_date:
conditions.append(f"trade_date >= ${len(params) + 1}") conditions.append("trade_date >= ?")
params.append(start_date.isoformat()) params.append(start_date.isoformat())
if end_date: if end_date:
conditions.append(f"trade_date <= ${len(params) + 1}") conditions.append("trade_date <= ?")
params.append(end_date.isoformat()) params.append(end_date.isoformat())
where_clause = "" where_clause = ""
@@ -150,7 +153,7 @@ class KLineRepository:
LIMIT {limit} OFFSET {offset} LIMIT {limit} OFFSET {offset}
""" """
try: try:
return self._db.conn.execute(sql, params).fetchdf() return self._query_parquet(sql, tuple(params) if params else None)
except Exception as e: except Exception as e:
logger.warning(f"Query failed for freq={freq.value}: {e}") logger.warning(f"Query failed for freq={freq.value}: {e}")
return pd.DataFrame(columns=STANDARD_COLS) return pd.DataFrame(columns=STANDARD_COLS)
@@ -203,19 +206,16 @@ class KLineRepository:
freq: Freq, freq: Freq,
ts_code: str | None = None, ts_code: str | None = None,
) -> pd.DataFrame: ) -> pd.DataFrame:
"""Get the latest K-line data for the most recent trade date.""" """Get the latest K-line for each stock (or one stock).
if freq == Freq.h2:
return self._read_2h(ts_code, limit=100)
Uses DuckDB ROW_NUMBER() window function to get the last N bars per stock.
"""
glob = partition_glob(freq) glob = partition_glob(freq)
parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}"
if not parquet_path.exists():
return pd.DataFrame(columns=STANDARD_COLS)
conditions = [] conditions = []
params = [] params = []
if ts_code: if ts_code:
conditions.append(f"ts_code = ${len(params) + 1}") conditions.append("ts_code = ?")
params.append(ts_code) params.append(ts_code)
where_clause = "" where_clause = ""
@@ -223,49 +223,67 @@ class KLineRepository:
where_clause = "WHERE " + " AND ".join(conditions) where_clause = "WHERE " + " AND ".join(conditions)
sql = f""" sql = f"""
WITH latest AS ( WITH ranked AS (
SELECT MAX(trade_date) AS max_date SELECT *, ROW_NUMBER() OVER (
PARTITION BY ts_code ORDER BY trade_time DESC
) as rn
FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true)
{where_clause}
) )
SELECT k.* FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) k SELECT * FROM ranked WHERE rn = 1
JOIN latest ON k.trade_date = latest.max_date ORDER BY ts_code
{where_clause}
ORDER BY k.ts_code
""" """
try: return self._query_parquet(sql, tuple(params) if params else None)
return self._db.conn.execute(sql, params).fetchdf()
except Exception as e:
logger.warning(f"get_latest failed for freq={freq.value}: {e}")
return pd.DataFrame(columns=STANDARD_COLS)
# ---- Stats ---- def get_latest_bars(
self,
def get_date_range(self, freq: Freq) -> tuple[date | None, date | None]: freq: Freq,
"""Get the min and max trade_date for a given frequency.""" bars: int = 100,
parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}" ts_code: str | None = None,
if not parquet_path.exists(): ) -> pd.DataFrame:
return None, None """Get the last N bars per stock."""
glob = partition_glob(freq) glob = partition_glob(freq)
try: conditions = []
row = self._db.conn.execute(f""" params = []
SELECT MIN(trade_date), MAX(trade_date)
if ts_code:
conditions.append("ts_code = ?")
params.append(ts_code)
where_clause = ""
if conditions:
where_clause = "WHERE " + " AND ".join(conditions)
sql = f"""
WITH ranked AS (
SELECT *, ROW_NUMBER() OVER (
PARTITION BY ts_code ORDER BY trade_time DESC
) as rn
FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true)
""").fetchone() {where_clause}
return row[0], row[1] )
except Exception: SELECT * FROM ranked WHERE rn <= {bars}
return None, None ORDER BY ts_code, trade_time ASC
"""
return self._query_parquet(sql, tuple(params) if params else None)
def count_records(self, freq: Freq) -> int: def count_records(self, freq: Freq) -> int:
"""Count total records for a given frequency.""" """Count total records for a frequency."""
parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}"
if not parquet_path.exists():
return 0
glob = partition_glob(freq) glob = partition_glob(freq)
try: result = self._query_parquet(
row = self._db.conn.execute(f""" f"SELECT COUNT(*) FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true)"
SELECT count(*) )
return int(result.iloc[0, 0]) if not result.empty else 0
def get_date_range(self, freq: Freq) -> tuple[str, str] | None:
"""Get min/max trade_date for a frequency."""
glob = partition_glob(freq)
result = self._query_parquet(
f"""
SELECT MIN(trade_date) as min_date, MAX(trade_date) as max_date
FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true)
""").fetchone() """
return row[0] )
except Exception: if result.empty:
return 0 return None
return str(result.iloc[0, 0]), str(result.iloc[0, 1])
+25
View File
@@ -33,4 +33,29 @@ DDL_STATEMENTS = [
month TINYINT NOT NULL month TINYINT NOT NULL
) )
""", """,
"""
CREATE SEQUENCE IF NOT EXISTS seq_ema52_id
""",
"""
CREATE TABLE IF NOT EXISTS ema52_screening (
id BIGINT PRIMARY KEY DEFAULT nextval('seq_ema52_id'),
trade_date DATE NOT NULL,
ts_code VARCHAR(9) NOT NULL,
name VARCHAR(40),
freq VARCHAR(3) NOT NULL,
close_price DOUBLE NOT NULL,
ema52 DOUBLE NOT NULL,
distance_pct DOUBLE NOT NULL,
amount DOUBLE,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
UNIQUE(trade_date, ts_code, freq)
)
""",
"""
CREATE INDEX IF NOT EXISTS idx_ema52_date ON ema52_screening(trade_date)
""",
# Migration: add amount column if upgrading from older schema
"""
ALTER TABLE ema52_screening ADD COLUMN IF NOT EXISTS amount DOUBLE
""",
] ]