diff --git a/tests/test_backtest.py b/tests/test_backtest.py index f642d4f..9f8d4a8 100644 --- a/tests/test_backtest.py +++ b/tests/test_backtest.py @@ -6,10 +6,12 @@ import pandas as pd import pytest from quant_engine.backtest import ( + BacktestResult, compare_to_benchmark, compute_nav_from_weights, compute_returns_from_nav, rebalance_periodic, + run_weight_backtest, weights_to_long_short, ) @@ -122,3 +124,86 @@ def test_compare_to_benchmark_rejects_non_overlapping_dates() -> None: with pytest.raises(ValueError, match="overlapping dates"): compare_to_benchmark(strategy, benchmark) + + +# ── 统一回测结果门面 ────────────────────────────────────── + + +def test_run_weight_backtest_returns_nav_returns_and_input_snapshot() -> None: + dates = pd.date_range("2026-01-05", periods=3, freq="B") + weights = pd.DataFrame({"A": [1.0]}, index=dates[:1]) + stock_returns = pd.DataFrame({"A": [0.10, -0.10, 0.20]}, index=dates) + + result = run_weight_backtest(weights, stock_returns, initial_capital=100.0) + + assert isinstance(result, BacktestResult) + pd.testing.assert_series_equal( + result.nav, + pd.Series([110.0, 99.0, 118.8], index=dates), + ) + pd.testing.assert_series_equal( + result.returns, + pd.Series([0.0, -0.1, 0.2], index=dates), + ) + pd.testing.assert_frame_equal(result.weights, weights) + + +def test_backtest_result_stats_reuses_standard_metrics_contract() -> None: + dates = pd.date_range("2026-01-05", periods=3, freq="B") + result = run_weight_backtest( + pd.DataFrame({"A": [1.0]}, index=dates[:1]), + pd.DataFrame({"A": [0.10, -0.10, 0.20]}, index=dates), + ) + + stats = result.stats(rf=0.02) + + assert stats["n_days"] == 3 + assert stats["ann_return"] == pytest.approx( + (1.0 * 0.9 * 1.2) ** (252 / 3) - 1.0 + ) + assert "sharpe" in stats + assert stats["drawback"] == stats["max_drawdown"] + + +def test_backtest_result_builds_benchmark_report() -> None: + dates = pd.date_range("2026-01-05", periods=3, freq="B") + benchmark = pd.Series([1.0, 1.05, 1.10], index=dates, name="benchmark") + result = run_weight_backtest( + pd.DataFrame({"A": [1.0]}, index=dates[:1]), + pd.DataFrame({"A": [0.10, -0.10, 0.20]}, index=dates), + benchmark_nav=benchmark, + ) + + report = result.benchmark_report() + + assert report.columns.tolist() == ["策略", "基准"] + assert report.loc["累计收益", "基准"] == pytest.approx(0.10) + + +def test_backtest_result_requires_benchmark_for_comparison() -> None: + dates = pd.date_range("2026-01-05", periods=2, freq="B") + result = run_weight_backtest( + pd.DataFrame({"A": [1.0]}, index=dates[:1]), + pd.DataFrame({"A": [0.0, 0.0]}, index=dates), + ) + + with pytest.raises(ValueError, match="benchmark_nav"): + result.benchmark_report() + + +def test_backtest_result_isolated_from_mutated_caller_inputs() -> None: + dates = pd.date_range("2026-01-05", periods=2, freq="B") + weights = pd.DataFrame({"A": [1.0]}, index=dates[:1]) + benchmark = pd.Series([1.0, 1.1], index=dates) + result = run_weight_backtest( + weights, + pd.DataFrame({"A": [0.0, 0.0]}, index=dates), + benchmark_nav=benchmark, + ) + + weights.iloc[0, 0] = 0.0 + benchmark.iloc[1] = 99.0 + + assert result.weights.iloc[0, 0] == 1.0 + assert result.benchmark_nav is not None + assert result.benchmark_nav.iloc[1] == 1.1