diff --git a/backend/app/backtest/strategy.py b/backend/app/backtest/strategy.py index edb4c30..4a3836e 100644 --- a/backend/app/backtest/strategy.py +++ b/backend/app/backtest/strategy.py @@ -600,6 +600,7 @@ class StrategyBacktestService: config.holding_days, config.minute_fill, json.dumps(config.overrides or {}, sort_keys=True, ensure_ascii=False, default=str), + json.dumps(config.regime_filter or {}, sort_keys=True, ensure_ascii=False, default=str), ) def _resolve_composite_feature_plan( @@ -854,6 +855,8 @@ class StrategyBacktestService: _rm = self._build_regime_mask( market_data.timestamp_labels, first.regime_filter, getattr(getattr(self.engine.repo, "store", None), "data_dir", None), + required_start=first.start, + required_end=first.end, ) if _rm is not None: entry_time_mask = entry_time_mask & _rm @@ -1139,10 +1142,15 @@ class StrategyBacktestService: config.end, ) # 市场环境过滤(强制 T-1): 只叠加 entry, 不影响 exit - _rm = self._build_regime_mask( - market_data.timestamp_labels, config.regime_filter, - getattr(getattr(self.engine.repo, "store", None), "data_dir", None), - ) + try: + _rm = self._build_regime_mask( + market_data.timestamp_labels, config.regime_filter, + getattr(getattr(self.engine.repo, "store", None), "data_dir", None), + required_start=config.start, + required_end=config.end, + ) + except ValueError as e: + return _err(str(e)) if _rm is not None: entry_time_mask = entry_time_mask & _rm exit_time_mask = self._matrix_date_range_mask( @@ -1238,10 +1246,15 @@ class StrategyBacktestService: config.start, config.end, ) - _rm = self._build_regime_mask( - market_data.timestamp_labels, config.regime_filter, - getattr(getattr(self.engine.repo, "store", None), "data_dir", None), - ) + try: + _rm = self._build_regime_mask( + market_data.timestamp_labels, config.regime_filter, + getattr(getattr(self.engine.repo, "store", None), "data_dir", None), + required_start=config.start, + required_end=config.end, + ) + except ValueError as e: + return _err(str(e)) if _rm is not None: entry_time_mask = entry_time_mask & _rm exit_time_mask = self._matrix_date_range_mask( @@ -1348,6 +1361,27 @@ class StrategyBacktestService: formal_candidate_mask = candidate_mask & formal_range entry_mask = self._build_entry_mask_from_candidate(panel, candidate_mask, s, entry_signals) entry_mask = entry_mask & formal_range + if config.regime_filter: + date_values = panel.get_column("date").unique().sort().to_list() + date_labels = tuple(str(value)[:10] for value in date_values) + try: + regime_time_mask = self._build_regime_mask( + date_labels, + config.regime_filter, + getattr(getattr(self.engine.repo, "store", None), "data_dir", None), + required_start=config.start, + required_end=config.end, + ) + except ValueError as e: + return _err(str(e)) + if regime_time_mask is not None: + allowed_dates = [ + value for value, allowed in zip(date_values, regime_time_mask, strict=True) + if allowed + ] + regime_row_mask = panel.get_column("date").is_in(allowed_dates).fill_null(False) + formal_candidate_mask = formal_candidate_mask & regime_row_mask + entry_mask = entry_mask & regime_row_mask raw_exit_mask = self._build_signal_mask(panel, exit_signals, "_exit") exit_range = self._date_range_mask(panel, config.start, load_end) if config.mode == "full" else formal_range exit_mask = raw_exit_mask & exit_range @@ -1649,25 +1683,31 @@ class StrategyBacktestService: timestamp_labels: tuple[str, ...], regime_filter: dict | None, data_dir: Path | None, + *, + required_start: date | None = None, + required_end: date | None = None, ) -> np.ndarray | None: """构造逐日 regime mask。强制 T-1 防未来函数: regime[T-1] 决定 entry[T]。 timestamp_labels[i] 的入场资格 = 它的"前一交易日"的 regime 是否满足条件。 "前一交易日"用 timestamp_labels 自身的顺序确定(回测时间轴上的前一天)。 边界: 首日无前一日环境 → 默认允许(不阻断)。 - regime_filter 为 None 或无 regime 数据时返回 None(不过滤)。 + regime_filter 为 None 时返回 None(不过滤)。启用过滤后缺少正式区间所需的 + 环境数据则 fail-closed, 避免界面显示已过滤但实际静默放行。 """ - if not regime_filter or data_dir is None: + if not regime_filter: return None allowed_states = set(regime_filter.get("states") or []) min_score = regime_filter.get("min_score") if not allowed_states and min_score is None: return None + if data_dir is None: + raise ValueError("市场环境过滤不可用: 未找到环境数据目录") from app.services import regime_builder regime_df = regime_builder.load_regime_history(data_dir) if regime_df.is_empty(): - return None + raise ValueError("市场环境数据为空, 请先在数据页完成市场环境计算后再回测") # 构建 date(ISO) → (state, score) 映射 regime_map: dict[str, tuple[str, int]] = {} @@ -1680,11 +1720,21 @@ class StrategyBacktestService: # 对每个 label, 找它的前一交易日的 regime(timestamp_labels 顺序里的前一天) n = len(timestamp_labels) mask = np.ones(n, dtype=bool) # 默认允许 + required_start_text = str(required_start) if required_start is not None else None + required_end_text = str(required_end) if required_end is not None else None + missing_dates: list[str] = [] for i in range(1, n): + current_label = timestamp_labels[i][:10] prev_label = timestamp_labels[i - 1][:10] entry = regime_map.get(prev_label) if entry is None: - continue # 无前一日环境数据 → 允许(不阻断) + required = ( + (required_start_text is None or current_label >= required_start_text) + and (required_end_text is None or current_label <= required_end_text) + ) + if required: + missing_dates.append(prev_label) + continue state, score = entry ok = True if allowed_states and state not in allowed_states: @@ -1692,6 +1742,13 @@ class StrategyBacktestService: if min_score is not None and score < min_score: ok = False mask[i] = ok + if missing_dates: + first_missing = missing_dates[0] + suffix = f" 等 {len(missing_dates)} 天" if len(missing_dates) > 1 else "" + raise ValueError( + f"市场环境数据覆盖不完整: 缺少前一交易日环境 {first_missing}{suffix}, " + "请先补算对应区间" + ) return mask def _build_candidate_filter_mask( @@ -1975,6 +2032,7 @@ class StrategyBacktestService: "mode": c.mode, "holding_days": c.holding_days, "minute_fill": c.minute_fill, + "regime_filter": c.regime_filter, } @staticmethod diff --git a/backend/tests/backtest/test_strategy_backtest_correctness.py b/backend/tests/backtest/test_strategy_backtest_correctness.py index d9d8489..61e080b 100644 --- a/backend/tests/backtest/test_strategy_backtest_correctness.py +++ b/backend/tests/backtest/test_strategy_backtest_correctness.py @@ -9,6 +9,7 @@ import polars as pl from app.backtest.engine import BacktestEngine, SimResult from app.backtest.matrix import build_market_data_matrix, make_signal_matrix, rolling_mean from app.backtest.strategy import StrategyBacktestConfig, StrategyBacktestService +from app.services import regime_builder from app.strategy.engine import StrategyDef @@ -43,14 +44,17 @@ class _StrategyEngineStub: class _RepoStub: + def __init__(self, data_dir=None) -> None: + self.store = SimpleNamespace(data_dir=data_dir) + def get_index_daily(self, *args, **kwargs) -> pl.DataFrame: return pl.DataFrame() class _EngineStub: - def __init__(self, panel: pl.DataFrame) -> None: + def __init__(self, panel: pl.DataFrame, data_dir=None) -> None: self.panel = panel - self.repo = _RepoStub() + self.repo = _RepoStub(data_dir) self.load_args = None self.load_count = 0 self.sim_panel: pl.DataFrame | None = None @@ -162,6 +166,55 @@ def test_basic_filter_only_limits_entries_not_panel_rows(): } +def test_non_matrix_strategy_applies_regime_filter_and_reports_config(tmp_path): + start = date(2024, 1, 1) + panel = pl.DataFrame([ + { + "symbol": "A", + "name": "A", + "date": start + timedelta(days=offset), + "open": 10.0, + "high": 10.0, + "low": 10.0, + "close": 10.0, + "volume": 1000.0, + "amount": 1000.0, + "signal_limit_up": False, + "signal_limit_down": False, + } + for offset in range(3) + ]).sort(["symbol", "date"]) + regime_builder.upsert_regime_history(tmp_path, pl.DataFrame({ + "date": [start, start + timedelta(days=1)], + "state": ["weak", "strong"], + "score": [10, 85], + })) + engine = _EngineStub(panel, data_dir=tmp_path) + service = StrategyBacktestService(engine=engine, strategy_engine=_StrategyEngineStub(_strategy())) + regime_filter = {"states": ["strong"]} + + result = service.run(StrategyBacktestConfig( + strategy_id="test", + symbols=None, + start=start, + end=start + timedelta(days=2), + matching="close_t", + mode="position", + regime_filter=regime_filter, + )) + + assert result.error is None + assert engine.sim_matrix is not None + assert engine.sim_matrix.entry[:, 0].tolist() == [1, 0, 1] + assert result.config["regime_filter"] == regime_filter + assert result.stats["selection"] == { + "strategy_matches": 2, + "entry_candidates": 2, + "entry_trigger_filtered": 0, + "entry_trigger_enabled": False, + } + + def test_selection_stats_explain_entry_trigger_filtering(): start = date(2024, 1, 1) panel = pl.DataFrame([ @@ -423,6 +476,21 @@ def test_matrix_optimizer_preparation_loads_and_builds_base_data_once(): assert all(result.stats["shared_market_data_bytes"] == prepared.market_data.nbytes for result in results) +def test_matrix_prepare_signature_includes_regime_filter(): + base = dict( + strategy_id="native", + symbols=None, + start=date(2024, 1, 1), + end=date(2024, 1, 2), + ) + without_filter = StrategyBacktestConfig(**base) + with_filter = StrategyBacktestConfig(**base, regime_filter={"states": ["strong"]}) + + assert StrategyBacktestService._matrix_prepare_signature(without_filter) != ( + StrategyBacktestService._matrix_prepare_signature(with_filter) + ) + + def test_matrix_cache_preserves_trades_daily_equity_and_core_stats(): start = date(2024, 1, 1) panel = pl.DataFrame([ diff --git a/backend/tests/test_regime_builder.py b/backend/tests/test_regime_builder.py index 947a349..c031546 100644 --- a/backend/tests/test_regime_builder.py +++ b/backend/tests/test_regime_builder.py @@ -13,6 +13,7 @@ import time from datetime import date import polars as pl +import pytest from app.services import regime_builder @@ -312,14 +313,32 @@ def test_build_regime_mask_none_when_no_filter(): assert StrategyBacktestService._build_regime_mask(("2026-01-01",), None, None) is None -def test_build_regime_mask_none_when_no_data(tmp_path): - """无 regime 历史数据 → 返回 None(不阻断回测)。""" +def test_build_regime_mask_fails_when_no_data(tmp_path): + """启用过滤但无 regime 历史数据时必须阻止回测。""" from app.backtest.strategy import StrategyBacktestService - mask = StrategyBacktestService._build_regime_mask( - ("2026-01-01", "2026-01-02"), {"states": ["strong"]}, tmp_path, - ) - assert mask is None + with pytest.raises(ValueError, match="市场环境数据为空"): + StrategyBacktestService._build_regime_mask( + ("2026-01-01", "2026-01-02"), {"states": ["strong"]}, tmp_path, + ) + + +def test_build_regime_mask_fails_when_required_t1_date_is_missing(tmp_path): + """正式区间内任一入场日缺少 T-1 环境时必须阻止回测。""" + from app.backtest.strategy import StrategyBacktestService + + regime_builder.upsert_regime_history(tmp_path, pl.DataFrame({ + "date": [date(2026, 1, 1)], + "state": ["strong"], + "score": [85], + })) + + with pytest.raises(ValueError, match="缺少前一交易日环境"): + StrategyBacktestService._build_regime_mask( + ("2026-01-01", "2026-01-02", "2026-01-03"), + {"states": ["strong"]}, + tmp_path, + ) def test_build_regime_mask_first_day_allowed(tmp_path): @@ -336,4 +355,3 @@ def test_build_regime_mask_first_day_allowed(tmp_path): ) # 1/1 首日 → True; 1/2 由 1/1(weak) → False assert mask.tolist() == [True, False] - diff --git a/frontend/src/lib/storage.ts b/frontend/src/lib/storage.ts index 1c14271..b6f9457 100644 --- a/frontend/src/lib/storage.ts +++ b/frontend/src/lib/storage.ts @@ -113,6 +113,8 @@ export const storage = { mode: 'position' | 'full' holdingDays: string minuteFill?: boolean + regimeStates?: string[] + regimeMinScore?: number | '' params?: Record overrides?: Record strategyConfigSignature?: string diff --git a/frontend/src/pages/backtest/StrategyBacktest.tsx b/frontend/src/pages/backtest/StrategyBacktest.tsx index 8f1d4b0..ae155ca 100644 --- a/frontend/src/pages/backtest/StrategyBacktest.tsx +++ b/frontend/src/pages/backtest/StrategyBacktest.tsx @@ -947,8 +947,8 @@ export function StrategyBacktest() { const [holdingDays, setHoldingDays] = useState(saved?.holdingDays ?? '5') const [highGranularity, setHighGranularity] = useState(saved?.minuteFill ?? false) // 市场环境过滤(空=不过滤) - const [regimeStates, setRegimeStates] = useState([]) - const [regimeMinScore, setRegimeMinScore] = useState('') + const [regimeStates, setRegimeStates] = useState(saved?.regimeStates ?? []) + const [regimeMinScore, setRegimeMinScore] = useState(saved?.regimeMinScore ?? '') const [settingsOpen, setSettingsOpen] = useState(false) // 分钟K成交价细化: 不改变信号日或成交日, 需 Pro+ 分钟K能力 const { data: caps } = useCapabilities() @@ -1097,6 +1097,8 @@ export function StrategyBacktest() { mode: simMode, holdingDays, minuteFill: highGranularity, + regimeStates, + regimeMinScore, params: strategyParams, overrides, strategyConfigSignature: strategyDetail.data @@ -1373,6 +1375,18 @@ export function StrategyBacktest() { const resultStartDate = result?.config?.start ?? result?.equity_curve?.[0]?.date ?? start const resultEndDate = result?.config?.end ?? result?.equity_curve?.[result.equity_curve.length - 1]?.date ?? end const resultTradeDays = result?.equity_curve?.length ?? 0 + const resultRegimeFilter = result?.config?.regime_filter as { + states?: string[] + min_score?: number + } | null | undefined + const resultRegimeSummary = resultRegimeFilter + ? [ + resultRegimeFilter.states?.length + ? resultRegimeFilter.states.map(state => REGIME_STATE_LABELS[state as keyof typeof REGIME_STATE_LABELS] ?? state).join('/') + : null, + resultRegimeFilter.min_score != null ? `最低 ${resultRegimeFilter.min_score} 分` : null, + ].filter(Boolean).join(' · ') + : '' const selectionStats = result?.stats?.selection as Record | undefined const selectionStages = selectionStats ? [ @@ -1925,6 +1939,12 @@ export function StrategyBacktest() {
{result.strategy_info?.name ?? '策略'} 全量模拟 + {resultRegimeSummary && ( + + + 环境 {resultRegimeSummary} + + )} 持有 {result.config?.holding_days ?? 5} 天 {String(result.config?.start).slice(0,10)} ~ {String(result.config?.end).slice(0,10)} @@ -1995,6 +2015,12 @@ export function StrategyBacktest() { {SRC_MAP[result.strategy_info.source] ?? ''} )} + {resultRegimeSummary && ( + + + 环境 {resultRegimeSummary} + + )}
{/* 叠加策略: 子策略构成归因 */} {result.strategy_info.composite_children && result.strategy_info.composite_children.length > 0 && (