Files
sanguo_vnpy_v2/sanguo_api/routes_portfolio.py
T

160 lines
6.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""组合策略 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 []