feat(backtest): 个股回测接入费用(cfg通道)+slippage+benchmark放宽至4基准 [vps]
This commit is contained in:
+17
-4
@@ -59,12 +59,25 @@ def login(req: LoginRequest):
|
||||
return {"token": create_token(req.username)}
|
||||
|
||||
|
||||
_BENCHMARKS = ("hs300", "zz500", "zz1000", "zz2000")
|
||||
|
||||
|
||||
def _build_fee_cfg(req: CtaBacktestRequest) -> dict:
|
||||
"""前端费用参数打包成 cfg,走现有 cfg 通道注入 AShareBacktestingEngine。"""
|
||||
return {
|
||||
"commission_rate": req.commission_rate,
|
||||
"min_commission": req.min_commission,
|
||||
"stamp_duty_rate": req.stamp_duty_rate,
|
||||
"transfer_fee_rate": req.transfer_fee_rate,
|
||||
"slippage": req.slippage,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/backtest/cta", dependencies=[Depends(verify_token)])
|
||||
async def submit_cta(req: CtaBacktestRequest):
|
||||
"""Submit CTA backtest task"""
|
||||
# Validate benchmark parameter
|
||||
if req.benchmark not in ("hs300", "zz500"):
|
||||
raise HTTPException(status_code=422, detail=f"Invalid benchmark: {req.benchmark}. Must be 'hs300' or 'zz500'")
|
||||
if req.benchmark not in _BENCHMARKS:
|
||||
raise HTTPException(status_code=422, detail=f"Invalid benchmark: {req.benchmark}. Must be one of {_BENCHMARKS}")
|
||||
|
||||
cls = get_strategy_class(req.strategy)
|
||||
if cls is None:
|
||||
@@ -75,7 +88,7 @@ async def submit_cta(req: CtaBacktestRequest):
|
||||
params=req.params,
|
||||
start=req.start,
|
||||
end=req.end,
|
||||
cfg=None,
|
||||
cfg=_build_fee_cfg(req),
|
||||
benchmark=req.benchmark,
|
||||
capital=req.capital,
|
||||
position_pct=req.position_pct,
|
||||
|
||||
@@ -21,6 +21,7 @@ class CtaBacktestRequest(BaseModel):
|
||||
min_commission: float = 5.0 # 最低 5 元
|
||||
stamp_duty_rate: float = 0.0005 # 卖方 0.05%
|
||||
transfer_fee_rate: float = 0.00001 # 沪市 0.001%
|
||||
slippage: float = 0.0 # 滑点(比率,万10=0.001);0=不加滑点
|
||||
|
||||
|
||||
class OptimizeRequest(BaseModel):
|
||||
|
||||
@@ -180,6 +180,14 @@ def run_cta_backtest(
|
||||
# 一致,走 vnpy 原生 load_data。
|
||||
engine.raw_interval = interval
|
||||
|
||||
# 前端透传费用(cfg 通道:routes.submit_cta _build_fee_cfg 打包;None/非 dict→用函数默认值)
|
||||
if cfg and isinstance(cfg, dict):
|
||||
commission_rate = float(cfg.get("commission_rate", commission_rate))
|
||||
min_commission = float(cfg.get("min_commission", min_commission))
|
||||
stamp_duty_rate = float(cfg.get("stamp_duty_rate", stamp_duty_rate))
|
||||
transfer_fee_rate = float(cfg.get("transfer_fee_rate", transfer_fee_rate))
|
||||
_slippage = float(cfg.get("slippage", 0.0)) if (cfg and isinstance(cfg, dict)) else 0.0
|
||||
|
||||
# A 股适配参数(定寸 + 费用)
|
||||
engine.position_pct = position_pct
|
||||
engine.commission_rate = commission_rate
|
||||
@@ -187,6 +195,8 @@ def run_cta_backtest(
|
||||
engine.stamp_duty_rate = stamp_duty_rate
|
||||
engine.transfer_fee_rate = transfer_fee_rate
|
||||
engine.is_sse = (_exchange.value == "SSE")
|
||||
if _slippage:
|
||||
engine.slippage = _slippage # 覆盖 set_parameters 的 slippage=0
|
||||
|
||||
# Add strategy
|
||||
engine.add_strategy(strategy_class, params)
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Tests for 个股回测费用接入 + benchmark 放宽 (Task D).
|
||||
|
||||
schema 层本地可跑(pydantic);_build_fee_cfg 走 routes→auth→jwt,需 NAS docker 或齐后端 venv。
|
||||
"""
|
||||
from sanguo_api.schemas import CtaBacktestRequest
|
||||
|
||||
|
||||
def test_slippage_field_present():
|
||||
req = CtaBacktestRequest(
|
||||
symbol="600519.SH", strategy="DoubleMaStrategy", start="2024-01-02", end="2024-06-28",
|
||||
slippage=0.001,
|
||||
)
|
||||
assert req.slippage == 0.001
|
||||
assert req.commission_rate == 0.00025 # 默认仍在
|
||||
|
||||
|
||||
def test_benchmark_accepts_four_values():
|
||||
for b in ("hs300", "zz500", "zz1000", "zz2000"):
|
||||
req = CtaBacktestRequest(
|
||||
symbol="600519.SH", strategy="DoubleMaStrategy",
|
||||
start="2024-01-02", end="2024-06-28", benchmark=b,
|
||||
)
|
||||
assert req.benchmark == b
|
||||
|
||||
|
||||
def test_build_fee_cfg_carries_all_fields():
|
||||
from sanguo_api.routes import _build_fee_cfg
|
||||
req = CtaBacktestRequest(
|
||||
symbol="600519.SH", strategy="DoubleMaStrategy", start="2024-01-02", end="2024-06-28",
|
||||
commission_rate=0.0003, stamp_duty_rate=0.001, slippage=0.001,
|
||||
)
|
||||
cfg = _build_fee_cfg(req)
|
||||
assert cfg["commission_rate"] == 0.0003
|
||||
assert cfg["stamp_duty_rate"] == 0.001
|
||||
assert cfg["slippage"] == 0.001
|
||||
assert "min_commission" in cfg and "transfer_fee_rate" in cfg
|
||||
Reference in New Issue
Block a user