From d2364a36c5e56f16b1883ec956f2bbe5c2964a1e Mon Sep 17 00:00:00 2001 From: claude_dev Date: Mon, 17 Aug 2026 09:12:47 +0800 Subject: [PATCH] =?UTF-8?q?refactor(provider):=20TET=20Phase3=E8=80=81?= =?UTF-8?q?=E6=8E=A5=E5=8F=A3=E5=86=85=E9=83=A8=E5=A7=94=E6=89=98Fetcher?= =?UTF-8?q?=E2=80=94=E2=80=94=E2=91=A0=E5=9B=9B=E8=80=81=E6=96=B9=E6=B3=95?= =?UTF-8?q?(get=5Fprice/get=5Fcloses=5Fpanel/get=5Findex=5Fstocks/get=5Ffu?= =?UTF-8?q?ndamentals=5Fdf)=E6=97=A7=E4=BD=93=E5=88=A0=E9=99=A4,=E6=94=B9?= =?UTF-8?q?=E8=B0=83=E5=90=8C=E6=AC=BEFetcher(=E4=B8=8E=5Fex=E4=B8=A4?= =?UTF-8?q?=E5=BC=A0=E7=9A=AE,strict=E5=A5=91=E7=BA=A6=E8=87=AA=E6=AD=A4?= =?UTF-8?q?=E5=AF=B9=E8=80=81=E6=8E=A5=E5=8F=A3=E7=94=9F=E6=95=88:?= =?UTF-8?q?=E9=9D=9E=E6=B3=95frequency/fq/count/=E6=97=A5=E6=9C=9F/?= =?UTF-8?q?=E7=A9=BA=E5=88=97=E8=A1=A8=E2=86=92ValueError,=E6=9F=A5?= =?UTF-8?q?=E8=AF=A2=E5=A4=B1=E8=B4=A5raise;Phase2=E5=89=AF=E6=9C=AC?= =?UTF-8?q?=E5=AF=B9=E7=85=A74/4=E8=AF=AD=E4=B9=89=E7=AD=89=E5=80=BC?= =?UTF-8?q?=E5=B7=B2=E9=AA=8Cissue#19)=E2=91=A1qfq=E5=9B=A0=E5=AD=90?= =?UTF-8?q?=E4=BA=8C=E6=AC=A1=E8=AF=BB=E5=BA=93=E5=BD=92=E4=B8=80:=5Fbuild?= =?UTF-8?q?=5Fqfq=5Ffactor/=5Fapply=5Fqfq=5Fbatch=E6=8B=86=E4=B8=BA=5Fread?= =?UTF-8?q?=5Fqfq=5Frows+=5Fqfq=5Ffactor=5Ffrom=5Frows=E4=B8=8E=5Fread=5Fq?= =?UTF-8?q?fq=5Fevents+=5Fapply=5Fqfq=5Fevents(IO/=E7=BA=AF=E8=AE=A1?= =?UTF-8?q?=E7=AE=97),=E5=9B=A0=E5=AD=90=E8=AF=BB=E5=85=A8=E6=8C=AA?= =?UTF-8?q?=E8=BF=9Bextract=5Fdata,transform=E9=9B=B6IO=E2=91=A2all=5Fweat?= =?UTF-8?q?her=E5=9B=9B=E5=A4=84get=5Ffundamentals=5Fdf(choice)=E8=A1=A5?= =?UTF-8?q?=E7=A9=BA=E5=AE=88=E5=8D=AB(=E7=A9=BA=E5=80=99=E9=80=89?= =?UTF-8?q?=E9=9B=86=3D=E8=B0=83=E7=94=A8=E6=96=B9=E4=B8=9A=E5=8A=A1?= =?UTF-8?q?=E6=80=81)=E2=91=A33=E4=B8=AA=E5=AE=BD=E6=9D=BE=E6=96=AD?= =?UTF-8?q?=E8=A8=80=E5=9B=9E=E5=BD=92=E6=B5=8B=E8=AF=95=E6=94=B9strict(1m?= =?UTF-8?q?=E9=A2=91=E7=8E=87/=E7=A9=BAstocks/=E7=A9=BAsymbols);938?= =?UTF-8?q?=E7=BB=BF=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../architecture/provider-tet-design.md | 2 + sanguo_portfolio/providers/fetchers/price.py | 83 +++-- .../providers/local_unified_provider.py | 327 ++++-------------- sanguo_portfolio/strategies/all_weather.py | 8 + .../portfolio/test_local_unified_provider.py | 36 +- tests/portfolio/test_provider_batch.py | 8 +- 6 files changed, 163 insertions(+), 301 deletions(-) diff --git a/docs/design/architecture/provider-tet-design.md b/docs/design/architecture/provider-tet-design.md index fe1f939..b1bd4af 100644 --- a/docs/design/architecture/provider-tet-design.md +++ b/docs/design/architecture/provider-tet-design.md @@ -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 契约自此对老接口生效)。 --- diff --git a/sanguo_portfolio/providers/fetchers/price.py b/sanguo_portfolio/providers/fetchers/price.py index 377a6b9..7c456e2 100644 --- a/sanguo_portfolio/providers/fetchers/price.py +++ b/sanguo_portfolio/providers/fetchers/price.py @@ -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) diff --git a/sanguo_portfolio/providers/local_unified_provider.py b/sanguo_portfolio/providers/local_unified_provider.py index 94e5505..18d1975 100644 --- a/sanguo_portfolio/providers/local_unified_provider.py +++ b/sanguo_portfolio/providers/local_unified_provider.py @@ -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) diff --git a/sanguo_portfolio/strategies/all_weather.py b/sanguo_portfolio/strategies/all_weather.py index 9bbd3ca..7eeae6a 100644 --- a/sanguo_portfolio/strategies/all_weather.py +++ b/sanguo_portfolio/strategies/all_weather.py @@ -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 [] diff --git a/tests/portfolio/test_local_unified_provider.py b/tests/portfolio/test_local_unified_provider.py index c7040eb..9155f4a 100644 --- a/tests/portfolio/test_local_unified_provider.py +++ b/tests/portfolio/test_local_unified_provider.py @@ -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= 按需短路 + 并发 ======================== diff --git a/tests/portfolio/test_provider_batch.py b/tests/portfolio/test_provider_batch.py index 6e1ce83..bc8e5df 100644 --- a/tests/portfolio/test_provider_batch.py +++ b/tests/portfolio/test_provider_batch.py @@ -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