fix: screener JIT 透传 turnover_rate (#187) + PullScheduler 线程安全 (#203)

- #187: _compute_enriched_full/_load_enriched_history 的 warmup 读取
  白名单补上存储列 turnover_rate —— instruments 不可用无从重算时,
  自定义 SQL 用该列做条件会 Binder Error 被吞成静默空结果
- #203: refresh 的 _tasks 增删 diff 从调用方线程移进主循环闭包,
  消除 TOCTOU 窗口; 行为测试覆盖增删/幂等

#187 测试已在修复前代码上验证会失败 (instruments 为空时列丢失)。
This commit is contained in:
shy3130
2026-09-03 13:45:12 +08:00
parent 91fef6d793
commit 4cb30e48aa
4 changed files with 179 additions and 16 deletions
+12 -14
View File
@@ -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)
+5 -2
View File
@@ -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 = (
+84
View File
@@ -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"}
@@ -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"]