Files
sanguo_vnpy_v2/sanguo_portfolio/providers/fetchers/base.py
T

307 lines
10 KiB
Python

"""TET Fetcher 基础设施(窄试点B;设计: docs/design/architecture/provider-tet-design.md)。
三段式契约(源自 OpenBB Fetcher,本项目化):
- ``transform_query``: pydantic 入参严格校验(fail-fast)。老接口对非法参数静默
返空表/取默认,问题下沉到回测跑完才发现;新接口非法即 ``ValidationError``。
- ``extract_data``: 唯一 IO 入口(读本地 dbbardata/parquet;复用 provider 惰性
连接)。SQL 原样搬迁自 LocalUnifiedProvider 对应方法,行为零变化。
- ``transform_data``: DataFrame 级 schema 校验 fail-fast(核心列缺失/全空即报错)。
与 OpenBB 的有意差异:
1. **DataFrame 级校验,非逐行 pydantic Data 模型** —— get_closes_panel 5128只×
全历史逐行构造模型性能不可接受;TET 精髓是「IO 集中 + 严格校验」,不是逐行模型。
2. **extract 读本地不打网络**(本项目落库导向)。
3. ``high_limit``/``low_limit``/``paused`` 的缺失补默认是**显式契约** (2026-08-15
用户定案),不是「兜底」——补 NaN 会让 bullet_trade 把 NaN 当停牌取消全部订单
(见 unified-provider-paused-nan-bug)。schema 校验不把这些可选列当脏数据。
已知妥协(Phase 3 切换时归一): qfq 复权因子的第二次读库 (bs_adjust_factor)
发生在 transform_data 内(经 ctx helper)——因子计算依赖 tail 后的 dates,拆到
extract 需两次往返;试点期注释标记,不为此过度设计。
"""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional, Sequence, Union
from pydantic import BaseModel, ConfigDict, field_validator, model_validator
# fq 合法值(与老接口 get_price/get_closes_panel 取值一致)
FQ_VALUES = ("raw", "qfq", "pre", "前复权")
# get_price 日线 frequency 合法值(老接口其余 frequency 返空 DataFrame=静默;
# _ex 接口 fail-fast 报错——15m 等周期请用 get_closes_panel_ex)
PRICE_FREQ_VALUES = ("daily", "day", "1d", "d")
class DataSchemaError(ValueError):
"""transform_data 检出脏数据(fail-fast)。核心列缺失/全空/类型不符。"""
def _norm_date(v: Optional[Union[str, datetime]]) -> Optional[str]:
"""str/datetime → ``YYYY-MM-DD``;非法类型/格式 raise ValueError。"""
if v is None:
return None
if isinstance(v, datetime):
return v.strftime("%Y-%m-%d")
if isinstance(v, str):
s = v.strip()[:10]
if len(s) == 10 and s[4] == "-" and s[7] == "-":
try:
datetime.strptime(s, "%Y-%m-%d")
return s
except ValueError:
pass
raise ValueError(f"非法日期(期望 YYYY-MM-DD 或 datetime): {v!r}")
class _QueryBase(BaseModel):
"""TET Query 基类: 未知参数直接报错(extra=forbid,防拼写错误静默吞参)。"""
model_config = ConfigDict(extra="forbid", validate_assignment=True)
class PriceQueryParams(_QueryBase):
"""get_price_ex 入参。"""
security: Union[str, Sequence[str]]
start_date: Optional[str] = None
end_date: Optional[str] = None
frequency: str = "daily"
fields: Optional[Sequence[str]] = None
skip_paused: bool = False
fq: str = "raw"
count: Optional[int] = None
panel: bool = True
fill_paused: bool = True
@field_validator("security")
@classmethod
def _v_security(cls, v):
secs = [v] if isinstance(v, str) else list(v)
if not secs:
raise ValueError("security 不能为空(老接口静默返空表,_ex fail-fast)")
for s in secs:
if not isinstance(s, str) or not s.strip():
raise ValueError(f"security 含非法代码: {s!r}")
return secs
@field_validator("start_date", "end_date")
@classmethod
def _v_dates(cls, v):
return _norm_date(v)
@field_validator("frequency")
@classmethod
def _v_frequency(cls, v):
f = str(v or "").lower()
if f not in PRICE_FREQ_VALUES:
raise ValueError(
f"frequency={v!r} 不支持(_ex 只做日线;15m 用 get_closes_panel_ex)"
)
return f
@field_validator("fq")
@classmethod
def _v_fq(cls, v):
if v not in FQ_VALUES:
raise ValueError(f"fq={v!r} 非法,合法值: {FQ_VALUES}")
return v
@field_validator("count")
@classmethod
def _v_count(cls, v):
if v is not None and v < 1:
raise ValueError(f"count={v!r} 必须 >=1")
return v
class PanelQueryParams(_QueryBase):
"""get_closes_panel_ex 入参。"""
symbols: Sequence[str]
start: Union[str, datetime]
end: Union[str, datetime]
interval: str = "d"
fq: str = "raw"
@field_validator("symbols")
@classmethod
def _v_symbols(cls, v):
syms = list(v)
if not syms:
raise ValueError("symbols 不能为空")
for s in syms:
if not isinstance(s, str) or not s.strip():
raise ValueError(f"symbols 含非法代码: {s!r}")
return syms
@field_validator("start", "end")
@classmethod
def _v_dates(cls, v):
return _norm_date(v)
@field_validator("interval")
@classmethod
def _v_interval(cls, v):
s = str(v or "")
if not s or len(s) > 16 or not all(c.isalnum() or c == "_" for c in s):
raise ValueError(f"interval={v!r} 非法(字母数字下划线,<=16)")
return s
@field_validator("fq")
@classmethod
def _v_fq(cls, v):
if v not in FQ_VALUES:
raise ValueError(f"fq={v!r} 非法,合法值: {FQ_VALUES}")
return v
class ConstituentQueryParams(_QueryBase):
"""get_constituent_ex 入参(date 参数保留但忽略——constituent_unified 并集模型无时点)。"""
index: str
date: Optional[str] = None
@field_validator("index")
@classmethod
def _v_index(cls, v):
if not isinstance(v, str) or not v.strip():
raise ValueError("index 不能为空")
return v.strip()
@field_validator("date")
@classmethod
def _v_date(cls, v):
return _norm_date(v)
class FundamentalsQueryParams(_QueryBase):
"""get_fundamentals_df_ex 入参。"""
stocks: Sequence[str]
date: Optional[str] = None
fields: Optional[Sequence[str]] = None
@field_validator("stocks")
@classmethod
def _v_stocks(cls, v):
stocks = list(v)
if not stocks:
raise ValueError("stocks 不能为空")
for s in stocks:
if not isinstance(s, str) or not s.strip():
raise ValueError(f"stocks 含非法代码: {s!r}")
return stocks
@field_validator("date")
@classmethod
def _v_date(cls, v):
return _norm_date(v)
@field_validator("fields")
@classmethod
def _v_fields(cls, v):
if v is None:
return None
fs = list(v)
for f in fs:
if not isinstance(f, str) or not f.strip():
raise ValueError(f"fields 含非法字段: {f!r}")
return fs
# 事件/快照 panel 白名单(2026-09-02 A 档使用层出口):与采集注册表
# (akshare_static_download.py)及 static_vintage_check.PANEL_TYPES 口径一致;
# hot_rank 2026-09-04 拍板挂载(墙测量5轮收官:emappdata 夜间开窗,19:30 档
# 失败零连带——单 unit 失败不触断路器,次夜自动重试)。
EVENT_PANEL_TYPES = frozenset({
"dragon_tiger", "block_trade", "margin_sse", "restricted",
"zt_pool", "zt_pool_zbgc", "zt_pool_dtgc",
"fund_flow_industry", "fund_flow_concept",
"ths_industry", "ths_concept",
"xueqiu_hot", "sina_sector", "gdhs",
"hot_rank",
})
class _DateRangeParams(_QueryBase):
"""date / start+end 二选一语义(事件 panel 族共用)。
- 单日: ``date``;区间: ``start``+``end``(含端点,二缺一/互斥/倒序 → 报错)
"""
date: Optional[str] = None
start: Optional[str] = None
end: Optional[str] = None
@field_validator("date", "start", "end")
@classmethod
def _v_dates(cls, v):
return _norm_date(v)
@model_validator(mode="after")
def _check_date_semantics(self):
if self.date and (self.start or self.end):
raise ValueError("date 与 start/end 互斥(单日用 date,区间用 start+end)")
if not self.date and not self.start:
raise ValueError("date 与 start 必须给其一")
if (self.start and not self.end) or (self.end and not self.start):
raise ValueError("区间查询必须同时给 start 和 end")
if self.start and self.end and self.start > self.end:
raise ValueError("start 不能晚于 end")
return self
class EventPanelQueryParams(_DateRangeParams):
"""get_event_panel 入参(泛型保真通道)。"""
event_type: str
trading_days_only: bool = True
@field_validator("event_type")
@classmethod
def _v_event_type(cls, v):
if v not in EVENT_PANEL_TYPES:
raise ValueError(
f"event_type={v!r} 不在白名单(合法: {sorted(EVENT_PANEL_TYPES)})")
return v
class LimitPoolQueryParams(_DateRangeParams):
"""get_limit_pool 入参(涨停池门面,kind 三合一 ← tushare limit_list_d U/D/Z)。"""
kind: str = "zt"
trading_days_only: bool = True
@field_validator("kind")
@classmethod
def _v_kind(cls, v):
if v not in ("zt", "zbgc", "dtgc"):
raise ValueError("kind 必须是 'zt'|'zbgc'|'dtgc'(涨停|炸板|跌停)")
return v
def validate_df_schema(
df: Any,
*,
required: Sequence[str],
non_empty: Sequence[str] = (),
context: str = "",
) -> None:
"""DataFrame 级 schema 校验(fail-fast)。
- 空 DataFrame 是**合法**情况(标的在库中无数据=缺失,策略已有处理;脏数据
指「有行但核心列空/缺列」)。
- ``required``: 列必须存在(dbbardata schema 变化/SQL 拼错在此拦截)。
- ``non_empty``: 核心数值列不允许全空(全空=源数据损坏,老接口会静默流出
NaN 列,下游 bool(NaN) 坑)。
"""
if df is None or len(df) == 0:
return
missing = [c for c in required if c not in df.columns]
if missing:
raise DataSchemaError(f"{context}: 缺失列 {missing}(df.columns={list(df.columns)})")
for c in non_empty:
if df[c].isna().all():
raise DataSchemaError(f"{context}: 核心列 {c!r} 全空(源数据损坏?)")