This commit is contained in:
2026-09-04 10:00:35 +08:00
parent 2735ff1fd9
commit 76b422320b
14 changed files with 611 additions and 218 deletions

View File

@@ -19,6 +19,7 @@ from datetime import datetime
import pandas as pd
from fastapi import APIRouter, Depends, File, HTTPException, Response, UploadFile
from fastapi.responses import StreamingResponse
from sqlalchemy import delete, func, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql.elements import TextClause
@@ -518,47 +519,88 @@ async def backtest_event(
# ---------- 智能选股 ----------
@router.post("/screener/run", response_model=ScreenerRunResponse)
@router.post("/screener/run")
async def screener_run(
req: ScreenerRunRequest,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> ScreenerRunResponse:
) -> StreamingResponse:
"""自然语言 -> LLM 解析条件 -> 全市场筛选。也可直传 conditions 跳过 LLM微调再跑
成功的提问(含解析出的条件与命中数)记录到 screener_queries供历史一键重跑。"""
try:
conds = req.conditions or await parse_conditions(req.text)
if not conds.indicator and not conds.snapshot:
raise HTTPException(status_code=400, detail="AI 未从描述中解析出任何筛选条件,请换种说法")
result = await engine.run_screen(session, conds, settings.screener_default_limit)
# 相同文本 + 相同条件的上一条不重复记录(一键重跑场景)
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),
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()
return ScreenerRunResponse(**result)
except HTTPException:
raise
except DataNotReadyError as e:
raise HTTPException(status_code=409, detail=str(e))
except ValueError as e: # 未知指标/字段、条件为空
raise HTTPException(status_code=400, detail=str(e))
except ScreenerError as e:
code = 503 if "未配置 LLM_API_KEY" in str(e) else 502
raise HTTPException(status_code=code, detail=str(e))
).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)