"""全市场数据同步(未复权,写入 candles 全量底座)。 设计:trade_cal 取近 N 个交易日 -> 逐日 pro.daily(trade_date=...) 一次返回全市场当日数据 -> upsert 进 candles(不复权底座,ON CONFLICT 幂等);daily_basic 仅同步最新交易日到 DailySnapshot(市值/PE/PB/换手率等截面字段)。 同步为进程内后台任务(MVP 不引入任务队列),前端轮询 /api/screener/sync/status。 daily 与 daily_basic 分步独立落库:daily_basic 积分不足时快照仍可用,错误写入状态不中断任务。 """ from __future__ import annotations import asyncio import time from datetime import datetime, timedelta from sqlalchemy import delete, func, insert, select from sqlalchemy.dialects.postgresql import insert as pg_insert from sqlalchemy.ext.asyncio import AsyncSession from .. import cache from ..config import settings from ..data.symbols import plain_code from ..models import AdjFactor, Candle, DailySnapshot, StockBasic, TradeCalendar from .llm import ScreenerError # 进程内单例任务状态(uvicorn --reload 单进程场景够用) _sync_state: dict = { "running": False, "step": None, "total_days": 0, "done_days": 0, "error": None, "started_at": None, "finished_at": None, } _sync_task: asyncio.Task | None = None _sync_lock = asyncio.Lock() _BATCH = 5000 # executemany 分批行数 # Tushare 积分/权限不足的特征文案(daily_basic 常见门槛) _PERM_MARKS = ("抱歉,您没有访问该项目权限", "积分", "权限") # 频率超限特征(等待 62s 重试一次) _RATE_MARKS = ("频率超限", "每分钟") def _call_retry(fn, *args, **kwargs): """同步调用 tushare 接口;「每分钟」级频率超限等 62s 重试一次(小时级限频直接抛)。""" try: return fn(*args, **kwargs) except Exception as e: # noqa: BLE001 msg = str(e) if any(m in msg for m in _RATE_MARKS) and "小时" not in msg: time.sleep(62) return fn(*args, **kwargs) raise def _get_pro(): """token 检查 + 返回 pro api 客户端(同步对象,调用需 to_thread 包裹)。 经 tushare_provider.get_pro 统一走镜像补丁(15000 积分档 token 只认 quicksync)。 """ if not settings.tushare_token: raise ScreenerError("未配置 TUSHARE_TOKEN,无法同步全市场数据(backend/.env)") from ..data.tushare_provider import get_pro return get_pro() def _parse_d(s: str) -> datetime: return datetime.strptime(str(s), "%Y%m%d") def _fetch_calendar_sync(pro) -> list[str]: """拉取宽范围交易日历(近 18 个月 + 未来 3 个月),返回 YYYYMMDD 列表。""" time.sleep(settings.screener_sync_interval) end = (datetime.now() + timedelta(days=90)).strftime("%Y%m%d") start = (datetime.now() - timedelta(days=550)).strftime("%Y%m%d") cal = _call_retry(pro.trade_cal, exchange="SSE", start_date=start, end_date=end, is_open="1") return sorted(cal["cal_date"].tolist()) async def _recent_trade_dates(session: AsyncSession, pro, days: int) -> list[str]: """近 N 个交易日(YYYYMMDD,倒序)。日历本地缓存,仅在覆盖不到当天时刷新一次。 trade_cal 低积分版限频 1 次/小时:刷新被限频时沿用缓存(日历略旧无害—— daily 对未生成日期返回空,同步会自然跳过)。 """ cached = (await session.execute(select(TradeCalendar.trade_date).order_by(TradeCalendar.trade_date.desc()))).scalars().all() today = datetime.now().strftime("%Y%m%d") have_today = bool(cached) and cached[0] >= today if not have_today: try: dates = await asyncio.to_thread(_fetch_calendar_sync, pro) await session.execute(delete(TradeCalendar)) await session.execute(insert(TradeCalendar), [{"trade_date": d} for d in dates]) await session.commit() cached = dates[::-1] except Exception as e: # noqa: BLE001 —— 限频且无缓存时才致命 if not cached: raise ScreenerError(f"获取交易日历失败(且本地无缓存): {str(e)[:150]}") from e _sync_state["step"] = "交易日历刷新受限,沿用本地缓存" recent = [d for d in cached if d <= today][:days] if not recent: raise ScreenerError("交易日历为空") return recent def _fetch_daily(pro, d: str) -> list[dict]: """拉取某交易日全市场日线(未复权)。当日数据未生成(盘前/盘中)返回空。""" time.sleep(settings.screener_sync_interval) df = _call_retry(pro.daily, trade_date=d) if df is None or df.empty: return [] rows = [] for _, r in df.iterrows(): rows.append({ "trade_date": _parse_d(d), "ts_code": r["ts_code"], "open": float(r["open"]), "high": float(r["high"]), "low": float(r["low"]), "close": float(r["close"]), "pre_close": float(r["pre_close"]), "change": None if r.get("change") != r.get("change") else float(r["change"]), "pct_chg": None if r.get("pct_chg") != r.get("pct_chg") else float(r["pct_chg"]), "vol": float(r["vol"]), # 手 "amount": float(r["amount"]), # 千元 }) return rows def _fetch_basic(pro, d: str) -> list[dict]: """拉取某交易日每日指标快照(daily_basic,低积分版限频 1 次/分钟)。 失败(积分不足等)时记录错误返回空,不拖垮日线同步。 """ time.sleep(settings.screener_sync_interval) try: df = _call_retry(pro.daily_basic, trade_date=d) except Exception as e: # noqa: BLE001 msg = str(e) if any(m in msg for m in _PERM_MARKS): _sync_state["error"] = ( f"Tushare 无法获取每日指标(daily_basic):{msg[:150]}。" "市值/市盈率等条件不可用;纯指标选股不受影响。" ) return [] raise if df is None or df.empty: return [] rows = [] for _, r in df.iterrows(): def _f(key: str) -> float | None: v = r.get(key) return None if v is None or v != v else float(v) rows.append({ "trade_date": _parse_d(d), "ts_code": r["ts_code"], "close": _f("close"), "turnover_rate": _f("turnover_rate"), "turnover_rate_f": _f("turnover_rate_f"), "volume_ratio": _f("volume_ratio"), "pe": _f("pe"), "pe_ttm": _f("pe_ttm"), "pb": _f("pb"), "total_mv": _f("total_mv"), "circ_mv": _f("circ_mv"), # 万元 }) return rows def _fetch_adj_factor(pro, d: str) -> list[dict]: """拉取某交易日全市场复权因子(K线 bfq->qfq/hfq 本地换算的底座)。""" time.sleep(settings.screener_sync_interval) df = _call_retry(pro.adj_factor, trade_date=d) if df is None or df.empty: return [] return [ {"trade_date": _parse_d(d), "ts_code": r["ts_code"], "adj_factor": float(r["adj_factor"])} for _, r in df.iterrows() ] def _sync_stock_list_sync(pro) -> list[dict]: """拉取在市股票列表。""" time.sleep(settings.screener_sync_interval) df = _call_retry(pro.stock_basic, exchange="", list_status="L", fields="ts_code,symbol,name,area,industry,market,exchange,list_status,list_date,delist_date") rows = [] for _, r in df.iterrows(): rows.append({ "ts_code": r["ts_code"], "symbol": r["symbol"], "name": r["name"], "area": r.get("area") or None, "industry": r.get("industry") or None, "market": r.get("market") or None, "exchange": r["exchange"] or "", "list_status": r["list_status"], "list_date": r.get("list_date") or "", "delist_date": r.get("delist_date") or None, }) return rows def _norm_date(v) -> str: """把 DB 读出的 trade_date(可能是 datetime 或 str)归一为 YYYYMMDD。""" if hasattr(v, "strftime"): return v.strftime("%Y%m%d") return str(v)[:10].replace("-", "") async def _existing_dates(session: AsyncSession, model) -> set[str]: """某表已落库的交易日集合(YYYYMMDD 字符串,便于比对)。""" res = await session.execute(select(func.distinct(model.trade_date))) return {_norm_date(r[0]) for r in res} async def _replace_day(session: AsyncSession, model, rows: list[dict], d_str: str) -> None: """按交易日删旧插新(幂等),executemany 分批。""" d = _parse_d(d_str) await session.execute(delete(model).where(model.trade_date == d)) for i in range(0, len(rows), _BATCH): await session.execute(insert(model), rows[i : i + _BATCH]) await session.commit() async def _existing_candle_dates(session: AsyncSession) -> set[str]: """candles 表已落库的交易日集合(YYYYMMDD 字符串,便于比对)。""" res = await session.execute( select(func.distinct(func.date(Candle.ts))).where(Candle.timeframe == "1d") ) return {r[0].strftime("%Y%m%d") for r in res if r[0] is not None} async def _upsert_candle_day(session: AsyncSession, rows: list[dict], listed: set[str], d_str: str) -> None: """把某交易日全市场日线 upsert 进 candles(不复权底座,幂等)。 rows 来自 _fetch_daily(ts_code/vol手/amount千元);只写 stock_basic 在市股票, 与 TDX 底座口径一致;amount 已有(TDX 回补)时保留旧值。 """ batch = [ { "symbol": plain_code(r["ts_code"]), "timeframe": "1d", "ts": _parse_d(d_str), "open": r["open"], "high": r["high"], "low": r["low"], "close": r["close"], "volume": r["vol"] * 100.0, # 手 -> 股 "amount": (r["amount"] * 1000.0) if r["amount"] is not None else None, # 千元 -> 元 "turnover": None, # 换手率另由 daily_basic 快照维护 } for r in rows if plain_code(r["ts_code"]) in listed ] if not batch: return # on_conflict 语句整批渲染为占位符(非 executemany),asyncpg 单语句参数上限 32766, # 10 列 x 3000 行 = 30000 参数留出余量 for i in range(0, len(batch), 3000): stmt = pg_insert(Candle).values(batch[i : i + 3000]) stmt = stmt.on_conflict_do_update( index_elements=["symbol", "timeframe", "ts"], set_={ "open": stmt.excluded.open, "high": stmt.excluded.high, "low": stmt.excluded.low, "close": stmt.excluded.close, "volume": stmt.excluded.volume, "amount": func.coalesce(Candle.amount, stmt.excluded.amount), }, ) await session.execute(stmt) await session.commit() async def _run_sync(days: int, force: bool) -> None: """后台任务主体:stock_basic -> 逐日日线 -> 最新交易日快照。异常写状态。 daily_basic 只拉最新交易日(快照条件仅作用于最新截面,且低积分 token 限频 1 次/分钟)。 """ from ..db import async_session # 延迟导入避免循环 try: pro = await asyncio.to_thread(_get_pro) # 1) 股票列表(已有数据则跳过——stock_basic 低积分版限频 1 次/小时) async with async_session() as session: stocks_now = int(await session.scalar(select(func.count()).select_from(StockBasic)) or 0) if stocks_now == 0 or force: _sync_state["step"] = "正在同步股票列表" try: rows = await asyncio.to_thread(_sync_stock_list_sync, pro) async with async_session() as session: await session.execute(delete(StockBasic)) for i in range(0, len(rows), _BATCH): await session.execute(insert(StockBasic), rows[i : i + _BATCH]) await session.commit() except Exception as e: # noqa: BLE001 —— 受限时沿用现有列表继续 if stocks_now > 0: _sync_state["step"] = f"股票列表同步受限(沿用现有 {stocks_now} 只)" else: raise # 2) 逐交易日全市场日线 -> candles(增量;当日未生成则跳过) async with async_session() as session: dates = await _recent_trade_dates(session, pro, days) have_daily = set() if force else await _existing_candle_dates(session) # 在市股票集合,限定写入范围(与 TDX 底座口径一致) listed = set( (await session.execute( select(StockBasic.symbol).where(StockBasic.list_status == "L") )).scalars() ) todo = [d for d in dates if d not in have_daily] _sync_state["total_days"] = len(todo) _sync_state["done_days"] = 0 for d in todo: _sync_state["step"] = f"正在同步 {d} 日线({_sync_state['done_days'] + 1}/{len(todo)})" daily_rows = await asyncio.to_thread(_fetch_daily, pro, d) if daily_rows: # 盘前/盘中等未生成数据的日期直接跳过 async with async_session() as session: await _upsert_candle_day(session, daily_rows, listed, d) _sync_state["done_days"] += 1 # 2.5) 复权因子(与日线同窗口增量;历史全量由 scripts/backfill_adj_factor.py 回补) async with async_session() as session: have_adj = set() if force else await _existing_dates(session, AdjFactor) for d in [d for d in dates if d not in have_adj]: _sync_state["step"] = f"正在同步 {d} 复权因子" adj_rows = await asyncio.to_thread(_fetch_adj_factor, pro, d) if adj_rows: async with async_session() as session: await _replace_day(session, AdjFactor, adj_rows, d) # 3) 最新「有数据」交易日的快照(daily_basic,仅 1 次调用) # 用 candles 实际最大交易日(今天的数据收盘后才生成,日历最新日会拉到空) async with async_session() as session: latest_dt = await session.scalar(select(func.max(Candle.ts))) latest = latest_dt.strftime("%Y%m%d") if latest_dt else None if latest: async with async_session() as session: have_snap = force or latest not in await _existing_dates(session, DailySnapshot) if have_snap: _sync_state["step"] = f"正在同步 {latest} 每日指标" basic_rows = await asyncio.to_thread(_fetch_basic, pro, latest) if basic_rows: async with async_session() as session: await _replace_day(session, DailySnapshot, basic_rows, latest) # candles/复权因子已更新:作废旧 K 线预览缓存(键含版本号,自增即全体失效) await cache.bump_version("candles") _sync_state["step"] = "同步完成" except Exception as e: # noqa: BLE001 _sync_state["error"] = f"同步失败:{str(e)[:300]}" _sync_state["step"] = "同步失败" finally: _sync_state["running"] = False _sync_state["finished_at"] = datetime.now() async def start_sync(session: AsyncSession, days: int, force: bool) -> dict: """幂等启动后台同步任务;已在跑则直接返回当前状态。""" global _sync_task async with _sync_lock: if _sync_state["running"] and _sync_task and not _sync_task.done(): return dict(_sync_state) _sync_state.update({ "running": True, "step": "准备同步", "total_days": days, "done_days": 0, "error": None, "started_at": datetime.now(), "finished_at": None, }) _sync_task = asyncio.create_task(_run_sync(days, force)) return dict(_sync_state) # candles 是千万行表,count 较重;前端每 2s 轮询状态,需 TTL 缓存降载 _status_stats_cache: dict = {"at": 0.0, "data": None} _STATS_TTL = 30.0 async def _db_stats(session: AsyncSession) -> dict: """candles/快照/股票列表实况(30s TTL 缓存)。""" now = time.time() if _status_stats_cache["data"] is not None and now - _status_stats_cache["at"] < _STATS_TTL: return _status_stats_cache["data"] stocks = int(await session.scalar(select(func.count()).select_from(StockBasic)) or 0) daily_rows = int(await session.scalar(select(func.count()).select_from(Candle)) or 0) snap_rows = int(await session.scalar(select(func.count()).select_from(DailySnapshot)) or 0) last_daily = await session.scalar( select(func.max(Candle.ts)).where(Candle.timeframe == "1d") ) n_dates = int(await session.scalar( select(func.count(func.distinct(func.date(Candle.ts)))).where(Candle.timeframe == "1d") ) or 0) data = { "stocks": stocks, "daily_rows": daily_rows, "snapshot_rows": snap_rows, "last_daily": last_daily, "dates": n_dates, } _status_stats_cache.update(at=now, data=data) return data async def get_sync_status(session: AsyncSession) -> dict: """合并任务状态 + DB 实况(最新交易日/行数/ready 标志),与 ScreenerSyncStatus DTO 对齐。""" stats = await _db_stats(session) status = dict(_sync_state) status.update({ "stats": {"stocks": stats["stocks"], "daily_rows": stats["daily_rows"], "snapshot_rows": stats["snapshot_rows"], "dates": stats["dates"]}, "last_trade_date": stats["last_daily"], "last_synced_at": _sync_state.get("finished_at") or _sync_state.get("started_at"), "ready": stats["daily_rows"] > 0, }) return status