From 40c2468cbc447aa4900310e09e128e811e75c967 Mon Sep 17 00:00:00 2001 From: kevin9327 <5299031+kevin9327@users.noreply.github.com> Date: Thu, 10 Sep 2026 19:35:32 +0900 Subject: [PATCH] =?UTF-8?q?fix(screener):=20=E6=B6=A8=E5=81=9C=E6=A2=AF?= =?UTF-8?q?=E9=98=9F=E7=9A=84=E6=97=B6=E5=BA=8F=E6=89=A9=E5=B1=95=E5=88=97?= =?UTF-8?q?=E5=8F=AA=E5=8F=96=E6=9C=80=E6=96=B0=E5=88=86=E5=8C=BA,=20?= =?UTF-8?q?=E4=B8=8D=E5=86=8D=E6=94=BE=E5=A4=A7=E8=A1=8C=E6=95=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ext_{config_id} 视图对 timeseries 模式覆盖 timeseries/**/*.parquet 全部分区 (app/api/ext_data._refresh_views), 一只票在 N 天快照里就有 N 行。梯队直接 LEFT JOIN 该视图, 同一只涨停股被复制 N 份, 各档 count 一并放大 N 倍。 改为与自选股列表 (app/api/watchlist) 同口径: 有配置时走 _read_ext_dataframe 取最新分区, 再按 symbol 去重后 JOIN; 无配置时保留视图查询兜底。 --- backend/app/api/screener.py | 40 +++-- .../tests/test_limit_ladder_ext_timeseries.py | 151 ++++++++++++++++++ 2 files changed, 179 insertions(+), 12 deletions(-) create mode 100644 backend/tests/test_limit_ladder_ext_timeseries.py diff --git a/backend/app/api/screener.py b/backend/app/api/screener.py index 63daa55..faa14e2 100644 --- a/backend/app/api/screener.py +++ b/backend/app/api/screener.py @@ -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: - 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") + # 扩展时序数据必须只取最新分区; 否则一个 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()) + 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 diff --git a/backend/tests/test_limit_ladder_ext_timeseries.py b/backend/tests/test_limit_ladder_ext_timeseries.py new file mode 100644 index 0000000..b79f2fc --- /dev/null +++ b/backend/tests/test_limit_ladder_ext_timeseries.py @@ -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"] == "人工智能"