"""回测域路由:策略回测(旧 API,K线+指标+买卖点+净值)+ 自然语言事件回测。""" 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.events import EventEngineError, run_event_backtest from ..backtest.strategies import build_strategy from ..data import fetcher, repository from ..data.aggregation import bars_per_year, resample_bars from ..data.symbols import is_etf_symbol from ..db import get_session from ..models import BacktestRun from ..schemas import ( BacktestRequest, BacktestResponse, CandleOut, EquityPoint, EventBacktestRequest, EventBacktestResponse, IndicatorOut, MetricsOut, SignalOut, ) from ..screener.llm import parse_event_spec, ScreenerError from ._deps import rows_to_bars, series_to_jsonable router = APIRouter() @router.post("/backtest", response_model=BacktestResponse) async def backtest( req: BacktestRequest, session: AsyncSession = Depends(get_session), ) -> BacktestResponse: # 真实数据:本地无缓存则先拉取 if 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), is_fund=is_etf_symbol(req.symbol), # ETF 免印花税/过户费 ) 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"], amount=r["amount"] if "amount" in df.columns else None, turnover=r["turnover"] if "turnover" in df.columns else None) 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("/backtest/event", response_model=EventBacktestResponse) async def backtest_event( req: EventBacktestRequest, session: AsyncSession = Depends(get_session), ) -> EventBacktestResponse: """自然语言事件回测:入场条件命中 -> 次日买入 -> 持有 N 日,单股或全市场汇总统计。 直传 spec 则跳过 LLM(前端调参重跑)。""" try: spec = req.spec or await parse_event_spec(req.text) result = await run_event_backtest( session, spec, ts_code=req.ts_code, start=req.start.date() if req.start else None, end=req.end.date() if req.end else None, ) except ScreenerError as e: raise HTTPException(status_code=502, detail=str(e)) except EventEngineError as e: raise HTTPException(status_code=400, detail=str(e)) except Exception as e: # noqa: BLE001 raise HTTPException(status_code=500, detail=f"事件回测失败: {e}") return EventBacktestResponse( text=req.text, spec=result["spec"], universe=result["universe"], start=result["start"], end=result["end"], stats=result["stats"], trades=result["trades"], total=result["total"], )