From 89115e6f737f2ec48f6013f5e1e8da53de6d65af Mon Sep 17 00:00:00 2001 From: shy3130 <415333856@qq.com> Date: Wed, 15 Jul 2026 19:00:56 +0800 Subject: [PATCH] fix(strategy): apply saved parameters during execution --- backend/app/strategy/engine.py | 6 +-- backend/tests/test_strategy_saved_params.py | 48 +++++++++++++++++++++ 2 files changed, 51 insertions(+), 3 deletions(-) create mode 100644 backend/tests/test_strategy_saved_params.py diff --git a/backend/app/strategy/engine.py b/backend/app/strategy/engine.py index a031d84..63ce7df 100644 --- a/backend/app/strategy/engine.py +++ b/backend/app/strategy/engine.py @@ -284,16 +284,16 @@ class StrategyEngine: strategy_id: 策略 ID as_of: 选股日期 pool: 限定股票池 - params: 策略参数 (用户在设置面板调的值) - overrides: 用户覆盖配置 (basic_filter/scoring/stop_loss 等) + params: 本次执行显式传入的策略参数 + overrides: 用户覆盖配置 (params/basic_filter/scoring/stop_loss 等) precomputed: 已加载的 enriched 数据 (run_all 场景复用) precomputed_history: 已加载的历史窗口数据 (run_all 场景复用) """ t0 = time.perf_counter() s = self.get(strategy_id) - params = params or {} overrides = overrides or {} + params = {**(overrides.get("params") or {}), **(params or {})} # 加载数据。普通策略只读目标日期;声明 filter_history 的策略读取历史窗口。 if s.filter_history_fn: diff --git a/backend/tests/test_strategy_saved_params.py b/backend/tests/test_strategy_saved_params.py new file mode 100644 index 0000000..2e1468c --- /dev/null +++ b/backend/tests/test_strategy_saved_params.py @@ -0,0 +1,48 @@ +from datetime import date + +import polars as pl + +from app.strategy.engine import StrategyDef, StrategyEngine + + +def _make_engine() -> StrategyEngine: + df = pl.DataFrame({"symbol": ["A", "B", "C"], "value": [1, 2, 3]}) + engine = StrategyEngine(enriched_loader=lambda _as_of: df) + engine._strategies["saved_params"] = StrategyDef( + meta={"id": "saved_params", "scoring": {}, "limit": 100}, + basic_filter={"enabled": False}, + entry_signals=[], + exit_signals=[], + stop_loss=None, + trailing_stop=None, + trailing_take_profit_activate=None, + trailing_take_profit_drawdown=None, + max_hold_days=None, + alerts=[], + filter_fn=lambda _df, params: pl.col("value") >= params.get("min_value", 1), + filter_history_fn=None, + lookback_days=1, + source="custom", + ) + return engine + + +def test_run_applies_saved_strategy_params(): + result = _make_engine().run( + "saved_params", + date(2026, 7, 15), + overrides={"params": {"min_value": 2}}, + ) + + assert [row["symbol"] for row in result.rows] == ["B", "C"] + + +def test_explicit_params_override_saved_strategy_params(): + result = _make_engine().run( + "saved_params", + date(2026, 7, 15), + params={"min_value": 3}, + overrides={"params": {"min_value": 2}}, + ) + + assert [row["symbol"] for row in result.rows] == ["C"]