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]
This commit is contained in:
@@ -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 契约自此对老接口生效)。
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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/pcf←valuation_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= 按需短路 + 并发 ========================
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user