refactor(provider): TET Phase3老接口内部委托Fetcher——①四老方法(get_price/get_closes_panel/get_index_stocks/get_fundamentals_df)旧体删除,改调同款Fetcher(与_ex两张皮,strict契约自此对老接口生效:非法frequency/fq/count/日期/空列表→ValueError,查询失败raise;Phase2副本对照4/4语义等值已验issue#19)②qfq因子二次读库归一:_build_qfq_factor/_apply_qfq_batch拆为_read_qfq_rows+_qfq_factor_from_rows与_read_qfq_events+_apply_qfq_events(IO/纯计算),因子读全挪进extract_data,transform零IO③all_weather四处get_fundamentals_df(choice)补空守卫(空候选集=调用方业务态)④3个宽松断言回归测试改strict(1m频率/空stocks/空symbols);938绿 [vps]
CI/CD / test (push) Successful in 14s
CI/CD / nas-deploy (push) Successful in 33s
CI/CD / nas-verify (push) Successful in 17s

This commit is contained in:
2026-08-17 09:12:47 +08:00
parent 5211d5a6ed
commit d2364a36c5
6 changed files with 163 additions and 301 deletions
@@ -3,6 +3,8 @@
> 设计日期 2026-07-30。来源:OpenBB Fetcher TET 三段式 + 本项目数据层历史踩坑。配套调研见 `docs/research/openbb-platform-research.md`。
>
> 一句话定位:**把数据层的容错「兜底」范式,改成 TET 三段式「严格校验、不行就报错」范式,让数据质量问题在取数时暴露,而不是被掩盖后在策略下单时记成大祸。**
>
> **落地状态(窄试点 B,全三期完成)**:Phase 1 ✅ 2026-08-15(f9322a7,4 个 `_ex` 新接口 Fetcher 化,老接口零改动);Phase 2 ✅ 2026-08-15(策略 session a45ab30,4 副本对照 4/4 语义等值,issue #19);Phase 3 ✅ 2026-08-17(老接口内部委托 Fetcher,qfq 因子二次读库归一到 extract,strict 契约自此对老接口生效)。
---
+52 -31
View File
@@ -42,9 +42,15 @@ class PriceFetcher:
return PriceQueryParams(**kwargs)
@staticmethod
def extract_data(query: PriceQueryParams, ctx: "LocalUnifiedProvider") -> Dict[str, pd.DataFrame]:
"""唯一 IO(①dbbardata): 逐只查 raw 日线列(原样搬 get_price :206-215)。"""
from ..local_unified_provider import jq_to_dbbardata
def extract_data(
query: PriceQueryParams, ctx: "LocalUnifiedProvider"
) -> tuple[Dict[str, pd.DataFrame], Dict[str, list]]:
"""唯一 IO(①dbbardata ②bs_adjust_factor[qfq 时]): 逐只 raw 日线 + 前复权事件行。
Phase 3 归一:qfq 因子读库从 transform 挪到本段(transform 纯 pandas,
因子 asof 计算按 tail(count) 后的 index 在 transform 里做,读库与 index 无关)。
"""
from ..local_unified_provider import _jq_to_bs_code, _read_qfq_rows, jq_to_dbbardata
conn = ctx._connect()
start_str = query.start_date or "1990-01-01"
@@ -55,26 +61,29 @@ class PriceFetcher:
"AND interval='d' AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? "
"ORDER BY datetime"
)
need_qfq = query.fq in ("qfq", "pre", "前复权")
frames: Dict[str, pd.DataFrame] = {}
qfq_rows: Dict[str, list] = {}
for jq_code in query.security:
sym, exc = jq_to_dbbardata(jq_code)
frames[jq_code] = pd.read_sql(
q, conn, params=(sym, exc, start_str, end_str)
)
return frames
if need_qfq:
qfq_rows[jq_code] = _read_qfq_rows(_jq_to_bs_code(jq_code), conn)
return frames, qfq_rows
@staticmethod
def transform_data(
query: PriceQueryParams,
ctx: "LocalUnifiedProvider",
raw: Dict[str, pd.DataFrame],
raw: tuple[Dict[str, pd.DataFrame], Dict[str, list]],
) -> pd.DataFrame:
"""清洗+schema 校验(原样搬 get_price :216-269)。qfq 因子二次读库=已知妥协(见 base 模块注释)"""
from ..local_unified_provider import _build_qfq_factor, _jq_to_bs_code
"""清洗+schema 校验(纯 pandas,零 IO)。qfq 因子=extract 读好的事件行 asof"""
from ..local_unified_provider import _qfq_factor_from_rows
conn = ctx._connect()
frames_in, qfq_rows = raw
frames: Dict[str, pd.DataFrame] = {}
for jq_code, df in raw.items():
for jq_code, df in frames_in.items():
# 缺数据标的: schema 校验跳过(空 df 合法),但 raw 列结构必须对
if not df.empty:
validate_df_schema(
@@ -93,8 +102,8 @@ class PriceFetcher:
if query.count:
df = df.tail(query.count)
if query.fq in ("qfq", "pre", "前复权") and not df.empty:
factor = _build_qfq_factor(
_jq_to_bs_code(jq_code), conn, pd.Series(df.index)
factor = _qfq_factor_from_rows(
qfq_rows.get(jq_code, []), pd.Series(df.index)
)
for col in ("open_price", "high_price", "low_price", "close_price"):
df[col] = df[col].values * factor.values
@@ -148,7 +157,7 @@ class PriceFetcher:
def fetch(cls, ctx: "LocalUnifiedProvider", **kwargs: Any) -> pd.DataFrame:
query = cls.transform_query(**kwargs)
raw = cls.extract_data(query, ctx)
return cls.transform_data(query, ctx, raw)
return cls.transform_data(query, raw)
class PanelFetcher:
@@ -161,13 +170,16 @@ class PanelFetcher:
@staticmethod
def extract_data(
query: PanelQueryParams, ctx: "LocalUnifiedProvider"
) -> pd.DataFrame:
"""唯一 IO(①dbbardata): chunk=400 UNION ALL per symbol(原样搬 :331-346)。
) -> tuple[pd.DataFrame, pd.DataFrame]:
"""唯一 IO(①dbbardata ②bs_adjust_factor[qfq 时])。
为何 UNION ALL 而非 OR chain: 大 OR 链打不动复合索引(实证全表扫 10s vs
UNION ALL 0.03s,340×);参数 2/symbol×400=800<999 上限。
①chunk=400 UNION ALL per symbol。为何 UNION ALL 而非 OR chain: 大 OR 链
打不动复合索引(实证全表扫 10s vs UNION ALL 0.03s,340×);参数
2/symbol×400=800<999 上限。
Phase 3 归一:qfq 批量因子读库从 transform 挪到本段(events 全量读,
apply 阶段按宽表实际列过滤,语义不变)。
"""
from ..local_unified_provider import jq_to_dbbardata
from ..local_unified_provider import LocalUnifiedProvider, jq_to_dbbardata
pairs: List[tuple[str, str, str]] = [] # (input_code, db_symbol, db_exchange)
for s in query.symbols:
@@ -194,16 +206,25 @@ class PanelFetcher:
for _inp, sym, exc in chunk:
params.extend([sym, exc])
chunk_frames.append(pd.read_sql(q, conn, params=params))
return pd.concat(chunk_frames, ignore_index=True) if chunk_frames else pd.DataFrame()
raw = (
pd.concat(chunk_frames, ignore_index=True)
if chunk_frames else pd.DataFrame()
)
events = (
LocalUnifiedProvider._read_qfq_events(pairs, conn)
if query.fq in ("qfq", "pre", "前复权") else pd.DataFrame()
)
return raw, events
@staticmethod
def transform_data(
query: PanelQueryParams,
ctx: "LocalUnifiedProvider",
raw: pd.DataFrame,
raw: tuple[pd.DataFrame, pd.DataFrame],
) -> pd.DataFrame:
"""清洗+schema 校验(原样搬 :348-390)。qfq 批量因子读库经 ctx helper(已知妥协)"""
from ..local_unified_provider import jq_to_dbbardata
"""清洗+schema 校验(纯 pandas,零 IO)。qfq=extract 读好的 events merge_asof"""
from ..local_unified_provider import LocalUnifiedProvider, jq_to_dbbardata
raw_df, events = raw
# pairs 重建(transform_data 纯函数需要的映射,不读库)
pairs: List[tuple[str, str, str]] = []
@@ -211,23 +232,23 @@ class PanelFetcher:
sym, exc = jq_to_dbbardata(str(s))
pairs.append((str(s), sym, exc))
if raw.empty:
if raw_df.empty:
return pd.DataFrame(
{inp: pd.Series(dtype=float) for inp in query.symbols},
index=pd.DatetimeIndex([]),
)
validate_df_schema(
raw, required=["datetime", "symbol", "close_price"],
raw_df, required=["datetime", "symbol", "close_price"],
non_empty=["close_price"], context="get_closes_panel_ex raw",
)
raw["datetime"] = pd.to_datetime(raw["datetime"], format="mixed")
raw_df["datetime"] = pd.to_datetime(raw_df["datetime"], format="mixed")
sym_to_input: Dict[str, str] = {}
for inp, sym, _exc in pairs:
sym_to_input.setdefault(sym, inp)
raw["symbol"] = raw["symbol"].map(sym_to_input).fillna(raw["symbol"])
raw = raw.drop_duplicates(subset=["datetime", "symbol"], keep="last")
wide = raw.pivot(index="datetime", columns="symbol", values="close_price")
raw_df["symbol"] = raw_df["symbol"].map(sym_to_input).fillna(raw_df["symbol"])
raw_df = raw_df.drop_duplicates(subset=["datetime", "symbol"], keep="last")
wide = raw_df.pivot(index="datetime", columns="symbol", values="close_price")
missing_cols = [inp for inp in query.symbols if inp not in wide.columns]
if missing_cols:
@@ -238,7 +259,7 @@ class PanelFetcher:
wide = wide[query.symbols]
if query.fq in ("qfq", "pre", "前复权"):
wide = ctx._apply_qfq_batch(wide, pairs, ctx._connect())
wide = LocalUnifiedProvider._apply_qfq_events(wide, events)
wide = wide.sort_index()
wide.index.name = None
@@ -248,4 +269,4 @@ class PanelFetcher:
def fetch(cls, ctx: "LocalUnifiedProvider", **kwargs: Any) -> pd.DataFrame:
query = cls.transform_query(**kwargs)
raw = cls.extract_data(query, ctx)
return cls.transform_data(query, ctx, raw)
return cls.transform_data(query, raw)
@@ -93,12 +93,17 @@ def _jq_to_bs_code(jq_code: str) -> str:
return f"{prefix}.{sym}"
def _build_qfq_factor(
bs_code: str,
conn: sqlite3.Connection,
dates: pd.Series,
) -> pd.Series:
"""构造每个 date 的前复权因子(asof 语义)。
def _read_qfq_rows(bs_code: str, conn: sqlite3.Connection) -> List[tuple]:
"""IO:读单只 ``bs_adjust_factor`` 事件行(按 dividOperateDate 升序)。"""
return conn.execute(
"SELECT dividOperateDate, foreAdjustFactor FROM bs_adjust_factor "
"WHERE code=? ORDER BY dividOperateDate",
(bs_code,),
).fetchall()
def _qfq_factor_from_rows(rows: List[tuple], dates: pd.Series) -> pd.Series:
"""纯计算:asof 前复权因子(与 ``_apply_qfq_batch`` 的 merge_asof 语义一致)。
规则:
- 找 ``<= d`` 的最大 dividOperateDate 的 foreAdjustFactor
@@ -108,11 +113,6 @@ def _build_qfq_factor(
``qfq[t] = raw[t] * factor[t]``
"""
rows = conn.execute(
"SELECT dividOperateDate, foreAdjustFactor FROM bs_adjust_factor "
"WHERE code=? ORDER BY dividOperateDate",
(bs_code,),
).fetchall()
dates_ts = pd.to_datetime(dates)
if not rows:
return pd.Series([1.0] * len(dates_ts), index=dates_ts)
@@ -179,94 +179,25 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
count: Optional[int] = None,
panel: bool = True,
fill_paused: bool = True,
**kwargs: Any,
) -> pd.DataFrame:
"""读 ``dbbardata('d')`` raw 日线,按需 ``bs_adjust_factor`` 算前复权。
策略契约(all_weather 实证):
- ``panel=False`` 返长表含 ``time`` + ``code`` 列(供 pivot)
- ``fields`` 里缺失列(如 ``high_limit``)补 NaN(降级)
- ``frequency`` 非 day/1d/d → 返空 DataFrame(1m 数据层无)
- ``frequency`` 只接受 daily/day/1d/d(1m 数据层无)
**Phase 3(2026-08-17)起内部委托 ``PriceFetcher``**(TET 三段式,与
``get_price_ex`` 同一实现)。契约随之切 strict:非法 frequency/fq/count/
日期/空列表 → ValueError fail-fast(老实现静默返空表);策略副本对照 4/4
语义等值已验(issue #19)。空候选集请调用方自行守卫(``if not choice``)。
"""
freq = str(frequency or "").lower()
if freq not in ("daily", "day", "1d", "d"):
return pd.DataFrame()
secs: List[str] = [security] if isinstance(security, str) else list(security or [])
if not secs:
return pd.DataFrame()
conn = self._connect()
start_str = self._to_date_str(start_date) or "1990-01-01"
end_str = self._to_date_str(end_date) or datetime.now().strftime("%Y-%m-%d")
frames: Dict[str, pd.DataFrame] = {}
for jq_code in secs:
sym, exc = jq_to_dbbardata(jq_code)
# substr(datetime,1,10) 取日期部分比 — datetime 列混合格式(有只日期有带时间),
# 纯字符串比 "2024-09-25" < "2024-09-25 00:00:00" 会漏边界行; 比日期(YYYY-MM-DD)规避
q = (
"SELECT datetime, open_price, high_price, low_price, close_price, "
"volume, turnover FROM dbbardata WHERE symbol=? AND exchange=? "
"AND interval='d' AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? "
"ORDER BY datetime"
)
df = pd.read_sql(q, conn, params=(sym, exc, start_str, end_str))
if df.empty:
frames[jq_code] = df
continue
# dbbardata datetime 混合格式(有的 "2024-09-26" 有的 "2024-09-26 00:00:00",
# 不同 schtask/迁移写入);pandas 2.3 严格模式要 format="mixed"
df["datetime"] = pd.to_datetime(df["datetime"], format="mixed")
df = df.set_index("datetime")
df.index.name = None
if count:
df = df.tail(count)
# 前复权
if fq in ("qfq", "pre", "前复权") and not df.empty:
factor = _build_qfq_factor(_jq_to_bs_code(jq_code), conn, pd.Series(df.index))
for col in ("open_price", "high_price", "low_price", "close_price"):
df[col] = df[col].values * factor.values
# jq 风格字段重命名
df = df.rename(columns={
"open_price": "open", "high_price": "high",
"low_price": "low", "close_price": "close",
})
# 缺失字段补默认: paused=False(避免 bool(NaN)=True 被 bullet_trade
# get_current_data 误判停牌→订单 cancel); high_limit/low_limit 按 close±10% 估
# (与 get_current_tick 同口径, 精确涨跌停/ST/创业科创规则 v2); 其他补 NaN
if fields:
for f in fields:
if f in df.columns:
continue
if f == "paused":
df[f] = False
elif f == "high_limit" and "close" in df.columns:
df[f] = (df["close"] * 1.1).round(2)
elif f == "low_limit" and "close" in df.columns:
df[f] = (df["close"] * 0.9).round(2)
else:
df[f] = float("nan")
df = df[[f for f in fields if f in df.columns]]
frames[jq_code] = df
if not frames or all(f.empty for f in frames.values()):
return pd.DataFrame()
if not panel:
parts: List[pd.DataFrame] = []
for jq_code, df in frames.items():
if df.empty:
continue
d = df.reset_index()
# index.name=None 时 reset_index 出 'index' 列; 统一改名 'time'
if "index" in d.columns and "time" not in d.columns:
d = d.rename(columns={"index": "time"})
elif "datetime" in d.columns:
d = d.rename(columns={"datetime": "time"})
d.insert(0, "code", jq_code)
parts.append(d)
return pd.concat(parts, ignore_index=True) if parts else pd.DataFrame()
if len(frames) == 1:
return next(iter(frames.values()))
return pd.concat(frames, axis=1)
from .fetchers.price import PriceFetcher
return PriceFetcher.fetch(
self, security=security, start_date=start_date, end_date=end_date,
frequency=frequency, fields=fields, skip_paused=skip_paused,
fq=fq, count=count, panel=panel, fill_paused=fill_paused,
)
# ==================== get_closes_panel (批量行情,G5 Phase1) ====================
def get_closes_panel(
@@ -299,124 +230,37 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
重复行 dedup (keep last)。
空列表 → 空 DataFrame(index=DatetimeIndex,columns=[])。
**Phase 3(2026-08-17)起内部委托 ``PanelFetcher``**(TET 三段式,与
``get_closes_panel_ex`` 同一实现)。契约随之切 strict:空 symbols/非法日期·
interval·fq → ValueError fail-fast(老实现静默返空骨架);策略副本对照 4/4
语义等值已验(issue #19)。空候选集请调用方自行守卫(``if not stocks``)。
"""
if not symbols:
return pd.DataFrame(index=pd.DatetimeIndex([]))
from .fetchers.price import PanelFetcher
return PanelFetcher.fetch(
self, symbols=symbols, start=start, end=end, interval=interval, fq=fq,
)
# jq 代码 → (db_symbol, db_exchange); 保留 input_code 作列名
# (symbol 不唯一: 932000 中证2000 / 北交所 920xxx 等, 必须按 exchange 配对)
pairs: List[tuple[str, str, str]] = [] # (input_code, db_symbol, db_exchange)
for s in symbols:
sym, exc = jq_to_dbbardata(str(s))
pairs.append((str(s), sym, exc))
conn = self._connect()
start_str = self._to_date_str(start) or "1990-01-01"
end_str = self._to_date_str(end) or datetime.now().strftime("%Y-%m-%d")
# 分批查询 dbbardata: 每 400 symbol 一批, 用 UNION ALL per symbol。
#
# 为何 UNION ALL 而非 OR chain: SQLite 对 ``(sym=? AND exc=?) OR ...`` 大 OR 链
# 无法用 ``dbbardata_symbol_exchange_interval_datetime`` 复合索引(实证 5 sym
# OR 用 idx_dbbardata_interval 全表扫 10s vs UNION ALL 0.03s, 340× 提速)。
# 每个 UNION ALL 子查询独立使用复合索引(SEEK by sym+exc+interval+dt 范围)。
#
# 参数: 2 per subquery (sym, exc) × 400 = 800 < 999 (老版 SQLite 上限);
# interval / start / end 作 *已校验字面量*注入(YYYY-MM-DD regex + alnum)避免每
# subquery 5 参数(会超 999 上限)。日期格式校验在下方 _safe_lit 完成。
start_lit = _safe_date_literal(start_str)
end_lit = _safe_date_literal(end_str)
interval_lit = _safe_interval_literal(interval)
CHUNK_SIZE = 400
chunk_frames: List[pd.DataFrame] = []
for i in range(0, len(pairs), CHUNK_SIZE):
chunk = pairs[i:i + CHUNK_SIZE]
# 每个 subquery 只 2 个参数 (sym, exc); UNION ALL 拼接
sub_template = (
"SELECT datetime, symbol, close_price FROM dbbardata "
f"WHERE symbol=? AND exchange=? AND interval={interval_lit} "
f"AND substr(datetime,1,10)>={start_lit} "
f"AND substr(datetime,1,10)<={end_lit}"
)
q = " UNION ALL ".join([sub_template] * len(chunk))
params: List[Any] = []
for _inp, sym, exc in chunk:
params.extend([sym, exc])
chunk_frames.append(pd.read_sql(q, conn, params=params))
df = pd.concat(chunk_frames, ignore_index=True) if chunk_frames else pd.DataFrame()
if df.empty:
# 全缺失: 返宽表骨架(全 NaN 列, 空 DatetimeIndex) — 与逐只 get_price 行为一致
return pd.DataFrame(
{inp: pd.Series(dtype=float) for inp in symbols},
index=pd.DatetimeIndex([]),
)
# 混合格式 datetime(与 get_price 一致: dbbardata 列混 "2024-09-26" 与
# "2024-09-26 00:00:00",pandas 2.3 严格模式要 format="mixed")
df["datetime"] = pd.to_datetime(df["datetime"], format="mixed")
# db_symbol → input_code 还原(列名回输入 jq code, 与策略其他接口一致)
sym_to_input: Dict[str, str] = {}
for inp, sym, _exc in pairs:
sym_to_input.setdefault(sym, inp)
df["symbol"] = df["symbol"].map(sym_to_input).fillna(df["symbol"])
# dedup: 同 (datetime, symbol) 重复行取最后一条(增量合并容错)
df = df.drop_duplicates(subset=["datetime", "symbol"], keep="last")
# pivot 宽表
wide = df.pivot(index="datetime", columns="symbol", values="close_price")
# 补缺失 symbol 列(全 NaN), 按 input 顺序对齐 columns
# (批量 concat 一次性补, 避免 5128 次 frame.insert 的 PerformanceWarning 碎片化)
missing_cols = [inp for inp in symbols if inp not in wide.columns]
if missing_cols:
wide = pd.concat(
[wide, pd.DataFrame(float("nan"), index=wide.index, columns=missing_cols)],
axis=1,
)
wide = wide[symbols]
# 前复权(批量, 向量化): 与 get_price(fq='qfq') 数值一致, 单次 WHERE code IN + merge_asof
if fq in ("qfq", "pre", "前复权"):
wide = self._apply_qfq_batch(wide, pairs, conn)
# 升序 + 清 index name(与 get_price 一致)
wide = wide.sort_index()
wide.index.name = None
return wide
def _apply_qfq_batch(
self,
wide: pd.DataFrame,
@staticmethod
def _read_qfq_events(
pairs: List[tuple[str, str, str]],
conn: sqlite3.Connection,
columns: Optional[List[str]] = None,
) -> pd.DataFrame:
"""批量前复权(merge_asof 向量化, 替代逐只 ``_build_qfq_factor`` N 次查询)。
"""IO:批量读 ``bs_adjust_factor`` → events DataFrame(供 ``_apply_qfq_events`` 纯计算)。
语义与 ``_build_qfq_factor`` 严格一致:
- 每个 date 用 ``<= d`` 的最大 dividOperateDate 的 foreAdjustFactor(asof backward)
- 早于首事件 → 用最早事件 factor(对应 ``_build_qfq_factor`` 全部 > d → idx=0)
- 无事件 symbol → factor=1.0 → 列不变
实现(高效关键, 避免 N 次 query):
1. 单次(或 400-chunked) ``WHERE code IN (...)`` 查所有 symbol 的全部事件
2. 按 input_code groupby, 每组 ``pd.merge_asof(direction='backward')`` onto
宽表 date index — pandas 内部 searchsorted 向量化, 远快于 Python 循环
3. ``qfq[t] = raw[t] * factor[t]`` 逐列乘
``columns`` 非空时只读宽表实际存在的列(全 NaN 列不浪费 query);None=全读
(Fetcher extract 阶段尚无宽表,读全量由 apply 阶段按存在性过滤,语义不变)。
返回列 ``input_code/event_date/factor``,按 (input_code, event_date) 升序。
"""
# input_code → bs_code 映射(只对宽表实际存在的列查, dropna 后全 NaN 列不浪费 query)
bs_to_input: Dict[str, str] = {}
for inp, sym, exc in pairs:
if inp in wide.columns:
if columns is None or inp in columns:
prefix = "sh" if exc == "SSE" else "sz"
bs_to_input.setdefault(f"{prefix}.{sym}", inp)
if not bs_to_input:
return wide
return pd.DataFrame()
# 批量查 bs_adjust_factor(chunk=400, 参数 < 999)
all_bs_codes = list(bs_to_input.keys())
@@ -431,7 +275,7 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
factor_rows.extend(conn.execute(q, chunk).fetchall())
if not factor_rows:
return wide # 无任何事件 → factor 全 1.0 → 等于 raw
return pd.DataFrame() # 无任何事件 → factor 全 1.0 → 等于 raw
events = pd.DataFrame(
factor_rows, columns=["bs_code", "event_date", "factor"]
@@ -439,10 +283,17 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
events["event_date"] = pd.to_datetime(events["event_date"])
events["input_code"] = events["bs_code"].map(bs_to_input)
events = events.dropna(subset=["input_code"])
if events.empty:
return wide
return events.sort_values(["input_code", "event_date"]).reset_index(drop=True)
events = events.sort_values(["input_code", "event_date"]).reset_index(drop=True)
@staticmethod
def _apply_qfq_events(wide: pd.DataFrame, events: pd.DataFrame) -> pd.DataFrame:
"""纯计算:批量前复权(merge_asof 向量化)。语义与 ``_qfq_factor_from_rows`` 严格一致:
- 每个 date 用 ``<= d`` 的最大 dividOperateDate 的 foreAdjustFactor(asof backward)
- 早于首事件 → 用最早事件 factor(对应逐只版全部 > d → idx=0)
- 无事件 symbol → factor=1.0 → 列不变(events 无该 input_code 即跳过)
"""
if events is None or events.empty:
return wide
# merge_asof 要求 left 升序; 这里先 sort_index, 末尾调用方再 sort_index 是 no-op
wide = wide.sort_index()
@@ -464,7 +315,7 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
)
factor_series = merged["factor"]
if factor_series.isna().any():
# 早于首事件 → 用最早事件 factor(对齐 _build_qfq_factor idx=0)
# 早于首事件 → 用最早事件 factor(对齐逐只版 idx=0)
factor_series = factor_series.fillna(g["factor"].iloc[0])
# 防御: 若 inp_code 在 wide.columns 有重复(caller 传 dup symbols),
# wide[inp_code] 返 DataFrame 而非 Series, 用 iloc[:, 0] 取首个一致列
@@ -485,29 +336,14 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
⚠️ 并集模型: 表无 date 列,治"纯当前"幸存者偏差(含已退市/被踢),
但有轻微前视(使用说明标注)。``date`` 参数忽略(无时点数据)。
**Phase 3(2026-08-17)起内部委托 ``ConstituentFetcher``**(TET 三段式,
与 ``get_constituent_ex`` 同一实现)。契约随之切 strict:查询失败(表缺/
库锁)→ DataSchemaError raise(老实现 log+返 ``[]`` 静默);空 index →
ValueError。未知 index 无记录 → 仍返 ``[]``(合法缺失)。
"""
idx = (
index_symbol.split(".")[0]
if "." in str(index_symbol) else str(index_symbol)
)
conn = self._connect()
try:
rows = conn.execute(
"SELECT code FROM constituent_unified WHERE index_code=? "
"AND (in_current=1 OR was_removed=1)",
(idx,),
).fetchall()
except sqlite3.Error as exc:
logger.warning("constituent_unified 查询失败 %s: %s", idx, exc)
return []
out: List[str] = []
for (code,) in rows:
code_str = str(code).strip()
if len(code_str) != 6 or not code_str.isdigit():
continue
exc_name = "SSE" if code_str.startswith("6") else "SZSE"
out.append(dbbardata_to_jq(code_str, exc_name))
return out
from .fetchers.constituent import ConstituentFetcher
return ConstituentFetcher.fetch(self, index=index_symbol, date=date)
def get_constituent(
self,
@@ -600,29 +436,14 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
- ``len(stocks) > _FUND_POOL_THRESHOLD`` 时 ThreadPool 并发逐只读(本地文件 I/O,
非 baostock 网络 → 并发安全, 不触"baostock 不并发"铁律)。
``fields=None`` 全列(向后兼容, all_weather/small_cap 现有调用零改动)。
**Phase 3(2026-08-17)起内部委托 ``FundamentalsFetcher``**(TET 三段式,
与 ``get_fundamentals_df_ex`` 同一实现)。契约随之切 strict:空 stocks →
ValueError(老实现返空表);策略副本对照 4/4 语义等值已验(issue #19)。
空候选集请调用方自行守卫(``if not choice: return []``)。
"""
from .local_parquet_provider import _FUNDAMENTAL_COLUMNS, jq_to_file_code
if not stocks:
return pd.DataFrame(columns=_FUNDAMENTAL_COLUMNS)
date_str = self._to_date_str(date) or datetime.now().strftime("%Y-%m-%d")
need = self._fields_to_need(fields) if fields else None
def _one(jq_code: str) -> Dict[str, Any]:
return self._build_fundamental_row(jq_code, date_str, need)
if len(stocks) <= _FUND_POOL_THRESHOLD:
rows: List[Dict[str, Any]] = [_one(s) for s in stocks]
else:
workers = min(8, os.cpu_count() or 4)
with ThreadPoolExecutor(max_workers=workers) as ex:
rows = list(ex.map(_one, stocks))
df = pd.DataFrame(rows, columns=_FUNDAMENTAL_COLUMNS)
if "code" in df.columns:
df = df.set_index("code", drop=False)
if fields:
keep = ["code"] + [f for f in fields if f in df.columns]
df = df[keep]
return df
from .fetchers.fundamentals import FundamentalsFetcher
return FundamentalsFetcher.fetch(self, stocks=stocks, date=date, fields=fields)
def get_value_metrics(
self,
@@ -1010,8 +831,10 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
# 设计: docs/design/architecture/provider-tet-design.md
# 契约: 合法参数下输出与老接口逐值一致(等价性测试保证);**非法参数 fail-fast
# 报错**(老接口静默返空/取默认)——这是有意的 strict 新行为。
# Phase 2: 策略 session copy 策略副本改调 _ex 在 NAS 回测对照;
# Phase 3: 验证通过后,老接口内部改为委托 Fetcher(届时本段注释更新)。
# Phase 2: 策略副本 4/4 语义等值(issue #19,a45ab30)。
# Phase 3 ✅(2026-08-17): 老接口(get_price/get_closes_panel/get_index_stocks/
# get_fundamentals_df)内部已改为委托同一 Fetcher——本段 _ex 方法与老方法现在是
# **同一实现的两张皮**,_ex 保留作未来 MCP 出口;strict 契约自此对老接口同样生效。
def get_price_ex(
self,
@@ -1026,7 +849,7 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
panel: bool = True,
fill_paused: bool = True,
) -> pd.DataFrame:
"""TET 版 get_price(日线)。差异 vs 老接口: 非法 frequency/fq/count/日期 → ValueError"""
"""TET 版 get_price(日线)。Phase 3 起与老 ``get_price`` 同一实现(strict 契约)"""
from .fetchers.price import PriceFetcher
return PriceFetcher.fetch(
self, security=security, start_date=start_date, end_date=end_date,
@@ -1042,7 +865,7 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
interval: str = "d",
fq: str = "raw",
) -> pd.DataFrame:
"""TET 版 get_closes_panel(批量 close 宽表)。差异 vs 老接口: 空 symbols/非法日期·interval·fq → ValueError"""
"""TET 版 get_closes_panel(批量 close 宽表)。Phase 3 起与老接口同一实现(strict 契约)"""
from .fetchers.price import PanelFetcher
return PanelFetcher.fetch(
self, symbols=symbols, start=start, end=end, interval=interval, fq=fq,
@@ -1053,7 +876,7 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
index: str,
date: Optional[Union[str, datetime]] = None,
) -> List[str]:
"""TET 版 get_constituent(constituent_unified 并集)。差异 vs 老接口: 查询失败 raise(老接口 log+返[])。"""
"""TET 版 get_constituent(constituent_unified 并集)。Phase 3 起与老接口同一实现(strict 契约)。"""
from .fetchers.constituent import ConstituentFetcher
return ConstituentFetcher.fetch(self, index=index, date=date)
@@ -1063,6 +886,6 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
date: Optional[Union[str, datetime]] = None,
fields: Optional[List[str]] = None,
) -> pd.DataFrame:
"""TET 版 get_fundamentals_df(多股财务/估值)。差异 vs 老接口: 空 stocks/非法字段 → ValueError"""
"""TET 版 get_fundamentals_df(多股财务/估值)。Phase 3 起与老接口同一实现(strict 契约)"""
from .fetchers.fundamentals import FundamentalsFetcher
return FundamentalsFetcher.fetch(self, stocks=stocks, date=date, fields=fields)
@@ -285,6 +285,8 @@ class AllWeatherStrategy:
降到 roe>0.05 & roa>0.02(中证1000 中位)确保选出足够股验证分支C链路
"""
cfg = self.config
if not choice:
return []
df = self.provider.get_fundamentals_df(choice, date=previous_date)
if df.empty:
return []
@@ -295,6 +297,8 @@ class AllWeatherStrategy:
def big(self, choice: List[str], current_dt: Any, previous_date: str) -> List[str]:
"""BIG: 多因子筛选,按 market_cap desc,取前 stock_num。"""
cfg = self.config
if not choice:
return []
df = self.provider.get_fundamentals_df(choice, date=previous_date)
if df.empty:
return []
@@ -316,6 +320,8 @@ class AllWeatherStrategy:
def roic_big(self, choice: List[str], current_dt: Any, previous_date: str) -> List[str]:
"""ROIC_BIG: 多因子筛选 + ROIC 过滤,按 retained_profit desc 取前 stock_num。"""
cfg = self.config
if not choice:
return []
df = self.provider.get_fundamentals_df(choice, date=previous_date)
if df.empty:
return []
@@ -340,6 +346,8 @@ class AllWeatherStrategy:
def bm(self, choice: List[str], current_dt: Any, previous_date: str) -> List[str]:
"""BM: 中市值价值股,按 market_cap asc 取前 stock_num。"""
cfg = self.config
if not choice:
return []
df = self.provider.get_fundamentals_df(choice, date=previous_date)
if df.empty:
return []
+22 -14
View File
@@ -3,7 +3,7 @@
Mac 本地 TDD: sqlite tmp_path + tmp parquet fixture, VPS 依赖,零网络
覆盖:
- ``jq_to_dbbardata`` / ``dbbardata_to_jq`` / ``_jq_to_bs_code`` 代码转换
- ``_build_qfq_factor`` 复权因子构造(asof)
- ``_build_qfq_factor`` 复权因子构造(asof;Phase 3 拆为 _read_qfq_rows IO + _qfq_factor_from_rows 纯计算)
- ``get_price`` dbbardata('d') raw + fq='qfq' 前复权 + panel=False 长表
- ``get_index_stocks`` constituent_unified 并集(治偏差, date 时点)
- ``get_fundamentals_df`` pe/pb/ps/pcfvaluation_baostock + 市值static akshare + 三表委托 LocalParquetProvider
@@ -21,11 +21,17 @@ from sanguo_portfolio.providers.local_unified_provider import (
jq_to_dbbardata,
dbbardata_to_jq,
_jq_to_bs_code,
_build_qfq_factor,
_read_qfq_rows,
_qfq_factor_from_rows,
LocalUnifiedProvider,
)
def _build_qfq_factor(bs_code, conn, dates):
"""Phase 3 拆分后的组合形式(IO 读 + 纯计算),保留原测试调用形状。"""
return _qfq_factor_from_rows(_read_qfq_rows(bs_code, conn), dates)
# ======================== Task 0: 代码转换 ========================
class TestCodeFormat:
def test_jq_to_dbbardata_sh(self):
@@ -242,16 +248,16 @@ class TestGetPrice:
assert abs(df.iloc[0]["high_limit"] - expected) < 1e-6
def test_minute_frequency_returns_empty(self, unified_provider):
# 1m 频率无数据 → 返空 DataFrame
df = unified_provider.get_price(
"600519.XSHG",
end_date="2024-06-20",
frequency="1m",
count=1,
panel=False,
)
assert isinstance(df, pd.DataFrame)
assert df.empty
# Phase 3(2026-08-17)起老接口=Fetcher strict 契约:
# 1m 频率 → ValueError fail-fast(老实现静默返空表;15m 走 get_closes_panel)
with pytest.raises(ValueError, match="frequency"):
unified_provider.get_price(
"600519.XSHG",
end_date="2024-06-20",
frequency="1m",
count=1,
panel=False,
)
def test_multi_stocks_panel_false(self, tmp_path):
# 多股 panel=False → 长表含 code 列区分
@@ -469,10 +475,12 @@ class TestGetFundamentals:
assert col in df.columns, f"missing col: {col}"
def test_empty_stocks_returns_empty(self, tmp_path):
# Phase 3(2026-08-17)起老接口=Fetcher strict 契约:
# 空 stocks → ValueError(老实现返空表);空候选集由调用方守卫(if not choice)
db = _make_fundamentals_fixture(tmp_path)
p = LocalUnifiedProvider({"db_path": str(db), "data_dir": str(tmp_path)})
df = p.get_fundamentals_df([], date="2024-09-30")
assert df.empty
with pytest.raises(ValueError, match="stocks"):
p.get_fundamentals_df([], date="2024-09-30")
# ======================== Task 3b: get_fundamentals_df fields= 按需短路 + 并发 ========================
+4 -4
View File
@@ -183,10 +183,10 @@ class TestDateRangeAndInterval:
# ======================== 4. 空 symbols + dedup ========================
class TestEmptyAndDedup:
def test_empty_symbols_returns_empty_dataframe(self, batch_provider):
df = batch_provider.get_closes_panel([], start="2024-06-18", end="2024-06-20")
assert isinstance(df, pd.DataFrame)
assert df.empty
assert list(df.columns) == []
# Phase 3(2026-08-17)起老接口=Fetcher strict 契约:
# 空 symbols → ValueError(老实现返空骨架);空候选集由调用方守卫(if not stocks)
with pytest.raises(ValueError, match="symbols"):
batch_provider.get_closes_panel([], start="2024-06-18", end="2024-06-20")
def test_duplicate_rows_deduped_keep_last(self, tmp_path):
# 同 (symbol, datetime) 两行 close 不同 → dedup keep last