- 首页双入口(智能选股/策略回测):引入 vue-router,顶部导航 - 智能选股:自然语言 -> LLM 解析结构化条件(智谱 GLM,OpenAI 兼容,/v4 兼容)-> SQL 快照预筛 + pandas 指标过滤(复用 indicators 单一事实源) - 条件模型:指标 vs 常数/指标(value_indicator,如 DIF>DEA、close<布林下轨)、lookback+match 表达连续N天/近N天任一天、市值/PE/PB/换手率快照条件、默认排除 ST/退市/北交所 - 全市场数据同步:按 trade_date 批量拉取未复权日线(与回测 candles qfq 隔离),交易日历/股票列表本地缓存,daily_basic 仅最新截面,Tushare 限频兜底(分钟级重试/小时级降级) - 存储:DATABASE_URL 切远程 PostgreSQL(cirry.cn/stock),本地 SQLite 已移除 - .env 入库(私有仓库);smoke_test 扩展选股链路 Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
222 lines
8.2 KiB
Python
222 lines
8.2 KiB
Python
"""HTTP 路由(OpenAPI 契约的载体)。
|
||
|
||
GET /api/health 健康检查
|
||
GET /api/candles/{sym} 取 K 线(支持 1d/1w/1M/1y 周期,日线为基底聚合)
|
||
POST /api/backtest 跑回测,返回 K线+指标+买卖点+净值+绩效
|
||
POST /api/screener/run 智能选股:自然语言 -> 条件 -> 全市场筛选
|
||
POST /api/screener/sync 启动全市场数据同步(后台任务)
|
||
GET /api/screener/sync/status 同步任务状态与数据实况
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
|
||
import pandas as pd
|
||
from fastapi import APIRouter, Depends, HTTPException
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from .backtest.engine import BacktestConfig, run_backtest
|
||
from .backtest.strategies import build_strategy
|
||
from .config import settings
|
||
from .data import fetcher, repository
|
||
from .data.aggregation import bars_per_year, resample_bars
|
||
from .data.synthetic import seed_if_empty
|
||
from .db import get_session
|
||
from .domain import Bar
|
||
from .models import BacktestRun
|
||
from .schemas import (
|
||
BacktestRequest,
|
||
BacktestResponse,
|
||
CandleOut,
|
||
EquityPoint,
|
||
IndicatorOut,
|
||
MetricsOut,
|
||
ScreenerRunRequest,
|
||
ScreenerRunResponse,
|
||
ScreenerSyncRequest,
|
||
ScreenerSyncStatus,
|
||
SignalOut,
|
||
SyncRequest,
|
||
SyncResponse,
|
||
)
|
||
from .screener import engine, market_sync
|
||
from .screener.engine import DataNotReadyError
|
||
from .screener.llm import ScreenerError, parse_conditions
|
||
|
||
router = APIRouter(prefix="/api")
|
||
|
||
|
||
def _series_to_jsonable(s: pd.Series) -> list[float | None]:
|
||
"""NaN -> None(lightweight-charts 的 whitespace data,跳过指标预热期)。"""
|
||
out: list[float | None] = []
|
||
for v in s.tolist():
|
||
if v is None or (isinstance(v, float) and v != v):
|
||
out.append(None)
|
||
else:
|
||
out.append(float(v))
|
||
return out
|
||
|
||
|
||
def _rows_to_bars(rows) -> list[Bar]:
|
||
return [Bar(ts=r.ts, open=r.open, high=r.high, low=r.low, close=r.close, volume=r.volume) for r in rows]
|
||
|
||
|
||
@router.get("/health")
|
||
async def health() -> dict:
|
||
return {"status": "ok"}
|
||
|
||
|
||
@router.get("/candles/{symbol}", response_model=list[CandleOut])
|
||
async def get_candles(
|
||
symbol: str,
|
||
timeframe: str = "1d",
|
||
limit: int = 5000,
|
||
session: AsyncSession = Depends(get_session),
|
||
) -> list[CandleOut]:
|
||
await seed_if_empty(session, symbol="DEMO")
|
||
# 始终以日线为基底,再聚合到目标周期
|
||
rows = await repository.get_candles(session, symbol, "1d", limit=limit)
|
||
bars = resample_bars(_rows_to_bars(rows), timeframe)
|
||
return [CandleOut(ts=b.ts, open=b.open, high=b.high, low=b.low, close=b.close, volume=b.volume) for b in bars]
|
||
|
||
|
||
@router.post("/data/sync", response_model=SyncResponse)
|
||
async def sync_data(req: SyncRequest, session: AsyncSession = Depends(get_session)) -> SyncResponse:
|
||
"""主动拉取并缓存某标的的日线(Tushare 主 -> AKShare 兜底)。"""
|
||
try:
|
||
res = await fetcher.sync_symbol(
|
||
session, req.symbol, start=req.start, end=req.end, source=req.source, force=req.force
|
||
)
|
||
return SyncResponse(**res)
|
||
except Exception as e: # noqa: BLE001
|
||
raise HTTPException(status_code=502, detail=str(e))
|
||
|
||
|
||
@router.post("/backtest", response_model=BacktestResponse)
|
||
async def backtest(
|
||
req: BacktestRequest,
|
||
session: AsyncSession = Depends(get_session),
|
||
) -> BacktestResponse:
|
||
await seed_if_empty(session, symbol="DEMO")
|
||
|
||
# 非演示标的:首次自动拉取真实数据并缓存
|
||
if req.symbol != "DEMO" and not await fetcher.is_cached(session, req.symbol):
|
||
try:
|
||
await fetcher.sync_symbol(session, req.symbol, source="auto")
|
||
except Exception as e: # noqa: BLE001
|
||
raise HTTPException(status_code=502, detail=f"数据拉取失败: {e}")
|
||
|
||
# 日线为基底,聚合到请求周期
|
||
rows = await repository.get_candles(
|
||
session, req.symbol, "1d", start=req.start, end=req.end, limit=100000
|
||
)
|
||
if not rows:
|
||
raise HTTPException(status_code=404, detail=f"无数据: symbol={req.symbol}")
|
||
|
||
bars = resample_bars(_rows_to_bars(rows), req.timeframe)
|
||
if len(bars) < 2:
|
||
raise HTTPException(status_code=400, detail=f"周期 {req.timeframe} 下数据不足,无法回测")
|
||
|
||
try:
|
||
strategy = build_strategy(req.strategy, req.params)
|
||
except Exception as e: # noqa: BLE001
|
||
raise HTTPException(status_code=400, detail=f"策略构建失败: {e}")
|
||
cfg = BacktestConfig(
|
||
initial_cash=req.initial_cash,
|
||
fast_mode=req.fast_mode,
|
||
bars_per_year=bars_per_year(req.timeframe),
|
||
)
|
||
result = run_backtest(bars, strategy, cfg)
|
||
|
||
df: pd.DataFrame = result["df"]
|
||
m = result["metrics"]
|
||
|
||
# 记录到回测运行注册表(可复现/可审计的基础)
|
||
session.add(
|
||
BacktestRun(
|
||
symbol=req.symbol,
|
||
strategy=req.strategy,
|
||
timeframe=req.timeframe,
|
||
params_json=json.dumps(req.params, ensure_ascii=False),
|
||
initial_cash=req.initial_cash,
|
||
total_return=m["total_return"],
|
||
max_drawdown=m["max_drawdown"],
|
||
sharpe=m["sharpe"],
|
||
num_trades=m["num_trades"],
|
||
)
|
||
)
|
||
await session.commit()
|
||
|
||
candles = [
|
||
CandleOut(ts=r["ts"], open=r["open"], high=r["high"], low=r["low"],
|
||
close=r["close"], volume=r["volume"])
|
||
for _, r in df.iterrows()
|
||
]
|
||
signals = [
|
||
SignalOut(ts=f.ts, side=f.side.value, price=f.price, qty=f.qty)
|
||
for f in result["fills"]
|
||
]
|
||
indicators = IndicatorOut(
|
||
strategy=req.strategy,
|
||
data={col: _series_to_jsonable(df[col]) for col in result["indicator_cols"]},
|
||
)
|
||
equity = [EquityPoint(ts=t.to_pydatetime(), value=float(v))
|
||
for t, v in result["equity"].items()]
|
||
|
||
return BacktestResponse(
|
||
symbol=req.symbol,
|
||
timeframe=req.timeframe,
|
||
strategy=req.strategy,
|
||
candles=candles,
|
||
indicators=indicators,
|
||
signals=signals,
|
||
equity=equity,
|
||
metrics=MetricsOut(**m),
|
||
final_cash=result["final_cash"],
|
||
final_position=result["final_position"],
|
||
initial_cash=req.initial_cash,
|
||
)
|
||
|
||
|
||
# ---------- 智能选股 ----------
|
||
@router.post("/screener/run", response_model=ScreenerRunResponse)
|
||
async def screener_run(
|
||
req: ScreenerRunRequest, session: AsyncSession = Depends(get_session)
|
||
) -> ScreenerRunResponse:
|
||
"""自然语言 -> LLM 解析条件 -> 全市场筛选。也可直传 conditions 跳过 LLM(微调再跑)。"""
|
||
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)
|
||
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))
|
||
|
||
|
||
@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})
|