153 lines
5.3 KiB
Python
153 lines
5.3 KiB
Python
"""回测域路由:策略回测(旧 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"],
|
||
)
|