This commit is contained in:
2026-09-09 15:07:58 +08:00
parent d656c05b3d
commit 71a0f6e404
31 changed files with 3657 additions and 9 deletions

422
backend/app/api/screener.py Normal file
View File

@@ -0,0 +1,422 @@
"""智能选股域路由自然语言选股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底座口径恒为 bfqTDX 全量 + 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 1candles 窗口(注入 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+ 信息卡(注入 sessionLATERAL 一条)并行 ---
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")