423 lines
20 KiB
Python
423 lines
20 KiB
Python
"""智能选股域路由:自然语言选股(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")
|