提交
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user