b2c41c73b8
benchmark 数据缺失时原实现 if 跳过整段 compute_metrics → {task_id}_metrics.json 不写 → 策略净值/回撤/波动图也空(不只基准图)。改为:benchmark 空时传空 Series + log WARN,compute_metrics 内部 reindex→fillna(0) 容空(基准类指标 NaN→None,策略指标正常),metrics.json 照写。效果:策略图照常出,只基准/Alpha/Beta 图空(前端 benchmark-curve/risk-series 已容错)。加测试 test_run_cta_backtest_benchmark_missing_degrades(mock read_index_daily 返空 → 验 compute_metrics 仍 called + metrics.json 仍写 + 策略指标有值)。6 测试全绿(Mac pytest)。
446 lines
18 KiB
Python
446 lines
18 KiB
Python
"""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)"
|
||
|
||
def test_run_cta_backtest_benchmark_missing_degrades(self, temp_db_path):
|
||
"""benchmark 缺失时 compute_metrics 仍执行、metrics.json 照写(降级)。
|
||
|
||
根因:read_index_daily 返空时原实现 if 跳过整段 compute_metrics →
|
||
{task_id}_metrics.json 不写 → 策略净值/回撤/波动图也空(不只基准图)。
|
||
降级:benchmark 空时传空 Series 给 compute_metrics(内部 reindex→fillna(0),
|
||
基准类指标返 NaN→None,策略指标正常),metrics.json 照写。
|
||
"""
|
||
mock_strategy_class = Mock()
|
||
mock_strategy_class.__name__ = "DegBmStrategy"
|
||
|
||
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
|
||
daily_df = pd.DataFrame(
|
||
{"balance": [1_000_000.0 * (1.001 ** i) for i in range(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_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
|
||
mock_engine.history_data = [] # 空 history → 定寸跳过
|
||
|
||
mock_cfg = Mock()
|
||
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
|
||
|
||
mock_ashare = MagicMock()
|
||
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||
|
||
mock_metrics_result = Mock()
|
||
mock_metrics_result.scalars = {
|
||
"total_return": 0.15, # 策略指标(不依赖基准)有值
|
||
"max_drawdown": -0.05,
|
||
"alpha": None, # 基准类指标 NaN→None(基准缺失)
|
||
"benchmark_return": None,
|
||
}
|
||
mock_metrics_result.series = {
|
||
"equity_curve": pd.Series([1.0, 1.05, 1.1]),
|
||
"drawdown": pd.Series([0.0, -0.01, -0.02]),
|
||
}
|
||
|
||
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(),
|
||
}):
|
||
# read_index_daily 返空 DataFrame = benchmark 缺失
|
||
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=pd.DataFrame()):
|
||
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result) as mock_cm:
|
||
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,
|
||
)
|
||
|
||
# 1. compute_metrics 仍被调用(降级,没因 benchmark 空跳过)
|
||
mock_cm.assert_called_once()
|
||
# 2. 传入的 benchmark_returns 是空 Series(降级标记)
|
||
benchmark_arg = mock_cm.call_args[0][1]
|
||
assert isinstance(benchmark_arg, pd.Series)
|
||
assert benchmark_arg.empty, "benchmark 缺失时应传空 Series 给 compute_metrics"
|
||
|
||
# 3. metrics.json 仍写(策略净值/回撤/波动图数据源)
|
||
file_dir = Path(temp_db_path).parent
|
||
metrics_file = file_dir / f"{result.task_id}_metrics.json"
|
||
assert metrics_file.exists(), f"降级时 metrics.json 应照写: {metrics_file}"
|
||
with open(metrics_file) as f:
|
||
metrics_data = json.load(f)
|
||
series_keys = list(metrics_data.get("series", {}).keys())
|
||
assert "equity_curve" in series_keys
|
||
assert "drawdown" in series_keys
|
||
|
||
# 4. 策略指标有值(基准类 None)
|
||
assert result.statistics.get("total_return") == 0.15
|
||
assert result.statistics.get("alpha") is None |