From 4ad505434b86e8c78e8e7cd6df3788dba667ef62 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Wed, 12 Aug 2026 23:56:14 +0800 Subject: [PATCH] =?UTF-8?q?feat(backtest):=20=E4=B8=AA=E8=82=A1=E5=9B=9E?= =?UTF-8?q?=E6=B5=8B=E6=8E=A5=E5=85=A5=E8=B4=B9=E7=94=A8(cfg=E9=80=9A?= =?UTF-8?q?=E9=81=93)+slippage+benchmark=E6=94=BE=E5=AE=BD=E8=87=B34?= =?UTF-8?q?=E5=9F=BA=E5=87=86=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_api/routes.py | 21 ++++++++++++++---- sanguo_api/schemas.py | 1 + sanguo_backtest/cta_engine.py | 10 +++++++++ tests/api/test_submit_cta_fees.py | 36 +++++++++++++++++++++++++++++++ 4 files changed, 64 insertions(+), 4 deletions(-) create mode 100644 tests/api/test_submit_cta_fees.py diff --git a/sanguo_api/routes.py b/sanguo_api/routes.py index b508749..06ad2c7 100644 --- a/sanguo_api/routes.py +++ b/sanguo_api/routes.py @@ -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, diff --git a/sanguo_api/schemas.py b/sanguo_api/schemas.py index ad7941a..587b39b 100644 --- a/sanguo_api/schemas.py +++ b/sanguo_api/schemas.py @@ -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): diff --git a/sanguo_backtest/cta_engine.py b/sanguo_backtest/cta_engine.py index 31904d6..3f71248 100644 --- a/sanguo_backtest/cta_engine.py +++ b/sanguo_backtest/cta_engine.py @@ -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) diff --git a/tests/api/test_submit_cta_fees.py b/tests/api/test_submit_cta_fees.py new file mode 100644 index 0000000..6a2bac0 --- /dev/null +++ b/tests/api/test_submit_cta_fees.py @@ -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