"""智能选股域路由:自然语言选股(NDJSON 流式)+ 提问历史 + 全市场同步 + 个股预览。""" from __future__ import annotations import asyncio import json from datetime import datetime import pandas as pd from fastapi import APIRouter, Depends, HTTPException, Response from fastapi.responses import StreamingResponse from sqlalchemy import func, select, text from sqlalchemy.ext.asyncio import AsyncSession from .. import cache from .. import indicators as ind from ..auth import require_user from ..config import settings from ..data import fetcher, repository from ..data.aggregation import resample_bars from ..data.symbols import plain_code from ..db import async_session, get_session from ..models import AdjFactor, Candle, ScreenerQuery from ..schemas import ( CandleOut, PreviewInfoOut, PreviewResponse, ScreenerQueryListResponse, ScreenerQueryOut, ScreenerRunRequest, ScreenerRunResponse, ScreenerSyncRequest, ScreenerSyncStatus, ) from ..screener import engine, market_sync from ..screener.engine import DataNotReadyError from ..screener.llm import ScreenerError, parse_conditions from ._deps import ( ADJUST_MODES, FACTOR_STEP_SQL, FULL_MA_SET, INFO_SQL, adjust_bars, cached_json_response, raw_json, rows_to_bars, series_to_jsonable, ) router = APIRouter() # ---------- 智能选股 ---------- @router.post("/screener/run") async def screener_run( req: ScreenerRunRequest, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> StreamingResponse: """自然语言 -> LLM 解析条件 -> 全市场筛选。也可直传 conditions 跳过 LLM(微调再跑)。 NDJSON 流式响应(每行一个 JSON 事件,前端逐行渲染进度): {"type":"stage","key":"llm|date|prefilter|bars|filter_done|done","msg":"…","ms":123} {"type":"parsed","conditions":{…},"ms":456} LLM 解析出的结构化条件 {"type":"candidates","count":5400,"msg":"…","ms":…} SQL 预筛后的候选数 {"type":"progress","done":500,"total":5400} 逐股指标过滤进度 {"type":"result","result":{…ScreenerRunResponse…},"ms":…} {"type":"error","message":"…","code":400} 流中途失败(HTTP 已 200) 成功的提问(含解析出的条件与命中数)记录到 screener_queries,供历史一键重跑。 """ limit = settings.screener_default_limit async def gen(): try: if req.conditions: conds = req.conditions else: yield _ndjson({"type": "stage", "key": "llm", "msg": f"AI 解析条件中({settings.llm_model})…"}) conds = await parse_conditions(req.text) if not conds.indicator and not conds.snapshot: yield _ndjson({"type": "error", "code": 400, "message": "AI 未从描述中解析出任何筛选条件,请换种说法"}) return yield _ndjson({"type": "parsed", "conditions": conds.model_dump()}) result = None async for ev in engine.run_screen_events(session, conds, limit): if ev["type"] == "result": result = ev["result"] yield _ndjson({"type": "stage", "key": "done", "ms": ev.get("ms"), "msg": f"筛选完成:{result['total']} 只命中(数据基准 {result['trade_date']:%Y-%m-%d})"}) else: yield _ndjson(ev) if result is None: yield _ndjson({"type": "error", "code": 500, "message": "选股流程未产出结果"}) return yield _ndjson({"type": "result", "result": ScreenerRunResponse(**result).model_dump(mode="json")}) # 相同文本 + 相同条件的上一条不重复记录(一键重跑场景) exists = ( await session.execute( select(ScreenerQuery.id).where( ScreenerQuery.user_id == user.id, ScreenerQuery.text == req.text.strip(), ScreenerQuery.conditions_json == json.dumps(conds.model_dump(), ensure_ascii=False), ) ) ).scalar_one_or_none() if exists is None: session.add(ScreenerQuery( user_id=user.id, text=req.text.strip(), conditions_json=json.dumps(conds.model_dump(), ensure_ascii=False), hit_count=result.get("total", 0), )) await session.commit() except DataNotReadyError as e: yield _ndjson({"type": "error", "code": 409, "message": str(e)}) except ValueError as e: # 未知指标/字段、条件为空 yield _ndjson({"type": "error", "code": 400, "message": str(e)}) except ScreenerError as e: code = 503 if "未配置 LLM_API_KEY" in str(e) else 502 yield _ndjson({"type": "error", "code": code, "message": str(e)}) except Exception as e: # noqa: BLE001 yield _ndjson({"type": "error", "code": 500, "message": f"选股失败: {e}"}) return StreamingResponse(gen(), media_type="application/x-ndjson", headers={"Cache-Control": "no-store", "X-Accel-Buffering": "no"}) def _ndjson(obj: dict) -> str: """dict -> NDJSON 行(json.dumps 保证 default=str 兜底 datetime 等)。""" return json.dumps(obj, ensure_ascii=False, default=str) + "\n" @router.get("/screener/queries", response_model=ScreenerQueryListResponse) async def screener_queries( limit: int = 20, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> ScreenerQueryListResponse: """当前用户的提问历史(最新在前,含解析出的条件与命中数,可一键重跑)。""" limit = max(1, min(limit, 100)) rows = ( await session.execute( select(ScreenerQuery) .where(ScreenerQuery.user_id == user.id) .order_by(ScreenerQuery.created_at.desc()) .limit(limit) ) ).scalars().all() items = [] for r in rows: conds = None if r.conditions_json: try: from ..schemas import ScreenConditions conds = ScreenConditions.model_validate_json(r.conditions_json) except Exception: # noqa: BLE001 —— 旧格式/解析失败则只展示文本 conds = None items.append(ScreenerQueryOut( id=r.id, text=r.text, conditions=conds, hit_count=r.hit_count, created_at=r.created_at )) return ScreenerQueryListResponse(items=items) @router.delete("/screener/queries/{query_id}", status_code=204) async def screener_query_delete( query_id: int, session: AsyncSession = Depends(get_session), user=Depends(require_user), ) -> None: await session.execute( text("DELETE FROM screener_queries WHERE id = :i AND user_id = :u"), {"i": query_id, "u": user.id}, ) await session.commit() # ---------- 全市场数据同步 ---------- @router.post("/screener/sync", response_model=ScreenerSyncStatus) async def screener_sync_start( req: ScreenerSyncRequest, session: AsyncSession = Depends(get_session) ) -> ScreenerSyncStatus: """启动全市场数据同步(后台任务,立即返回状态)。""" try: await market_sync.start_sync(session, req.days, req.force) except ScreenerError as e: raise HTTPException(status_code=503, detail=str(e)) status = await market_sync.get_sync_status(session) return ScreenerSyncStatus(**{k: status.get(k) for k in ScreenerSyncStatus.model_fields}) @router.get("/screener/sync/status", response_model=ScreenerSyncStatus) async def screener_sync_status(session: AsyncSession = Depends(get_session)) -> ScreenerSyncStatus: """同步任务状态 + 数据实况(最新交易日/行数/ready)。""" status = await market_sync.get_sync_status(session) return ScreenerSyncStatus(**{k: status.get(k) for k in ScreenerSyncStatus.model_fields}) # ---------- 个股详情预览 ---------- @router.get("/screener/preview/{ts_code}", response_model=PreviewResponse) async def screener_preview( ts_code: str, limit: int = 500, adjust: str = "qfq", timeframe: str = "1d", mas: str = "5,10,20,60", zx: str = "10,20,30,60", end: str | None = None, session: AsyncSession = Depends(get_session), ) -> Response: """个股详情预览:日线(candles 不复权底座 + adj_factor 本地换算 bfq/qfq/hfq, 未缓存自动拉取,落后全市场最新交易日则强制刷新)+ 全套指标 + 最新截面信息卡。 timeframe 聚合到周/月/年(先复权再聚合);mas 指定主图 MA 周期(逗号分隔)。 end=YYYY-MM-DD 时为「向前翻页」:返回该日之前最近 limit 根(含预热计算指标), has_more 标记窗口前是否还有更早历史,前端据此继续向左滚动加载。""" if adjust not in ADJUST_MODES: raise HTTPException(status_code=400, detail=f"adjust 仅支持 {'/'.join(ADJUST_MODES)}") if timeframe not in ("1d", "1w", "1M", "1y"): raise HTTPException(status_code=400, detail="timeframe 仅支持 1d/1w/1M/1y") try: ma_periods = sorted({int(p) for p in mas.split(",") if p.strip().isdigit() and 1 <= int(p) <= 500}) except ValueError: raise HTTPException(status_code=400, detail="mas 格式应为逗号分隔的数字,如 5,10,20,60") if not ma_periods: ma_periods = [5, 10, 20, 60] try: zx_periods = sorted({int(p) for p in zx.split(",") if p.strip().isdigit() and 1 <= int(p) <= 500}) except ValueError: raise HTTPException(status_code=400, detail="zx 格式应为逗号分隔的数字,如 10,20,30,60") if not zx_periods: zx_periods = [10, 20, 30, 60] limit = max(30, min(limit, 5000)) end_dt: datetime | None = None if end: try: end_dt = datetime.strptime(end.strip()[:10], "%Y-%m-%d") except ValueError: raise HTTPException(status_code=400, detail="end 格式应为 YYYY-MM-DD") symbol = plain_code(ts_code) # --- 两级读缓存:历史窗口(end 翻页)只增不改,最新窗口每日由全市场同步推进; # 键含 ver:candles 版本号(同步完成后自增,旧缓存全部失效),TTL 兜底(cache.py)。 # 存序列化好的 JSON 直返(j: 前缀),跳过 json.loads + pydantic 校验/序列化(热路径数百 ms → 个位数)。 # 注:ma_periods 不参与缓存键 —— 前端已改为本地计算 MA,后端始终返回全量 MA 集合 cache_key = cache.digest( "preview", ts_code, timeframe, limit, adjust, end_dt.strftime("%Y-%m-%d") if end_dt else None, await cache.get_version("candles"), ) cached = await cached_json_response(f"pvj:{cache_key}") if cached is not None: return cached # --- 日线:candles(全量不复权底座);未缓存拉取,落后于全市场最新交易日则强制刷新 --- # fetcher 只做「不复权」增量 upsert,底座口径恒为 bfq(TDX 全量 + Tushare 增量), # 复权(qfq/hfq)读取时按 adj_factor 表本地换算。 # 每次只取「窗口 + 400 根预热」行(MA250/MACD EMA 在 400 根内充分收敛),不拉全量: # 首屏 ~500 根秒开,向左滚动时按 end 参数逐页向前翻。 frame_mult = {"1d": 1, "1w": 6, "1M": 24, "1y": 280}[timeframe] fetch_n = min(100000, limit * frame_mult + 400) source = "bfq" mode = "bfq" # 并发约定:注入 session 与 s2 各占一条连接,每次 gather 里每个 session 恰好跑一条查询 # (AsyncSession 单连接非并发安全),把 ~6 次串行 DB RTT 折叠成 2 个波次。 async with async_session() as s2: if end_dt is not None: # 向前翻页:取 end 之前的历史窗口,不触发同步(历史浏览);max(ts) 用不到 rows = await repository.get_candles_before(session, symbol, "1d", before=end_dt, limit=fetch_n) global_latest = None else: # Wave 1:candles 窗口(注入 session)+ 全市场最新交易日(s2)并行 rows, global_latest = await asyncio.gather( repository.get_recent_candles(session, symbol, "1d", limit=fetch_n), s2.scalar(select(func.max(Candle.ts)).where(Candle.timeframe == "1d")), ) try: if not rows: await fetcher.sync_symbol(session, symbol, source="auto") rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n) elif global_latest is not None and rows[-1].ts.date() < global_latest.date(): await fetcher.sync_symbol(session, symbol, source="auto", force=True) rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n) except Exception: # noqa: BLE001 —— tushare/写库失败时回滚会话(否则毒化后兜底查询 500) await session.rollback() if not rows: rows = [] bars = rows_to_bars(rows) if not bars and end_dt is None: raise HTTPException(status_code=404, detail=f"无数据: {ts_code}(可先点「同步市场数据」)") # 信息卡取未聚合的日线最新 bar(聚合后 ts 是周期起点,不适用于「最新交易日」) last_daily = bars[-1] if bars else None prev_daily = bars[-2] if len(bars) > 1 else None # 翻页到底(end 之前无数据):返回空页 + has_more=False,前端停止向前翻页 # --- Wave 2:复权因子(s2,覆盖索引 Index Only Scan)+ 信息卡(注入 session,LATERAL 一条)并行 --- async def _fetch_factors() -> list | None: if adjust == mode or not bars: return None # 只取因子「变化点」行(覆盖索引 Index Only Scan,免堆访问——adj_factor 堆碎片化 # 严重);bisect 在阶梯函数上取值与日级序列逐字节一致 if end_dt is not None: # 分页:窗口 ≤ end 的变化点 + 全局最新因子(qfq 以最新因子归一) win = list((await s2.execute(FACTOR_STEP_SQL, {"code": ts_code, "upto": end_dt})).all()) if win: latest_f = (await s2.execute( select(AdjFactor.trade_date, AdjFactor.adj_factor) .where(AdjFactor.ts_code == ts_code) .order_by(AdjFactor.trade_date.desc()).limit(1) )).first() if latest_f is not None: win.append(latest_f) return win or None # 非分页:上界 global_latest(≥ 最新 bar),末项变化点即全局最新因子,比 # 「窗口 ≤ bars[-1].ts + 单独 latest」少一次查询 return list((await s2.execute( FACTOR_STEP_SQL, {"code": ts_code, "upto": global_latest} )).all()) or None factors, info_row = await asyncio.gather( _fetch_factors(), session.execute(INFO_SQL, {"code": ts_code, "target": last_daily.ts if last_daily else None}), ) # --- 复权换算:请求模式与底座模式不同时按 adj_factor 本地换算(无因子则维持原样) --- if factors: bars = adjust_bars(bars, factors, mode, adjust) mode = adjust source = adjust # --- 周期聚合:复权之后按日历聚合到周/月/年,指标在聚合后的序列上计算 --- bars = resample_bars(bars, timeframe) # --- 指标(在预热窗口上计算后截尾,保证预热正确;翻页到底的空页跳过) --- has_more = len(bars) > limit # 返回窗口之前还有更早历史(含预热行) indicators: dict[str, dict[str, list[float | None]]] = {} if bars: df = pd.DataFrame({"close": [b.close for b in bars], "high": [b.high for b in bars], "low": [b.low for b in bars]}) closes, highs, lows = df["close"], df["high"], df["low"] macd = ind.macd(closes) kdj = ind.kdj(highs, lows, closes) boll = ind.bollinger(closes) indicators = { # MA 始终返回全量集合(前端本地计算 MA,此处仅保留兼容;缓存键不依赖 ma_periods) "ma": {f"ma{p}": series_to_jsonable(ind.ma(closes, p)) for p in FULL_MA_SET}, "macd": { "dif": series_to_jsonable(macd["macd"]), "dea": series_to_jsonable(macd["signal"]), "hist": series_to_jsonable(macd["hist"]), }, "kdj": {k: series_to_jsonable(kdj[k]) for k in ("k", "d", "j")}, "rsi": { "rsi6": series_to_jsonable(ind.rsi(closes, 6)), "rsi12": series_to_jsonable(ind.rsi(closes, 12)), "rsi24": series_to_jsonable(ind.rsi(closes, 24)), }, "boll": {k: series_to_jsonable(boll[k]) for k in ("upper", "mid", "lower")}, "zx": { "short": series_to_jsonable(ind.ema2(closes)), "duokong": series_to_jsonable(ind.avg_ma(closes, tuple(zx_periods))), }, } limit = max(30, min(limit, len(bars))) for group in indicators.values(): for key in group: group[key] = group[key][-limit:] # --- 信息卡:Wave 2 已并行取回(stock_basic + 与行情同日对齐的快照、缺则最新日,见 INFO_SQL) --- row = info_row.first() if row is not None: m = row._mapping sb_name, sb_industry, sb_area, sb_market, sb_list_date = ( m["name"], m["industry"], m["area"], m["market"], m["list_date"] ) ds_turnover, ds_pe, ds_pb, ds_tmv, ds_cmv = ( m["turnover_rate"], m["pe_ttm"], m["pb"], m["total_mv"], m["circ_mv"] ) else: sb_name = sb_industry = sb_area = sb_market = sb_list_date = None ds_turnover = ds_pe = ds_pb = ds_tmv = ds_cmv = None def _yi(v) -> float | None: if v is None: return None v = float(v) return None if v != v else round(v / 1e4, 2) # 万元 -> 亿元 info = PreviewInfoOut( ts_code=ts_code, symbol=symbol, name=sb_name or ts_code, industry=sb_industry, area=sb_area, market=sb_market, list_date=sb_list_date, trade_date=last_daily.ts if last_daily else None, open=last_daily.open if last_daily else None, high=last_daily.high if last_daily else None, low=last_daily.low if last_daily else None, close=last_daily.close if last_daily else None, pre_close=prev_daily.close if prev_daily else None, pct_chg=((last_daily.close / prev_daily.close - 1) * 100) if last_daily and prev_daily and prev_daily.close else None, volume_hand=round(last_daily.volume / 100, 0) if last_daily else None, # 股 -> 手 amount_yi=round(last_daily.amount / 1e8, 2) if last_daily and last_daily.amount else None, # 元 -> 亿元 turnover_rate=ds_turnover, pe_ttm=ds_pe, pb=ds_pb, total_mv=_yi(ds_tmv), circ_mv=_yi(ds_cmv), ) candles = [ CandleOut(ts=b.ts, open=b.open, high=b.high, low=b.low, close=b.close, volume=b.volume, amount=b.amount, turnover=b.turnover) for b in bars[-limit:] ] resp = PreviewResponse(ts_code=ts_code, symbol=symbol, source=source, info=info, candles=candles, indicators=indicators, has_more=has_more) # 只序列化一次:本地(同步,120s)+ Redis(后台写,600s TTL 兜底跨进程/重启) raw = raw_json(resp) cache.local_set(f"pvj:{cache_key}", raw, ttl=120) cache.set_bg(f"pvj:{cache_key}", raw, ttl=600) return Response(content=raw, media_type="application/json")