提交
This commit is contained in:
152
backend/app/api/backtest.py
Normal file
152
backend/app/api/backtest.py
Normal file
@@ -0,0 +1,152 @@
|
||||
"""回测域路由:策略回测(旧 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"],
|
||||
)
|
||||
Reference in New Issue
Block a user