Files
sanguo_vnpy_v2/sanguo_api/routes_portfolio.py
T
claude_dev e22c8ee939
CI/CD / test (push) Successful in 8s
CI/CD / nas-deploy (push) Successful in 21s
CI/CD / nas-verify (push) Successful in 3s
feat(portfolio): max_pool 字段透传(model+handler+前端,解除30硬编码)
model 加 max_pool 字段(默认30,0=全市场不限);handler 用 req.max_pool 替换硬编码30;
前端加选股池上限 input + onSubmit 透传。下游 submit_portfolio→spec→argv--max-pool
链路异步化时已存在,本次只补入口。全周期验收可传 max_pool=0 对比 docker exec 旁证
2026-08-01 23:20:40 +08:00

93 lines
3.3 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
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验证用)")
@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,
)
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。"""
r = get_orchestrator().get_result(task_id)
if r is None:
raise HTTPException(status_code=404, detail="result not ready")
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,
}
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 []