160 lines
6.7 KiB
Python
160 lines
6.7 KiB
Python
"""组合策略 API 路由(异步化 P0)。
|
||
|
||
POST /portfolio/backtest: 异步提交组合回测任务,返回 task_id 供前端轮询。
|
||
GET /portfolio/task/{task_id}/result: 取组合回测结果。
|
||
|
||
回测子进程执行(三机自适应 NAS/VPS/Mac-SSH)搬至 portfolio_worker.py,
|
||
超时从 600s 提升至 3600s。复用 sanguo_orchestrator task 框架。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from typing import Any
|
||
|
||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||
from pydantic import BaseModel, Field
|
||
|
||
from .auth import verify_token as verify_token_impl
|
||
from .routes import get_orchestrator
|
||
from .validation import validate_portfolio_request
|
||
|
||
logger = logging.getLogger(__name__)
|
||
router = APIRouter()
|
||
|
||
|
||
async def verify_token(authorization: str | None = Header(None)):
|
||
if authorization is None or not authorization.startswith("Bearer "):
|
||
raise HTTPException(401, "Missing/invalid authorization")
|
||
return verify_token_impl(authorization.split(" ", 1)[1])
|
||
|
||
|
||
class PortfolioBacktestRequest(BaseModel):
|
||
pool: str = Field(default="all", description="标的池(all=全市场)")
|
||
start_date: str = Field(default="2024-01-01", description="YYYY-MM-DD")
|
||
end_date: str = Field(default="2024-02-29", description="YYYY-MM-DD")
|
||
initial_cash: float = Field(default=1_000_000.0, description="初始资金(元)")
|
||
benchmark: str = Field(default="000300.XSHG", description="基准代码")
|
||
strategy: str = Field(
|
||
default="all_weather",
|
||
description="策略: all_weather/momentum_timing/value_selection/small_cap",
|
||
)
|
||
max_pool: int = Field(default=30, description="选股池上限: 0=全市场不限, N=前N只(MVP验证用)")
|
||
# A 股费用(对齐个股回测)
|
||
commission_rate: float = Field(default=0.0003, description="佣金率双边(万3=0.0003)")
|
||
stamp_duty_rate: float = Field(default=0.001, description="印花税率卖出(千1=0.001)")
|
||
min_commission: float = Field(default=5.0, description="单笔最低佣金(元)")
|
||
slippage: float = Field(default=0.0, description="滑点比率(万10=0.001,0=不加)")
|
||
interval: str = Field(default="d", description="K线周期:组合回放暂仅日线(d)")
|
||
instance_id: int | None = Field(default=None, description="§12.6 可选关联策略档案(带了则结果回写)")
|
||
|
||
|
||
@router.post("/portfolio/backtest", dependencies=[Depends(verify_token)])
|
||
async def run_portfolio_backtest(req: PortfolioBacktestRequest):
|
||
"""异步提交组合回测任务,返回 task_id。前端轮询 GET /task/{id} 再取结果。"""
|
||
if req.interval != "d":
|
||
raise HTTPException(400, "组合回放暂仅支持日线(影子柜台将支持全周期分钟档)")
|
||
validate_portfolio_request(req) # 400 中文提示(格式/未来/区间/资金/费率),拦在进队列前
|
||
# §12.6 补:带档案发起 → 当时代码快照(完成回写 run_meta.code_hash)
|
||
code_hash = None
|
||
if req.instance_id:
|
||
from .instance_store import get_instance_params_snapshot
|
||
from .code_versions import snapshot_code
|
||
|
||
cf = (get_instance_params_snapshot(req.instance_id) or {}).get("code_file") or ""
|
||
snap = snapshot_code(cf)
|
||
code_hash = snap["code_hash"] if snap else None
|
||
tid = await get_orchestrator().submit_portfolio(
|
||
start=req.start_date,
|
||
end=req.end_date,
|
||
cash=req.initial_cash,
|
||
benchmark=req.benchmark,
|
||
strategy=req.strategy,
|
||
max_pool=req.max_pool,
|
||
pool=req.pool,
|
||
provider_config=None,
|
||
commission_rate=req.commission_rate,
|
||
stamp_duty_rate=req.stamp_duty_rate,
|
||
min_commission=req.min_commission,
|
||
slippage=req.slippage,
|
||
interval=req.interval,
|
||
instance_id=req.instance_id,
|
||
code_hash=code_hash,
|
||
)
|
||
return {"task_id": tid}
|
||
|
||
|
||
@router.get("/portfolio/task/{task_id}/result", dependencies=[Depends(verify_token)])
|
||
def get_portfolio_result(task_id: str):
|
||
"""取组合回测结果。从 BacktestResult.statistics 取 metrics/stocks_selected/period,
|
||
从 equity_curve/trades(DataFrame)转 records。"""
|
||
# 区分 failed(422+error)/running·pending(404 not ready)/done(result)/not found(404)。
|
||
# 原实现对跑中/失败都 404,排查困难(分不清慢/死/失败)。
|
||
orch = get_orchestrator()
|
||
state = orch.get_status(task_id)
|
||
if state is None:
|
||
r = orch.get_result(task_id) # 不在内存池(历史 task 持久化 DB)→直接查 DB
|
||
if r is None:
|
||
raise HTTPException(status_code=404, detail="task not found")
|
||
elif state.value == "failed":
|
||
task = orch.pool.get_task(task_id)
|
||
raise HTTPException(
|
||
status_code=422,
|
||
detail=f"task failed: {task.error_msg if task else 'unknown'}",
|
||
)
|
||
elif state.value != "done":
|
||
raise HTTPException(
|
||
status_code=404, detail=f"result not ready (state={state.value})"
|
||
)
|
||
else:
|
||
r = orch.get_result(task_id)
|
||
if r is None:
|
||
raise HTTPException(status_code=404, detail="done but result missing")
|
||
|
||
stats: dict = r.statistics or {}
|
||
metrics = stats.get("metrics", {})
|
||
stocks_selected = stats.get("stocks_selected", [])
|
||
period = stats.get("period", {})
|
||
|
||
return {
|
||
"task_id": task_id,
|
||
"strategy": r.strategy,
|
||
"period": period,
|
||
"metrics": metrics,
|
||
"equity_curve": _df_to_records(r.equity_curve),
|
||
"trades": _df_to_records(r.trades),
|
||
"stocks_selected": stocks_selected,
|
||
"benchmark_curve": stats.get("benchmark_curve", []),
|
||
"drawdown_curve": stats.get("drawdown_curve", []),
|
||
"holdings_curve": stats.get("holdings_curve", []),
|
||
}
|
||
|
||
|
||
@router.get("/portfolio/task/{task_id}/status", dependencies=[Depends(verify_token)])
|
||
def get_portfolio_status(task_id: str):
|
||
"""组合回测任务状态(区分 pending/running/done/failed + error 详情)。
|
||
通用 GET /task/{id} 也返 status,本接口聚焦 portfolio 命名空间 + error 字段。"""
|
||
orch = get_orchestrator()
|
||
state = orch.get_status(task_id)
|
||
if state is None:
|
||
raise HTTPException(status_code=404, detail="task not found")
|
||
task = orch.pool.get_task(task_id)
|
||
return {
|
||
"task_id": task_id,
|
||
"state": state.value if hasattr(state, "value") else str(state),
|
||
"stage": orch.pool.get_stage(task_id) or "",
|
||
"error": task.error_msg if (task and isinstance(task.error_msg, str)) else None,
|
||
}
|
||
|
||
|
||
def _df_to_records(df: Any) -> list[dict]:
|
||
"""DataFrame/list → list[dict] (empty-safe)."""
|
||
if df is None:
|
||
return []
|
||
if hasattr(df, "empty") and df.empty:
|
||
return []
|
||
if hasattr(df, "to_dict"):
|
||
return df.to_dict(orient="records")
|
||
if isinstance(df, list):
|
||
return df
|
||
return []
|