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"],
|
|
)
|