Files
sanguo_vnpy_v2/tests/backtest/test_cta_engine.py
T
claude_dev 8d55e414fa fix(backtest): A股适配层—定寸/做空拦截/真实费用/口径统一(Phase1+2)
审计发现包装层系统性失真(2 CRITICAL+7 HIGH),vnpy底座可信但A股场景未适配:
- C1 定寸: engine.size=N(满仓手数),策略volume=1手=N股,开平对称(pos归零)
- C2 做空拦截: SHORT+OPEN拒单,long-only,SHORT+CLOSE平多允许
- H3 A股费用: AShareDailyResult重算(佣金保底5元/印花税卖方/过户费沪市)
- H4 收益口径: simple return从balance算(不再用vnpy log return喂empyrical)
- H5+口径: benchmark ffill对齐不缩样本; sizing_shares_per_lot暴露
- H7 退化检测: 零成交/空数据标degenerate不静默done
- H8 task_id: optimize/factor用uuid4(原id()内存地址)
- 静默except改warning

验证: 容器内真实vnpy DoubleMa 600000 2022-2024, total_return 1e-6→42.3%,
end_balance 100万→142万, SHORT+OPEN成交0笔, N=7800股/手.
22 backtest测试全绿(含集成测试), API健康200.
2026-07-12 23:39:45 +08:00

363 lines
15 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.
"""Tests for sanguo_backtest.cta_engine module."""
# Mock vnpy and tzlocal modules before importing anything that depends on them
import sys
import importlib.util
from unittest.mock import MagicMock
# Only mock these when the real module is NOT importable (e.g. local env missing empyrical).
# In the container (authoritative test env) everything imports fine, so NO mock is installed
# → zero sys.modules pollution leaking into other test modules at collection time.
# (Previous unconditional sys.modules[...]=MagicMock() here broke datareader/spike/factor/metrics
# tests because it ran at collection time and never restored.)
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = MagicMock(return_value="UTC")
_MOCK_FACTORIES = {
"tzlocal": lambda: mock_tzlocal,
"vnpy.trader.setting": MagicMock,
"vnpy.trader.constant": MagicMock,
"vnpy.trader.object": MagicMock,
"vnpy.trader.database": MagicMock,
"vnpy_ctastrategy.backtesting": MagicMock,
"empyrical": MagicMock,
}
_SAVED_MODULES = {}
for _name, _factory in _MOCK_FACTORIES.items():
try:
_found = importlib.util.find_spec(_name)
except ModuleNotFoundError:
_found = None
if _found is None:
_SAVED_MODULES[_name] = sys.modules.get(_name)
sys.modules[_name] = _factory()
import pytest
import json
from unittest.mock import Mock, patch, MagicMock
from datetime import datetime
from pathlib import Path
import pandas as pd
from sanguo_backtest.cta_engine import run_cta_backtest
@pytest.fixture(scope="module", autouse=True)
def _restore_modules_after():
"""Restore sys.modules after this module's tests, so downstream tests import the
REAL vnpy/empyrical instead of the mocks we installed above. Also drop cta_engine/
metrics from the module cache (they were imported while mocks were active) so they
re-import fresh with real dependencies."""
yield
for cached in ("sanguo_backtest.cta_engine", "sanguo_backtest.metrics", "sanguo_backtest.ashare_engine"):
sys.modules.pop(cached, None)
for k, orig in _SAVED_MODULES.items():
if orig is None:
sys.modules.pop(k, None)
else:
sys.modules[k] = orig
class TestRunCtaBacktest:
"""Test suite for run_cta_backtest function."""
def test_run_cta_backtest_returns_result(self, temp_db_path):
"""Test successful CTA backtest execution returns proper result."""
# Mock strategy class
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "TestStrategy"
# Mock BacktestingEngine — calculate_result() returns daily_df (DataFrame),
# calculate_statistics(df) returns the stats dict (vnpy API, matches cta_engine)
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = pd.DataFrame(
{"balance": [1_000_000.0, 1_010_000.0], "net_pnl": [0.0, 10000.0]},
index=pd.date_range("2024-01-01", periods=2, freq="D"),
)
mock_engine.calculate_statistics.return_value = {
"total_return": 0.15,
"sharpe_ratio": 1.2,
"max_drawdown": -0.08,
"win_rate": 0.55
}
mock_engine.get_all_daily_results.return_value = [
{"date": "2024-01-01", "balance": 100000},
{"date": "2024-01-02", "balance": 101000},
]
# Mock config
mock_cfg = Mock()
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={"window": 20},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify result structure
assert result.type == "cta"
assert result.status == "done"
assert result.strategy == "TestStrategy"
assert result.symbol == "600000"
assert result.params == {"window": 20}
assert result.start == "2024-01-01"
assert result.end == "2024-03-31"
# statistics 是 dict——vnpy calculate_statistics 的原始键会被 compute_metrics
# 的 scalars 覆盖/扩充(容器内 config 加载成功时 metrics 会跑,覆盖 mock 的 0.15
# Mac 无 config 时 metrics 跳过,保留 mock 值)。具体数值随环境,真实值由
# test_integration_ashare 验证;此处只验结构 + mock 特有字段。
assert isinstance(result.statistics, dict)
assert result.statistics.get("sizing_shares_per_lot") == 0
assert result.error_msg is None
# Verify BacktestingEngine methods were called
mock_engine.set_parameters.assert_called_once()
mock_engine.add_strategy.assert_called_once_with(mock_strategy_class, {"window": 20})
mock_engine.load_data.assert_called_once()
mock_engine.run_backtesting.assert_called_once()
mock_engine.calculate_result.assert_called_once()
mock_engine.calculate_statistics.assert_called_once()
def test_run_cta_backtest_handles_exception(self, temp_db_path):
"""Test that exceptions during backtest are handled properly."""
# Mock strategy class
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "FailingStrategy"
# Mock BacktestingEngine to raise exception
mock_engine = MagicMock()
mock_engine.load_data.side_effect = RuntimeError("Data loading failed")
# Mock config
mock_cfg = Mock()
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="000001",
params={"window": 20},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify failed result structure
assert result.type == "cta"
assert result.status == "failed"
assert result.strategy == "FailingStrategy"
assert result.symbol == "000001"
assert result.error_msg is not None
assert "RuntimeError" in result.error_msg
assert "Data loading failed" in result.error_msg
def test_run_cta_backtest_generates_task_id(self, temp_db_path):
"""Test that run_cta_backtest generates unique task IDs."""
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "IdTestStrategy"
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = {"sharpe_ratio": 1.0}
mock_engine.get_all_daily_results.return_value = []
mock_cfg = Mock()
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
result1 = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
result2 = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600001",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify unique task IDs
assert result1.task_id != result2.task_id
assert result1.task_id.startswith("cta_")
assert result2.task_id.startswith("cta_")
def test_run_cta_backtest_computes_relative_metrics(self, temp_db_path):
"""Test that run_cta_backtest computes relative metrics against benchmark."""
# Mock strategy class
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "BmTestStrategy"
# Mock daily_df with 'return' column (required by compute_metrics)
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
daily_df = pd.DataFrame({
"return": [0.001] * len(dates)
}, index=dates)
# Mock vnpy BacktestingEngine
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = daily_df
mock_engine.calculate_statistics.return_value = {
"total_return": 0.15,
"sharpe_ratio": 1.2,
"max_drawdown": -0.08,
}
# Mock config with benchmark
mock_cfg = Mock()
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
# Mock read_index_daily to return benchmark data
mock_bench_df = pd.DataFrame({
"date": dates,
"close": [100.0] * len(dates)
})
# Mock compute_metrics result
mock_metrics_result = Mock()
mock_metrics_result.scalars = {
"alpha": 0.05,
"beta": 0.95,
"sharpe_ratio": 1.3,
"total_return": 0.15,
"benchmark_return": 0.10
}
mock_metrics_result.series = {
"equity_curve": pd.Series([1.0, 1.1, 1.2]),
"benchmark_curve": pd.Series([1.0, 1.05, 1.1]),
"alpha": pd.Series([0.01, 0.02, 0.03]),
"beta": pd.Series([0.9, 0.95, 1.0]),
"drawdown": pd.Series([0.0, -0.01, -0.02])
}
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
# Mock tzlocal and vnpy modules to avoid import errors
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
with patch.dict("sys.modules", {
"sanguo_backtest.ashare_engine": mock_ashare,
"tzlocal": mock_tzlocal,
"vnpy.trader.setting": MagicMock()
}):
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=mock_bench_df):
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={"window": 20},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify result contains relative metrics (scalars merged into statistics)
assert result.statistics.get("alpha") == 0.05
assert result.statistics.get("beta") == 0.95
assert result.statistics.get("sharpe_ratio") == 1.3 # Should be present
assert result.statistics.get("total_return") == 0.15 # Should be present
# Verify metrics JSON file was written
file_dir = Path(temp_db_path).parent
metrics_file = file_dir / f"{result.task_id}_metrics.json"
assert metrics_file.exists(), f"Metrics file not found: {metrics_file}"
# Verify metrics file can be loaded and contains expected keys
with open(metrics_file, "r") as f:
metrics_data = json.load(f)
# Check that we have 5 series keys
series_keys = list(metrics_data.get("series", {}).keys())
assert len(series_keys) == 5
assert "equity_curve" in series_keys
assert "benchmark_curve" in series_keys
assert "alpha" in series_keys
assert "beta" in series_keys
assert "drawdown" in series_keys
def test_run_cta_backtest_default_benchmark_hs300(self, temp_db_path):
"""Test that default benchmark is hs300 when not specified in config."""
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "DefaultBmStrategy"
# Mock daily_df
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
daily_df = pd.DataFrame({
"return": [0.001] * len(dates)
}, index=dates)
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = daily_df
mock_engine.calculate_statistics.return_value = {"total_return": 0.15}
# Mock config without benchmark (should use default hs300)
mock_cfg = Mock()
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
mock_bench_df = pd.DataFrame({
"date": dates,
"close": [100.0] * len(dates)
})
mock_metrics_result = Mock()
mock_metrics_result.scalars = {"alpha": 0.05}
mock_metrics_result.series = {}
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
# Mock tzlocal and vnpy modules to avoid import errors
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
with patch.dict("sys.modules", {
"sanguo_backtest.ashare_engine": mock_ashare,
"tzlocal": mock_tzlocal,
"vnpy.trader.setting": MagicMock()
}):
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=mock_bench_df) as mock_read:
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify read_index_daily was called with hs300 code (sh000300)
mock_read.assert_called_once()
call_args = mock_read.call_args
assert call_args[0][0] == "sh000300", "Default benchmark should be hs300 (sh000300)"