diff --git a/backend/app/api/screener.py b/backend/app/api/screener.py index d81c962..4ffb30f 100644 --- a/backend/app/api/screener.py +++ b/backend/app/api/screener.py @@ -283,10 +283,23 @@ def run_preset(req: PresetRequest, request: Request): if dl is None and overrides and "display_limit" in overrides: dl = 0 + engine = getattr(request.app.state, "strategy_engine", None) + # 内置策略 if req.strategy_id in PRESET_STRATEGIES: + filter_fn = None + if req.asset_type == "stock" and engine and engine.has(req.strategy_id): + filter_fn = engine.get(req.strategy_id).filter_fn try: - result = svc.run_preset(req.strategy_id, as_of=as_of, pool=req.pool, basic_filter=bf, display_limit=dl) + result = svc.run_preset( + req.strategy_id, + as_of=as_of, + pool=req.pool, + basic_filter=bf, + filter_fn=filter_fn, + strategy_params=overrides.get("params") if overrides else None, + display_limit=dl, + ) except ValueError as e: raise HTTPException(status_code=404, detail=str(e)) from e safe_data = _safe(asdict(result)) @@ -294,7 +307,6 @@ def run_preset(req: PresetRequest, request: Request): return _result_with_ext(safe_data, ext_values) # 自定义/AI 策略 — 通过 StrategyEngine 执行 - engine = getattr(request.app.state, "strategy_engine", None) if not engine: raise HTTPException(status_code=404, detail=f"策略引擎未初始化或策略 {req.strategy_id} 不存在") @@ -468,7 +480,16 @@ def run_all(request: Request, body: Optional[dict] = None): dl = 0 if sid in PRESET_STRATEGIES: - r = svc.run_preset(sid, as_of=as_of, precomputed=precomputed, basic_filter=bf, display_limit=dl) + filter_fn = engine.get(sid).filter_fn if engine and engine.has(sid) else None + r = svc.run_preset( + sid, + as_of=as_of, + precomputed=precomputed, + basic_filter=bf, + filter_fn=filter_fn, + strategy_params=overrides.get("params") if overrides else None, + display_limit=dl, + ) else: r = engine.run( sid, as_of, overrides=overrides or None, diff --git a/backend/app/services/screener.py b/backend/app/services/screener.py index a38aef9..e8f71c8 100644 --- a/backend/app/services/screener.py +++ b/backend/app/services/screener.py @@ -9,6 +9,7 @@ from __future__ import annotations import logging import time +from collections.abc import Callable from dataclasses import dataclass, field from datetime import date, timedelta @@ -525,6 +526,8 @@ class ScreenerService: pool: list[str] | None = None, precomputed: pl.DataFrame | None = None, basic_filter: dict | None = None, + filter_fn: Callable[[pl.DataFrame, dict], pl.Expr] | None = None, + strategy_params: dict | None = None, display_limit: int | None = None, ) -> ScreenerResult: """预设策略选股 — 从 enriched 读取预计算好的指标列后过滤。 @@ -532,6 +535,7 @@ class ScreenerService: - precomputed 不为空: 直接复用(run_all 场景) - precomputed 为空: 从 enriched 读目标日期 - basic_filter: 用户保存的基础参数过滤(boards、价格等) + - filter_fn/strategy_params: 内置策略文件的参数化过滤;未传时兼容旧预设表达式 """ t0 = time.perf_counter() @@ -555,7 +559,8 @@ class ScreenerService: df = self._apply_basic_filter(df, basic_filter) # 应用策略过滤 - df = df.filter(strat["filter"]) + filter_expr = filter_fn(df, strategy_params or {}) if filter_fn else strat["filter"] + df = df.filter(filter_expr) # 应用 pool if pool: diff --git a/backend/tests/test_screener_builtin_params.py b/backend/tests/test_screener_builtin_params.py new file mode 100644 index 0000000..c4aa425 --- /dev/null +++ b/backend/tests/test_screener_builtin_params.py @@ -0,0 +1,166 @@ +from __future__ import annotations + +import random +import types +from datetime import date +from pathlib import Path +from typing import ClassVar + +import polars as pl + +from app.api import screener as screener_api +from app.services.screener import PRESET_STRATEGIES, ScreenerResult, ScreenerService +from app.strategy.engine import StrategyEngine + + +def _builtin_engine() -> StrategyEngine: + builtin_dir = Path(__file__).parents[1] / "app" / "strategy" / "builtin" + return StrategyEngine( + enriched_loader=lambda _as_of: pl.DataFrame(), + strategy_dirs=[builtin_dir], + ) + + +def _comparison_frame() -> pl.DataFrame: + rng = random.Random(20260715) + size = 200 + return pl.DataFrame({ + "symbol": [f"S{i:03d}" for i in range(size)], + "close": [rng.uniform(5, 100) for _ in range(size)], + "open": [rng.uniform(5, 100) for _ in range(size)], + "ma5": [rng.uniform(5, 100) for _ in range(size)], + "ma10": [rng.uniform(5, 100) for _ in range(size)], + "ma20": [rng.uniform(5, 100) for _ in range(size)], + "ma60": [rng.uniform(5, 100) for _ in range(size)], + "vol_ratio_5d": [rng.uniform(0, 4) for _ in range(size)], + "momentum_20d": [rng.uniform(-0.5, 0.5) for _ in range(size)], + "momentum_60d": [rng.uniform(-0.5, 0.5) for _ in range(size)], + "annual_vol_20d": [rng.uniform(0, 0.6) for _ in range(size)], + "change_pct": [rng.uniform(-0.1, 0.1) for _ in range(size)], + "rsi_14": [rng.uniform(0, 100) for _ in range(size)], + "consecutive_limit_ups": [rng.randrange(0, 5) for _ in range(size)], + "signal_n_day_high": [rng.choice([True, False, None]) for _ in range(size)], + "signal_ma_golden_5_20": [rng.choice([True, False, None]) for _ in range(size)], + "signal_macd_golden": [rng.choice([True, False, None]) for _ in range(size)], + "signal_ma20_breakout": [rng.choice([True, False, None]) for _ in range(size)], + "signal_limit_up": [rng.choice([True, False, None]) for _ in range(size)], + "signal_boll_breakout_upper": [rng.choice([True, False, None]) for _ in range(size)], + "signal_n_day_low": [rng.choice([True, False, None]) for _ in range(size)], + }) + + +def test_builtin_default_filters_match_legacy_presets(): + engine = _builtin_engine() + df = _comparison_frame() + + for strategy_id, preset in PRESET_STRATEGIES.items(): + strategy = engine.get(strategy_id) + defaults = {item["id"]: item["default"] for item in strategy.meta["params"]} + expected = df.filter(preset["filter"])["symbol"].to_list() + actual = df.filter(strategy.filter_fn(df, defaults))["symbol"].to_list() + assert actual == expected, strategy_id + + +def test_run_preset_applies_numeric_and_boolean_strategy_params(): + engine = _builtin_engine() + strategy = engine.get("trend_breakout") + df = pl.DataFrame({ + "symbol": ["A", "B", "C"], + "close": [10.0, 10.0, 10.0], + "ma60": [9.0, 9.0, 9.0], + "signal_n_day_high": [True, True, True], + "vol_ratio_5d": [1.0, 2.5, 3.5], + "momentum_60d": [0.1, 0.2, 0.3], + }) + service = ScreenerService(types.SimpleNamespace()) + + strict = service.run_preset( + "trend_breakout", + as_of=date(2026, 7, 15), + precomputed=df, + filter_fn=strategy.filter_fn, + strategy_params={"vol_ratio_min": 3.0}, + ) + volume_disabled = service.run_preset( + "trend_breakout", + as_of=date(2026, 7, 15), + precomputed=df, + filter_fn=strategy.filter_fn, + strategy_params={"use_volume_filter": False}, + ) + + assert [row["symbol"] for row in strict.rows] == ["C"] + assert [row["symbol"] for row in volume_disabled.rows] == ["C", "B", "A"] + + +class _CapturingScreenerService: + calls: ClassVar[list[dict]] = [] + + def __init__(self, repo, asset_type="stock"): + self.repo = repo + self.asset_type = asset_type + + def latest_date(self): + return date(2026, 7, 15) + + def _load_enriched_for_date(self, _as_of): + return pl.DataFrame({"symbol": ["A"]}) + + def run_preset(self, strategy_id, as_of, **kwargs): + self.calls.append({"strategy_id": strategy_id, **kwargs}) + return ScreenerResult(as_of=as_of, strategy=strategy_id) + + +def _api_request(tmp_path, engine): + repo = types.SimpleNamespace(store=types.SimpleNamespace(data_dir=tmp_path)) + state = types.SimpleNamespace(repo=repo, strategy_engine=engine) + return types.SimpleNamespace(app=types.SimpleNamespace(state=state)) + + +def test_single_run_passes_saved_params_to_builtin_filter(monkeypatch, tmp_path): + engine = _builtin_engine() + request = _api_request(tmp_path, engine) + _CapturingScreenerService.calls = [] + monkeypatch.setattr(screener_api, "ScreenerService", _CapturingScreenerService) + monkeypatch.setattr( + screener_api.strategy_config, + "load_override", + lambda *_args: {"params": {"vol_ratio_min": 3.0}}, + ) + monkeypatch.setattr(screener_api, "_load_ext_value_maps", lambda *_args: {}) + monkeypatch.setattr(screener_api, "_update_cache_strategy", lambda *_args: None) + + screener_api.run_preset( + screener_api.PresetRequest( + strategy_id="trend_breakout", + as_of=date(2026, 7, 15), + ), + request, + ) + + call = _CapturingScreenerService.calls[0] + assert call["filter_fn"] is not None + assert call["strategy_params"] == {"vol_ratio_min": 3.0} + + +def test_batch_run_passes_saved_params_to_builtin_filter(monkeypatch, tmp_path): + engine = _builtin_engine() + request = _api_request(tmp_path, engine) + _CapturingScreenerService.calls = [] + monkeypatch.setattr(screener_api, "ScreenerService", _CapturingScreenerService) + monkeypatch.setattr( + screener_api.strategy_config, + "list_overrides", + lambda *_args: {"trend_breakout": {"params": {"vol_ratio_min": 3.0}}}, + ) + monkeypatch.setattr(screener_api.strategy_cache, "write_cache", lambda *_args: None) + monkeypatch.setattr(screener_api, "_load_ext_value_maps", lambda *_args: {}) + + screener_api.run_all( + request, + body={"as_of": "2026-07-15", "strategy_ids": ["trend_breakout"]}, + ) + + call = _CapturingScreenerService.calls[0] + assert call["filter_fn"] is not None + assert call["strategy_params"] == {"vol_ratio_min": 3.0}