From 6bb6eadb123bc9eacce5d52b6e26bbcf26f908a8 Mon Sep 17 00:00:00 2001 From: wshy Date: Thu, 9 Jul 2026 17:38:44 +0800 Subject: [PATCH] =?UTF-8?q?fix(storage):=20=E5=85=BC=E5=AE=B9=20parquet=20?= =?UTF-8?q?=E5=88=86=E5=8C=BA=20schema=20=E6=BC=94=E8=BF=9B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 兼容 quote_ts 等新增存储列导致的新旧 parquet 分区 schema 差异,恢复老用户 enriched 缓存刷新和看板/选股页面显示。 --- backend/app/backtest/engine.py | 3 +- backend/app/indicators/pipeline.py | 15 ++-- backend/app/parquet.py | 55 ++++++++++++++ backend/app/services/backtest.py | 3 +- backend/app/services/quote_service.py | 5 +- backend/app/services/screener.py | 5 +- backend/app/tickflow/repository.py | 17 +++-- backend/tests/test_parquet_schema_compat.py | 81 +++++++++++++++++++++ 8 files changed, 163 insertions(+), 21 deletions(-) create mode 100644 backend/app/parquet.py create mode 100644 backend/tests/test_parquet_schema_compat.py diff --git a/backend/app/backtest/engine.py b/backend/app/backtest/engine.py index d9be4d6..1bb0679 100644 --- a/backend/app/backtest/engine.py +++ b/backend/app/backtest/engine.py @@ -19,6 +19,7 @@ from typing import Literal import numpy as np import polars as pl +from app.parquet import scan_enriched_parquet from app.tickflow.repository import KlineRepository logger = logging.getLogger(__name__) @@ -223,7 +224,7 @@ class BacktestEngine: enriched_glob = str(self.repo.store.data_dir / enriched_dirname(asset_type) / "**" / "*.parquet") try: - lf = pl.scan_parquet(enriched_glob) + lf = scan_enriched_parquet(enriched_glob) if symbols is not None: lf = lf.filter(pl.col("symbol").is_in(symbols)) if columns is not None: diff --git a/backend/app/indicators/pipeline.py b/backend/app/indicators/pipeline.py index 3f7de8a..f7b27b9 100644 --- a/backend/app/indicators/pipeline.py +++ b/backend/app/indicators/pipeline.py @@ -22,6 +22,7 @@ from pathlib import Path import polars as pl from app.config import settings +from app.parquet import scan_daily_parquet, scan_enriched_parquet, scan_parquet_compat logger = logging.getLogger(__name__) @@ -938,7 +939,7 @@ def run_pipeline(data_dir: Path | None = None, # 加载 instruments (涨跌停+换手率需要) instruments = pl.DataFrame() try: - instruments = pl.scan_parquet(inst_glob, cast_options=_cast).collect() + instruments = scan_parquet_compat(inst_glob, cast_options=_cast).collect() except Exception as e: # noqa: BLE001 logger.warning("instruments 读取失败: %s", e) @@ -963,9 +964,9 @@ def run_pipeline(data_dir: Path | None = None, # 2. 为新日期计算 enriched (所有标的) if new_date_dirs: - raw_new = pl.scan_parquet(new_date_dirs[0] / "*.parquet", cast_options=_cast) + raw_new = scan_daily_parquet(new_date_dirs[0] / "*.parquet", cast_options=_cast) for nd in new_date_dirs[1:]: - raw_new = pl.concat([raw_new, pl.scan_parquet(nd / "*.parquet", cast_options=_cast)], how="diagonal_relaxed") + raw_new = pl.concat([raw_new, scan_daily_parquet(nd / "*.parquet", cast_options=_cast)], how="diagonal_relaxed") raw_new = raw_new.sort(["symbol", "date"]).collect(streaming=True) # 增量模式: 只算新日期, 但指标需要历史窗口 @@ -1013,7 +1014,7 @@ def run_pipeline(data_dir: Path | None = None, # 3. 受除权因子影响的个股: 重算全部已有日期 (累积因子链变了) if symbols: sym_set = set(symbols) - raw_sym = pl.scan_parquet(daily_glob, cast_options=_cast).sort(["symbol", "date"]) + raw_sym = scan_daily_parquet(daily_glob, cast_options=_cast).sort(["symbol", "date"]) raw_sym = raw_sym.filter(pl.col("symbol").is_in(list(sym_set))) raw_sym = raw_sym.collect(streaming=True) if not raw_sym.is_empty(): @@ -1053,7 +1054,7 @@ def run_pipeline(data_dir: Path | None = None, # ── 按 symbol 分批处理: 每只股只有 ~244 行, 无冗余计算 ── # 先获取全部 symbol 列表 - lf_all = pl.scan_parquet(daily_glob, cast_options=_cast) + lf_all = scan_daily_parquet(daily_glob, cast_options=_cast) if symbols: sym_set = set(symbols) lf_all = lf_all.filter(pl.col("symbol").is_in(list(sym_set))) @@ -1090,7 +1091,7 @@ def run_pipeline(data_dir: Path | None = None, batch_syms = all_symbols[batch_start:batch_end] # 只读取本批 symbol 的数据 - lf_batch = pl.scan_parquet(daily_glob, cast_options=_cast) + lf_batch = scan_daily_parquet(daily_glob, cast_options=_cast) lf_batch = lf_batch.filter(pl.col("symbol").is_in(batch_syms)) raw = lf_batch.sort(["symbol", "date"]).collect(streaming=True) @@ -1190,7 +1191,7 @@ def _load_recent_history(enriched_base: Path, symbols: list[str], days: int) -> try: lf = ( - pl.scan_parquet(str(enriched_base / "**" / "*.parquet"), cast_options=_cast) + scan_enriched_parquet(str(enriched_base / "**" / "*.parquet"), cast_options=_cast) .filter( (pl.col("symbol").is_in(symbols)) & (pl.col("date") >= cutoff) diff --git a/backend/app/parquet.py b/backend/app/parquet.py new file mode 100644 index 0000000..c268ae5 --- /dev/null +++ b/backend/app/parquet.py @@ -0,0 +1,55 @@ +"""Polars parquet helpers.""" +from __future__ import annotations + +from typing import Any + +import polars as pl + +DAILY_STORAGE_SCHEMA: dict[str, pl.DataType] = { + "symbol": pl.Utf8, + "date": pl.Date, + "open": pl.Float64, + "high": pl.Float64, + "low": pl.Float64, + "close": pl.Float64, + "volume": pl.Float64, + "amount": pl.Float64, + "quote_ts": pl.Int64, +} + +ENRICHED_STORAGE_SCHEMA: dict[str, pl.DataType] = { + "symbol": pl.Utf8, + "date": pl.Date, + "open": pl.Float64, + "high": pl.Float64, + "low": pl.Float64, + "close": pl.Float64, + "volume": pl.Float64, + "amount": pl.Float64, + "raw_close": pl.Float64, + "raw_high": pl.Float64, + "raw_low": pl.Float64, + "turnover_rate": pl.Float64, + "consecutive_limit_ups": pl.UInt32, + "consecutive_limit_downs": pl.UInt32, + "quote_ts": pl.Int64, +} + + +def scan_parquet_compat(source: Any, **kwargs: Any) -> pl.LazyFrame: + """Scan partitioned parquet while tolerating additive schema changes.""" + kwargs.setdefault("missing_columns", "insert") + kwargs.setdefault("extra_columns", "ignore") + return pl.scan_parquet(source, **kwargs) + + +def scan_daily_parquet(source: Any, **kwargs: Any) -> pl.LazyFrame: + kwargs.setdefault("schema", DAILY_STORAGE_SCHEMA) + kwargs.setdefault("cast_options", pl.ScanCastOptions(integer_cast="allow-float")) + return scan_parquet_compat(source, **kwargs) + + +def scan_enriched_parquet(source: Any, **kwargs: Any) -> pl.LazyFrame: + kwargs.setdefault("schema", ENRICHED_STORAGE_SCHEMA) + kwargs.setdefault("cast_options", pl.ScanCastOptions(integer_cast="allow-float")) + return scan_parquet_compat(source, **kwargs) diff --git a/backend/app/services/backtest.py b/backend/app/services/backtest.py index 3054d7d..d3c9640 100644 --- a/backend/app/services/backtest.py +++ b/backend/app/services/backtest.py @@ -15,6 +15,7 @@ import pandas as pd import polars as pl from app.config import settings +from app.parquet import scan_enriched_parquet from app.tickflow.repository import KlineRepository logger = logging.getLogger(__name__) @@ -137,7 +138,7 @@ class BacktestService: from app.tickflow.repository import enriched_dirname enriched_glob = str(self.repo.store.data_dir / enriched_dirname(asset_type) / "**" / "*.parquet") df = ( - pl.scan_parquet(enriched_glob) + scan_enriched_parquet(enriched_glob) .filter( (pl.col("symbol").is_in(symbols)) & (pl.col("date") >= start) diff --git a/backend/app/services/quote_service.py b/backend/app/services/quote_service.py index 2457a5c..688f153 100644 --- a/backend/app/services/quote_service.py +++ b/backend/app/services/quote_service.py @@ -33,6 +33,7 @@ from datetime import date, time as dt_time import polars as pl from app.market_time import cn_now, cn_today +from app.parquet import scan_daily_parquet logger = logging.getLogger(__name__) @@ -1231,9 +1232,9 @@ class QuoteService: cutoff = today - timedelta(days=90) table = "kline_etf_daily" if asset_type == "etf" else "kline_daily" daily_glob = str(self._repo.store.data_dir / table / "**" / "*.parquet") - ohlcv_cols = ["symbol", "date", "open", "high", "low", "close", "volume", "amount"] + ohlcv_cols = ["symbol", "date", "open", "high", "low", "close", "volume", "amount", "quote_ts"] hist_df = ( - pl.scan_parquet(daily_glob) + scan_daily_parquet(daily_glob) .filter(pl.col("date") >= cutoff) .sort(["symbol", "date"]) .collect() diff --git a/backend/app/services/screener.py b/backend/app/services/screener.py index cbb82a9..a38aef9 100644 --- a/backend/app/services/screener.py +++ b/backend/app/services/screener.py @@ -14,6 +14,7 @@ from datetime import date, timedelta import polars as pl +from app.parquet import scan_enriched_parquet from app.tickflow.repository import KlineRepository logger = logging.getLogger(__name__) @@ -320,7 +321,7 @@ class ScreenerService: try: lf = ( - pl.scan_parquet(str(enriched_dir / "**" / "*.parquet")) + scan_enriched_parquet(str(enriched_dir / "**" / "*.parquet")) .filter( (pl.col("date") >= start) & (pl.col("date") <= target_date) @@ -403,7 +404,7 @@ class ScreenerService: try: lf = ( - pl.scan_parquet(str(enriched_dir / "**" / "*.parquet")) + scan_enriched_parquet(str(enriched_dir / "**" / "*.parquet")) .filter((pl.col("date") >= start) & (pl.col("date") <= target_date)) .sort(["symbol", "date"]) ) diff --git a/backend/app/tickflow/repository.py b/backend/app/tickflow/repository.py index f022ff6..ebbc3bc 100644 --- a/backend/app/tickflow/repository.py +++ b/backend/app/tickflow/repository.py @@ -24,6 +24,7 @@ import duckdb import polars as pl from app.config import settings +from app.parquet import scan_enriched_parquet logger = logging.getLogger(__name__) @@ -504,7 +505,7 @@ class KlineRepository: "volume", "amount", "raw_close", "raw_high", "raw_low"] if c in df_latest.columns] lf = ( - pl.scan_parquet(self._enriched_glob) + scan_enriched_parquet(self._enriched_glob) .filter(pl.col("date") >= start_full) .sort(["symbol", "date"]) ) @@ -707,7 +708,7 @@ class KlineRepository: # 昨日连板数: 从 enriched parquet 取 (用于增量计算同向 +1) step = time.perf_counter() logger.info("live agg step start: consecutive state") - lf = pl.scan_parquet(self._enriched_glob).filter(pl.col("date") == latest) + lf = scan_enriched_parquet(self._enriched_glob).filter(pl.col("date") == latest) consec_cols = [c for c in ["symbol", "consecutive_limit_ups", "consecutive_limit_downs"] if c in lf.collect_schema().names()] if len(consec_cols) == 3: @@ -784,7 +785,7 @@ class KlineRepository: from app.indicators.pipeline import compute_indicators lf = ( - pl.scan_parquet(self._enriched_glob) + scan_enriched_parquet(self._enriched_glob) .filter(pl.col("date") >= start_60d) .filter(pl.col("date") <= latest) .sort(["symbol", "date"]) @@ -839,7 +840,7 @@ class KlineRepository: "volume", "amount", "raw_close", "raw_high", "raw_low"] if c in df_latest.columns] df_hist = ( - pl.scan_parquet(self._etf_enriched_glob, + scan_enriched_parquet(self._etf_enriched_glob, cast_options=pl.ScanCastOptions(integer_cast="allow-float")) .filter(pl.col("date") >= start_full) .select(read_cols) @@ -1324,7 +1325,7 @@ class KlineRepository: def _scan_daily_symbol(self, symbol: str, start: date, end: date, columns: list[str] | None) -> pl.DataFrame: try: - lf = pl.scan_parquet(self._enriched_glob, + lf = scan_enriched_parquet(self._enriched_glob, cast_options=pl.ScanCastOptions(integer_cast="allow-float")).filter( (pl.col("symbol") == symbol) & (pl.col("date") >= start) @@ -1341,7 +1342,7 @@ class KlineRepository: def _scan_daily_batch(self, symbols: list[str], start: date, end: date, columns: list[str] | None) -> pl.DataFrame: try: - lf = pl.scan_parquet(self._enriched_glob, + lf = scan_enriched_parquet(self._enriched_glob, cast_options=pl.ScanCastOptions(integer_cast="allow-float")).filter( (pl.col("symbol").is_in(symbols)) & (pl.col("date") >= start) @@ -1358,7 +1359,7 @@ class KlineRepository: def _scan_index_daily_symbol(self, symbol: str, start: date, end: date, columns: list[str] | None) -> pl.DataFrame: try: - lf = pl.scan_parquet(self._index_enriched_glob, + lf = scan_enriched_parquet(self._index_enriched_glob, cast_options=pl.ScanCastOptions(integer_cast="allow-float")).filter( (pl.col("symbol") == symbol) & (pl.col("date") >= start) @@ -1375,7 +1376,7 @@ class KlineRepository: def _scan_etf_daily_symbol(self, symbol: str, start: date, end: date, columns: list[str] | None) -> pl.DataFrame: try: - lf = pl.scan_parquet(self._etf_enriched_glob, + lf = scan_enriched_parquet(self._etf_enriched_glob, cast_options=pl.ScanCastOptions(integer_cast="allow-float")).filter( (pl.col("symbol") == symbol) & (pl.col("date") >= start) diff --git a/backend/tests/test_parquet_schema_compat.py b/backend/tests/test_parquet_schema_compat.py new file mode 100644 index 0000000..0fb611b --- /dev/null +++ b/backend/tests/test_parquet_schema_compat.py @@ -0,0 +1,81 @@ +from datetime import date + +import polars as pl + +from app.parquet import scan_daily_parquet, scan_enriched_parquet + + +def test_partitioned_daily_scan_tolerates_added_quote_ts(tmp_path): + old_part = tmp_path / "kline_daily" / "date=2026-07-08" / "part.parquet" + new_part = tmp_path / "kline_daily" / "date=2026-07-09" / "part.parquet" + old_part.parent.mkdir(parents=True) + new_part.parent.mkdir(parents=True) + + pl.DataFrame({ + "symbol": ["600000.SH"], + "date": [date(2026, 7, 8)], + "open": [10.0], + "high": [10.5], + "low": [9.8], + "close": [10.2], + "volume": [1000.0], + "amount": [10200.0], + }).write_parquet(old_part) + + pl.DataFrame({ + "symbol": ["600000.SH"], + "date": [date(2026, 7, 9)], + "open": [10.2], + "high": [10.8], + "low": [10.1], + "close": [10.6], + "volume": [1200], + "amount": [12720.0], + "quote_ts": [1783560600000], + }).write_parquet(new_part) + + df = scan_daily_parquet(str(tmp_path / "kline_daily" / "**" / "*.parquet")).sort("date").collect() + + assert df.height == 2 + assert df.schema["volume"] == pl.Float64 + assert df.schema["quote_ts"] == pl.Int64 + assert df["quote_ts"].to_list() == [None, 1783560600000] + + +def test_partitioned_enriched_scan_tolerates_added_quote_ts(tmp_path): + base = tmp_path / "kline_daily_enriched" + old_part = base / "date=2026-07-08" / "part.parquet" + new_part = base / "date=2026-07-09" / "part.parquet" + old_part.parent.mkdir(parents=True) + new_part.parent.mkdir(parents=True) + + common_old = { + "symbol": ["600000.SH"], + "date": [date(2026, 7, 8)], + "open": [10.0], + "high": [10.5], + "low": [9.8], + "close": [10.2], + "volume": [1000.0], + "amount": [10200.0], + "raw_close": [10.2], + "raw_high": [10.5], + "raw_low": [9.8], + "turnover_rate": [1.1], + "consecutive_limit_ups": pl.Series([0], dtype=pl.UInt32), + "consecutive_limit_downs": pl.Series([0], dtype=pl.UInt32), + } + pl.DataFrame(common_old).write_parquet(old_part) + + common_new = dict(common_old) + common_new["date"] = [date(2026, 7, 9)] + common_new["volume"] = [1200] + common_new["quote_ts"] = [1783560600000] + pl.DataFrame(common_new).write_parquet(new_part) + + df = scan_enriched_parquet(str(base / "**" / "*.parquet")).sort("date").collect() + + assert df.height == 2 + assert df.schema["volume"] == pl.Float64 + assert df.schema["quote_ts"] == pl.Int64 + assert df["quote_ts"].to_list() == [None, 1783560600000]