mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 17:54:15 +08:00
- 全部 K 线取数点改走 CompactKlineData 列式直转 polars (无 pandas 中转, 全市场单轮本地转换 ~1.1s → ~0.1s), 时区口径统一 _timestamp_to_beijing_datetime - 除权因子 trade_date 原取 UTC 日期, 北京零点事件整体早一天 → 转北京墙钟后取日期 - 盘中分钟增量间隔上限 300s → 120s (universe 仅回最新 3 根, 超限必留缺口)
111 lines
4.4 KiB
Python
111 lines
4.4 KiB
Python
"""Normalize provider responses into internal Polars schemas."""
|
|
from __future__ import annotations
|
|
|
|
import polars as pl
|
|
|
|
from app.indicators.pipeline import filter_halt_days
|
|
|
|
DAILY_COLS = ["symbol", "date", "open", "high", "low", "close", "volume", "amount", "quote_ts"]
|
|
ADJ_FACTOR_COLS = ["symbol", "trade_date", "ex_factor"]
|
|
INSTRUMENT_COLS = ["symbol", "name", "code", "exchange", "asset_type", "source"]
|
|
|
|
|
|
def to_polars(data) -> pl.DataFrame:
|
|
if data is None:
|
|
return pl.DataFrame()
|
|
if isinstance(data, pl.DataFrame):
|
|
return data
|
|
if isinstance(data, dict):
|
|
rows: list[dict] = []
|
|
for sym, values in data.items():
|
|
for item in values or []:
|
|
row = dict(item or {})
|
|
row.setdefault("symbol", sym)
|
|
rows.append(row)
|
|
return pl.DataFrame(rows) if rows else pl.DataFrame()
|
|
if hasattr(data, "reset_index"):
|
|
return pl.from_pandas(data.reset_index())
|
|
try:
|
|
return pl.DataFrame(data)
|
|
except Exception: # noqa: BLE001
|
|
return pl.DataFrame()
|
|
|
|
|
|
def normalize_daily(data, default_symbol: str | None = None, source: str = "tickflow") -> pl.DataFrame: # noqa: ARG001
|
|
df = to_polars(data)
|
|
if df.is_empty():
|
|
return df
|
|
rename_map = {
|
|
"ts_code": "symbol",
|
|
"trade_date": "date",
|
|
"datetime": "date",
|
|
"vol": "volume",
|
|
"amt": "amount",
|
|
"timestamp": "quote_ts",
|
|
}
|
|
df = df.rename({k: v for k, v in rename_map.items() if k in df.columns})
|
|
if "symbol" not in df.columns and default_symbol:
|
|
df = df.with_columns(pl.lit(default_symbol).alias("symbol"))
|
|
if "date" in df.columns and df.schema["date"] != pl.Date:
|
|
df = df.with_columns(pl.col("date").cast(pl.Date, strict=False))
|
|
# quote_ts: 毫秒级行情时间戳, 用于盘后校验/量比折算。保留为 Int64, 缺失则置 null。
|
|
if "quote_ts" in df.columns:
|
|
df = df.with_columns(pl.col("quote_ts").cast(pl.Int64, strict=False))
|
|
for col in ("open", "high", "low", "close", "volume", "amount"):
|
|
if col in df.columns:
|
|
df = df.with_columns(pl.col(col).cast(pl.Float64, strict=False))
|
|
df = filter_halt_days(df)
|
|
keep = [c for c in DAILY_COLS if c in df.columns]
|
|
return df.select(keep) if keep else pl.DataFrame()
|
|
|
|
|
|
def normalize_adj_factors(data, source: str = "tickflow") -> pl.DataFrame: # noqa: ARG001
|
|
df = to_polars(data)
|
|
if df.is_empty():
|
|
return df
|
|
rename_map = {
|
|
"timestamp": "trade_date",
|
|
"date": "trade_date",
|
|
"adj_factor": "ex_factor",
|
|
}
|
|
df = df.rename({k: v for k, v in rename_map.items() if k in df.columns})
|
|
if "trade_date" in df.columns:
|
|
if df.schema["trade_date"] in {pl.Int64, pl.Int32, pl.UInt64, pl.UInt32, pl.Float64, pl.Float32}:
|
|
# 毫秒时间戳 → 北京墙钟日期 (直接 from_epoch().dt.date() 是 UTC 日期,
|
|
# 除权事件时间戳为北京零点 = UTC 前一日 16:00, 会整体早一天)。
|
|
df = df.with_columns(
|
|
pl.from_epoch(pl.col("trade_date").cast(pl.Int64), time_unit="ms")
|
|
.dt.replace_time_zone("UTC")
|
|
.dt.convert_time_zone("Asia/Shanghai")
|
|
.dt.replace_time_zone(None)
|
|
.dt.date()
|
|
.alias("trade_date")
|
|
)
|
|
else:
|
|
df = df.with_columns(pl.col("trade_date").cast(pl.Date, strict=False))
|
|
if "ex_factor" in df.columns:
|
|
df = df.with_columns(pl.col("ex_factor").cast(pl.Float64, strict=False))
|
|
keep = [c for c in ADJ_FACTOR_COLS if c in df.columns]
|
|
return df.select(keep).drop_nulls() if len(keep) == len(ADJ_FACTOR_COLS) else pl.DataFrame()
|
|
|
|
|
|
def normalize_instruments(rows: list[dict], asset_type: str, source: str = "tickflow") -> pl.DataFrame:
|
|
if not rows:
|
|
return pl.DataFrame()
|
|
out: list[dict] = []
|
|
for item in rows:
|
|
symbol = item.get("symbol")
|
|
if not symbol:
|
|
continue
|
|
out.append({
|
|
"symbol": str(symbol),
|
|
"name": item.get("name") or str(symbol),
|
|
"code": item.get("code") or str(symbol).split(".")[0],
|
|
"exchange": item.get("exchange"),
|
|
"asset_type": asset_type,
|
|
"source": source,
|
|
})
|
|
if not out:
|
|
return pl.DataFrame()
|
|
return pl.DataFrame(out).select(INSTRUMENT_COLS).unique(subset=["symbol"], keep="last").sort("symbol")
|