Merge pull request #295 from kevin9327/fix/limit-ladder-ext-timeseries-latest

fix(screener): 涨停梯队的时序扩展列只取最新分区, 不再放大行数
This commit is contained in:
wshy
2026-09-11 09:02:11 +08:00
committed by GitHub
2 changed files with 179 additions and 12 deletions
+25 -9
View File
@@ -933,32 +933,48 @@ def limit_ladder(
if ext_specs:
db = repo.store.db
data_dir = repo.store.data_dir
from app.api.ext_data import _read_ext_dataframe
from app.services.ext_data import ExtConfigStore
ext_store = ExtConfigStore(data_dir)
configs = {c.id: c for c in ext_store.load_all()}
def _dedup_ext(frame: pl.DataFrame, field: str, out_col: str) -> pl.DataFrame | None:
"""(symbol, 字段) 两列并按 symbol 去重; 缺列时返回 None。"""
if frame.is_empty() or "symbol" not in frame.columns or field not in frame.columns:
return None
return (
frame
.select(["symbol", field])
.unique(subset=["symbol"], keep="last")
.rename({field: out_col})
)
for config_id, field_name in ext_specs:
view_name = f"ext_{config_id}"
ext_col_name = f"{config_id}__{field_name}"
try:
# 扩展时序数据必须只取最新分区; 否则一个 symbol 会按历史分区数被 JOIN 放大
# (ext_{id} 视图覆盖 timeseries/**), 与自选股列表同口径。
cfg = configs.get(config_id)
if cfg:
ext_df, _ = _read_ext_dataframe(cfg, data_dir)
else:
ext_df = pl.from_arrow(db.query(
f"SELECT symbol, {quote_ident(field_name)} FROM {view_name}"
).arrow())
if not ext_df.is_empty() and "symbol" in ext_df.columns:
ext_df = ext_df.rename({field_name: ext_col_name})
df = df.join(ext_df.select(["symbol", ext_col_name]), on="symbol", how="left")
joined = _dedup_ext(ext_df, field_name, ext_col_name)
if joined is not None:
df = df.join(joined, on="symbol", how="left")
ext_col_names.append(ext_col_name)
except Exception:
cfg = configs.get(config_id)
if cfg:
try:
from app.api.ext_data import _parquet_glob
glob = _parquet_glob(cfg, data_dir)
ext_df = pl.read_parquet(glob)
if not ext_df.is_empty() and "symbol" in ext_df.columns and field_name in ext_df.columns:
ext_df = ext_df.select(["symbol", field_name]).rename({field_name: ext_col_name})
df = df.join(ext_df, on="symbol", how="left")
ext_df, _ = _read_ext_dataframe(cfg, data_dir)
joined = _dedup_ext(ext_df, field_name, ext_col_name)
if joined is not None:
df = df.join(joined, on="symbol", how="left")
ext_col_names.append(ext_col_name)
except Exception:
pass
@@ -0,0 +1,151 @@
"""涨停梯队挂时序扩展列时, 一只票不能按历史分区数被 JOIN 放大。
ext_{config_id} DuckDB 视图对 timeseries 模式覆盖 timeseries/**/*.parquet 全部
分区 (见 app/api/ext_data._refresh_views), 一只票在 N 天快照里就有 N 行。
自选股列表 (app/api/watchlist) 对同一场景先走 _read_ext_dataframe 取最新分区
再按 symbol 去重, 梯队这边直接查视图, 于是同一只涨停股在梯队里出现 N 次、
档位 count 被放大 N 倍。
"""
from __future__ import annotations
import json
from datetime import date
from types import SimpleNamespace
import polars as pl
import pytest
from app.api import screener as screener_api
from app.services.screener import ScreenerService
_AS_OF = date(2026, 9, 10)
class _FakeQuery:
def __init__(self, df: pl.DataFrame) -> None:
self._df = df
def arrow(self):
return self._df.to_arrow()
class _FakeDB:
"""模拟 ext_{id} 视图: timeseries 模式下返回全部分区的行。"""
def __init__(self, df: pl.DataFrame) -> None:
self._df = df
def query(self, sql: str) -> _FakeQuery:
return _FakeQuery(self._df)
class _NoDepth:
def get_sealed_map(self, _as_of, is_down: bool = False) -> dict:
return {}
def is_sealed_ready(self, _as_of) -> bool:
return False
def get_sealed_age(self, _as_of):
return None
def _write_timeseries_ext(data_dir, partitions: dict[str, pl.DataFrame]) -> pl.DataFrame:
"""写 timeseries 扩展表; 返回视图口径 (全部分区拼接) 的 DataFrame。"""
cfg_dir = data_dir / "ext_data" / "concept_ts"
cfg_dir.mkdir(parents=True, exist_ok=True)
cfg_dir.joinpath("config.json").write_text(
json.dumps({
"id": "concept_ts",
"label": "概念时序",
"mode": "timeseries",
"fields": [{"name": "concept", "dtype": "string", "label": "概念"}],
}, ensure_ascii=False),
encoding="utf-8",
)
for day, frame in partitions.items():
part = cfg_dir / "timeseries" / f"date={day}"
part.mkdir(parents=True, exist_ok=True)
frame.write_parquet(part / "part.parquet")
return pl.concat([partitions[d] for d in sorted(partitions)])
def _request(data_dir, view_df: pl.DataFrame):
repo = SimpleNamespace(
store=SimpleNamespace(data_dir=data_dir, db=_FakeDB(view_df)),
)
state = SimpleNamespace(repo=repo, depth_service=_NoDepth())
return SimpleNamespace(app=SimpleNamespace(state=state))
@pytest.fixture
def _stub_enriched(monkeypatch):
"""单只涨停股的 enriched 快照 + 空的昨日连板表。"""
snapshot = pl.DataFrame({
"symbol": ["600000.SH"],
"name": ["示例"],
"close": [11.0],
"change_pct": [0.1],
"signal_limit_up": [True],
"signal_limit_down": [False],
"signal_broken_limit_up": [False],
"consecutive_limit_ups": [2],
})
monkeypatch.setattr(
ScreenerService, "_load_enriched_for_date", lambda self, d: snapshot
)
monkeypatch.setattr(
ScreenerService, "load_prior_consecutive", lambda self, d, c: pl.DataFrame()
)
return snapshot
def test_timeseries_ext_column_does_not_duplicate_ladder_rows(tmp_path, _stub_enriched):
view_df = _write_timeseries_ext(tmp_path, {
"2026-09-08": pl.DataFrame({"symbol": ["600000.SH"], "concept": ["旧概念"]}),
"2026-09-09": pl.DataFrame({"symbol": ["600000.SH"], "concept": ["次新概念"]}),
"2026-09-10": pl.DataFrame({"symbol": ["600000.SH"], "concept": ["最新概念"]}),
})
payload = screener_api.limit_ladder(
_request(tmp_path, view_df),
as_of=_AS_OF,
direction="up",
ext_columns="concept_ts.concept",
)
tiers = payload["tiers"]
assert len(tiers) == 1
stocks = tiers[0]["stocks"]
assert [s["symbol"] for s in stocks] == ["600000.SH"], "同一只票被历史分区放大了"
assert tiers[0]["count"] == 1
# 取最新分区的值, 不是任意一期的历史值
assert stocks[0]["concept_ts__concept"] == "最新概念"
def test_snapshot_ext_column_still_joins(tmp_path, _stub_enriched):
"""snapshot 模式扩展表照常挂列 (回归保护)。"""
cfg_dir = tmp_path / "ext_data" / "concept_snap"
cfg_dir.mkdir(parents=True, exist_ok=True)
cfg_dir.joinpath("config.json").write_text(
json.dumps({
"id": "concept_snap",
"label": "概念快照",
"mode": "snapshot",
"fields": [{"name": "concept", "dtype": "string", "label": "概念"}],
}, ensure_ascii=False),
encoding="utf-8",
)
snap = pl.DataFrame({"symbol": ["600000.SH"], "concept": ["人工智能"]})
snap.write_parquet(cfg_dir / "part.parquet")
payload = screener_api.limit_ladder(
_request(tmp_path, snap),
as_of=_AS_OF,
direction="up",
ext_columns="concept_snap.concept",
)
stocks = payload["tiers"][0]["stocks"]
assert len(stocks) == 1
assert stocks[0]["concept_snap__concept"] == "人工智能"