"""组合策略 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 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="hs300_subset", description="标的池(占位,MVP 用默认)") 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=不加)") @router.post("/portfolio/backtest", dependencies=[Depends(verify_token)]) async def run_portfolio_backtest(req: PortfolioBacktestRequest): """异步提交组合回测任务,返回 task_id。前端轮询 GET /task/{id} 再取结果。""" 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, provider_config=None, commission_rate=req.commission_rate, stamp_duty_rate=req.stamp_duty_rate, min_commission=req.min_commission, slippage=req.slippage, ) 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, } @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 []