From 4cb30e48aad0eed69774093627130fb2a0c819da Mon Sep 17 00:00:00 2001 From: shy3130 <415333856@qq.com> Date: Thu, 3 Sep 2026 13:44:43 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20screener=20JIT=20=E9=80=8F=E4=BC=A0=20tu?= =?UTF-8?q?rnover=5Frate=20(#187)=20+=20PullScheduler=20=E7=BA=BF=E7=A8=8B?= =?UTF-8?q?=E5=AE=89=E5=85=A8=20(#203)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - #187: _compute_enriched_full/_load_enriched_history 的 warmup 读取 白名单补上存储列 turnover_rate —— instruments 不可用无从重算时, 自定义 SQL 用该列做条件会 Binder Error 被吞成静默空结果 - #203: refresh 的 _tasks 增删 diff 从调用方线程移进主循环闭包, 消除 TOCTOU 窗口; 行为测试覆盖增删/幂等 #187 测试已在修复前代码上验证会失败 (instruments 为空时列丢失)。 --- backend/app/services/ext_pull.py | 26 +++---- backend/app/services/screener.py | 7 +- backend/tests/test_ext_pull_refresh.py | 84 +++++++++++++++++++++ backend/tests/test_screener_jit_turnover.py | 78 +++++++++++++++++++ 4 files changed, 179 insertions(+), 16 deletions(-) create mode 100644 backend/tests/test_ext_pull_refresh.py create mode 100644 backend/tests/test_screener_jit_turnover.py diff --git a/backend/app/services/ext_pull.py b/backend/app/services/ext_pull.py index d1e777c..a57054d 100644 --- a/backend/app/services/ext_pull.py +++ b/backend/app/services/ext_pull.py @@ -221,22 +221,21 @@ class PullScheduler: configs = store.load_all() active_ids: set[str] = set() - new_configs: list[ExtConfig] = [] + enabled_configs: list[ExtConfig] = [] for config in configs: if not config.pull or not config.pull.enabled or not config.pull.url: continue active_ids.add(config.id) - if config.id not in self._tasks: - new_configs.append(config) + enabled_configs.append(config) - # 需要移除的 id (快照当前 task 字典的键, 避免遍历时改字典) - remove_ids = [cid for cid in list(self._tasks) if cid not in active_ids] - - # 所有对 _tasks 的修改都提交到主循环里执行, 保证线程安全 + # 对 _tasks 的一切读判断 (含增删 diff) 都放进主循环闭包里执行: + # refresh 可能从工作线程调用, 若在调用方线程读 _tasks 再把决策 + # 提交回主循环, 两步之间主循环可能已改动字典 (TOCTOU, #203)。 + # 此处只携带与 _tasks 无关的 config 数据跨线程。 def _apply() -> None: - for config in new_configs: - if config.id not in self._tasks: # 二次校验, 防重复 + for config in enabled_configs: + if config.id not in self._tasks: self._tasks[config.id] = self._loop.create_task( self._run_loop(config) ) @@ -244,11 +243,10 @@ class PullScheduler: "PullScheduler: scheduled %s (every %d min)", config.id, config.pull.schedule_minutes, ) - for cid in remove_ids: - task = self._tasks.pop(cid, None) - if task is not None: - task.cancel() - logger.info("PullScheduler: removed %s", cid) + for cid in [c for c in self._tasks if c not in active_ids]: + task = self._tasks.pop(cid) + task.cancel() + logger.info("PullScheduler: removed %s", cid) self._submit(_apply) diff --git a/backend/app/services/screener.py b/backend/app/services/screener.py index 9e1ae0e..0db6966 100644 --- a/backend/app/services/screener.py +++ b/backend/app/services/screener.py @@ -154,8 +154,10 @@ class ScreenerService: # 加载 warmup 历史 (目标日期前 ~120 天) enriched_dir = self.repo.store.data_dir / self._enriched_dirname start = target_date - timedelta(days=150) + # turnover_rate 是 enriched 存储列, 必须随行透传: 否则即时计算后该列 + # 丢失, 自定义 SQL 用它做条件会 Binder Error 被吞成空结果 (#187) read_cols = ["symbol", "date", "open", "high", "low", "close", "volume", - "amount", "raw_close", "raw_high", "raw_low"] + "amount", "raw_close", "raw_high", "raw_low", "turnover_rate"] try: lf = ( @@ -245,8 +247,9 @@ class ScreenerService: start = target_date - timedelta(days=min((lookback_days + warmup) * 2, 180)) enriched_dir = self.repo.store.data_dir / self._enriched_dirname + # 同 _compute_enriched_full: turnover_rate 存储列随行透传 (#187) read_cols = ["symbol", "date", "open", "high", "low", "close", "volume", - "amount", "raw_close", "raw_high", "raw_low"] + "amount", "raw_close", "raw_high", "raw_low", "turnover_rate"] try: lf = ( diff --git a/backend/tests/test_ext_pull_refresh.py b/backend/tests/test_ext_pull_refresh.py new file mode 100644 index 0000000..6f1fb26 --- /dev/null +++ b/backend/tests/test_ext_pull_refresh.py @@ -0,0 +1,84 @@ +"""#203 回归: PullScheduler.refresh 的 _tasks 增删 diff 全部在主循环闭包内完成。 + +refresh 可能从 FastAPI worker 线程调用; 旧实现先在调用方线程读 _tasks +做 diff 再提交闭包, 两步之间主循环可能已改动字典 (TOCTOU)。 +本测试验证行为正确性 (增/删/幂等); 竞态本身无法确定性复现。 +""" +from __future__ import annotations + +import asyncio +import contextlib + +import pytest + +from app.services import ext_pull +from app.services.ext_data import ExtConfig, PullConfig +from app.services.ext_pull import PullScheduler + + +def _cfg(cid: str, enabled: bool = True) -> ExtConfig: + return ExtConfig( + id=cid, label=cid, mode="snapshot", fields=[], + pull=PullConfig( + enabled=enabled, url=f"https://example.test/{cid}", + schedule_minutes=60, + ) if enabled else None, + ) + + +@pytest.fixture() +def scheduler(monkeypatch, tmp_path): + """跑在真实事件循环上的调度器; _run_loop 打桩避免真实网络。""" + loop = asyncio.new_event_loop() + s = PullScheduler() + s._loop = loop + s._running = True + + async def _idle(self, config): # 打桩: 挂起不退出, 等 cancel (self 因类级打桩传入) + with contextlib.suppress(asyncio.CancelledError): + await asyncio.Event().wait() + + monkeypatch.setattr(PullScheduler, "_run_loop", _idle) + configs: list[ExtConfig] = [] + + def _load_all(self): + return list(configs) + + monkeypatch.setattr(ext_pull.ExtConfigStore, "load_all", _load_all) + yield s, loop, configs + for t in list(s._tasks.values()): + t.cancel() + loop.run_until_complete(asyncio.sleep(0)) # 让取消送达协程 + loop.close() + + +def _drain(loop): + loop.run_until_complete(asyncio.sleep(0)) + + +def test_refresh_schedules_enabled_configs(scheduler, tmp_path) -> None: + s, loop, configs = scheduler + configs.extend([_cfg("a"), _cfg("b"), _cfg("c", enabled=False)]) + + s.refresh(tmp_path) # data_dir 仅透传, 配置来自打桩的 load_all + _drain(loop) + assert set(s._tasks) == {"a", "b"} + + # 幂等: 重复 refresh 不重复建任务 + s.refresh(tmp_path) + _drain(loop) + assert set(s._tasks) == {"a", "b"} + + +def test_refresh_removes_disabled_configs(scheduler, tmp_path) -> None: + s, loop, configs = scheduler + configs.extend([_cfg("a"), _cfg("b")]) + s.refresh(tmp_path) + _drain(loop) + assert set(s._tasks) == {"a", "b"} + + # b 被禁用 → 下次 refresh 移除其任务 + configs[:] = [_cfg("a")] + s.refresh(tmp_path) + _drain(loop) + assert set(s._tasks) == {"a"} diff --git a/backend/tests/test_screener_jit_turnover.py b/backend/tests/test_screener_jit_turnover.py new file mode 100644 index 0000000..d7b2c35 --- /dev/null +++ b/backend/tests/test_screener_jit_turnover.py @@ -0,0 +1,78 @@ +"""#187 回归: screener JIT 即时计算路径不得丢失 turnover_rate 存储列。 + +历史日期走 _compute_enriched_full (scan_parquet + compute_indicators), +warmup 读取白名单漏掉 turnover_rate 时, 即时计算后该列丢失 —— 自定义 SQL +用 turnover_rate 做条件的请求在 DuckDB 注册视图里找不到列, Binder Error +被 except 吞掉返回空结果 (无任何报错提示)。 +""" +from __future__ import annotations + +from datetime import date, timedelta +from unittest.mock import MagicMock + +import polars as pl + +from app.services.screener import ScreenerService + + +def _write_enriched(tmp_path, days: int, turnover_by_day: dict[str, float]) -> None: + base = tmp_path / "kline_daily_enriched" + base.mkdir(parents=True, exist_ok=True) + start = date(2026, 9, 1) - timedelta(days=days) + for i in range(days + 1): + d = start + timedelta(days=i) + part = base / f"date={d.isoformat()}" + part.mkdir(exist_ok=True) + pl.DataFrame( + { + "symbol": ["600000.SH", "000001.SZ"], + "date": [d, d], + "open": [10.0, 20.0], + "high": [11.0, 21.0], + "low": [9.0, 19.0], + "close": [10.5, 20.5], + "volume": [100.0, 200.0], + "amount": [1050.0, 4100.0], + "raw_close": [10.5, 20.5], + "raw_high": [11.0, 21.0], + "raw_low": [9.0, 19.0], + "turnover_rate": [turnover_by_day.get(d.isoformat(), 1.0), 2.0], + "consecutive_limit_ups": [0, 0], + "consecutive_limit_downs": [0, 0], + } + ).write_parquet(part / "part.parquet") + + +def _service(tmp_path, instruments: pl.DataFrame | None = None) -> ScreenerService: + repo = MagicMock() + repo.store.data_dir = tmp_path + # 最新日缓存与 repo 级历史缓存均未命中 → 走 scan_parquet + 即时计算慢路径 + repo.get_enriched_latest_asset.return_value = (pl.DataFrame(), None) + repo.get_enriched_history.return_value = None + # instruments 不可用 (空维表) 时 turnover_rate 无从重算, 只能靠存储列透传 + repo.get_instruments_asset.return_value = instruments or pl.DataFrame() + repo.get_historical_shares.return_value = pl.DataFrame() + return ScreenerService(repo, asset_type="stock") + + +def test_historical_jit_frame_keeps_turnover_rate(tmp_path) -> None: + target = date(2026, 9, 1) + _write_enriched(tmp_path, 30, {target.isoformat(): 5.5}) + svc = _service(tmp_path) # 无 instruments → 无重算路径, 纯存储列透传 + + df = svc._load_enriched_for_date(target) + assert not df.is_empty() + assert "turnover_rate" in df.columns, "JIT 即时计算后 turnover_rate 列不应丢失" + # 目标日的值来自存储列透传, 不是置 null + row = df.filter(pl.col("symbol") == "600000.SH").row(0, named=True) + assert row["turnover_rate"] == 5.5 + + +def test_custom_sql_can_filter_on_turnover_rate(tmp_path) -> None: + target = date(2026, 9, 1) + _write_enriched(tmp_path, 30, {target.isoformat(): 5.5}) + svc = _service(tmp_path) + + result = svc.run(target, ["turnover_rate > 3"], limit=10) + # 600000 (5.5) 命中; 000001 恒为 2.0 被过滤 — 条件真正生效而非空结果 + assert [r["symbol"] for r in result.rows] == ["600000.SH"]