fix(storage): 兼容 parquet 分区 schema 演进

兼容 quote_ts 等新增存储列导致的新旧 parquet 分区 schema 差异,恢复老用户 enriched 缓存刷新和看板/选股页面显示。
This commit is contained in:
wshy
2026-07-09 17:38:44 +08:00
committed by GitHub
parent 61d1f28fd8
commit 6bb6eadb12
8 changed files with 163 additions and 21 deletions
+2 -1
View File
@@ -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:
+8 -7
View File
@@ -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)
+55
View File
@@ -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)
+2 -1
View File
@@ -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)
+3 -2
View File
@@ -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()
+3 -2
View File
@@ -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"])
)
+9 -8
View File
@@ -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)
@@ -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]