first commit
This commit is contained in:
167
backend/app/api.py
Normal file
167
backend/app/api.py
Normal file
@@ -0,0 +1,167 @@
|
||||
"""HTTP 路由(OpenAPI 契约的载体)。
|
||||
|
||||
GET /api/health 健康检查
|
||||
GET /api/candles/{sym} 取 K 线(支持 1d/1w/1M/1y 周期,日线为基底聚合)
|
||||
POST /api/backtest 跑回测,返回 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.strategies import build_strategy
|
||||
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,
|
||||
SignalOut,
|
||||
SyncRequest,
|
||||
SyncResponse,
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
Reference in New Issue
Block a user