Files
claude_dev b2c41c73b8
CI/CD / test (push) Successful in 10s
CI/CD / nas-deploy (push) Successful in 42s
CI/CD / nas-verify (push) Successful in 5s
fix(backtest): CTA metrics benchmark 缺失降级(不阻塞整组指标图)
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)。
2026-08-02 20:37:50 +08:00

446 lines
18 KiB
Python
Raw Permalink 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)"
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