From c8f26bef80c653145631ec92d5d25a1ef6ac8a32 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 28 Jul 2026 20:36:11 +0800 Subject: [PATCH] =?UTF-8?q?feat(portfolio):=20LocalUnifiedProvider=20?= =?UTF-8?q?=E6=89=B9=E9=87=8F=E8=A1=8C=E6=83=85=E6=8E=A5=E5=8F=A3=20get=5F?= =?UTF-8?q?closes=5Fpanel=20(G5-Phase1)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 单条 dbbardata 参数化查询 (symbol,exchange) OR + pivot 返宽表, 替代 N 次 get_price, 为策略向量化提速铺路。raw close 口径一致, 缺失 NaN 列, symbol-exchange 配对防歧义。22 tests 含 batch-vs-逐只回归。本文件另含 get_value_metrics 透传(策略移植)。 --- .../providers/local_unified_provider.py | 102 +++++ tests/portfolio/test_provider_batch.py | 351 ++++++++++++++++++ 2 files changed, 453 insertions(+) create mode 100644 tests/portfolio/test_provider_batch.py diff --git a/sanguo_portfolio/providers/local_unified_provider.py b/sanguo_portfolio/providers/local_unified_provider.py index 90a435b..321f84e 100644 --- a/sanguo_portfolio/providers/local_unified_provider.py +++ b/sanguo_portfolio/providers/local_unified_provider.py @@ -241,6 +241,96 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc] return next(iter(frames.values())) return pd.concat(frames, axis=1) + # ==================== get_closes_panel (批量行情,G5 Phase1) ==================== + def get_closes_panel( + self, + symbols: List[str], + start: Union[str, datetime], + end: Union[str, datetime], + interval: str = "d", + ) -> pd.DataFrame: + """批量取多只股票 close,返宽表 (index=datetime, columns=symbol, values=close)。 + + 单条 dbbardata 参数化查询 ``(symbol, exchange) OR (...)`` + ``pivot``, + 替代选股时 N 次 ``get_price``(为 Phase 2 策略向量化提速 10-50× 铺路)。 + + 与逐只 ``get_price(fq='raw')`` 等价(raw close, 不复权),区别仅在 IO 次数: + - 逐只: N 次 SQL(N+1 query pattern) + - 批量: 1 次 SQL + pivot + + Args: + symbols: jq 风格代码列表 (``["600519.XSHG", "000001.XSHE"]``) 或纯 6 位 + start/end: 日期 ``YYYY-MM-DD`` (或 datetime,取 date 部分) + interval: ``"d"``=日线 (dbbardata ``interval`` 字段; ``"15m"`` 等 v2 扩展) + + Returns: + DataFrame, ``index=datetime`` (升序), ``columns=symbols``(按输入顺序), + ``values=close_price``。缺失股票 → 该列全 NaN; 重复行 dedup (keep last)。 + + 空列表 → 空 DataFrame(index=DatetimeIndex,columns=[])。 + """ + if not symbols: + return pd.DataFrame(index=pd.DatetimeIndex([])) + + # 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") + + # 参数化 (symbol, exchange) OR 子句 — 防注入 + 不字符串拼接(长度无上限) + where_parts = " OR ".join(["(symbol=? AND exchange=?)" for _ in pairs]) + params: List[Any] = [interval, start_str, end_str] + for _inp, sym, exc in pairs: + params.extend([sym, exc]) + + q = ( + "SELECT datetime, symbol, close_price FROM dbbardata " + "WHERE interval=? " + "AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? " + f"AND ({where_parts})" + ) + df = pd.read_sql(q, conn, params=params) + + 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 + for inp in symbols: + if inp not in wide.columns: + wide[inp] = float("nan") + wide = wide[symbols] + + # 升序 + 清 index name(与 get_price 一致) + wide = wide.sort_index() + wide.index.name = None + return wide + # ==================== get_index_stocks (constituent_unified 并集) ==================== def get_index_stocks( self, @@ -334,6 +424,18 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc] df = df.set_index("code", drop=False) return df + def get_value_metrics( + self, + stock: str, + date: Union[str, datetime], + ) -> Optional[Dict[str, Any]]: + """委托 LocalParquetProvider 读三表多期 + NOTICE_DATE 过滤(供 ValueSelectionStrategy)。 + + LocalParquetProvider 实现见 ``local_parquet_provider.get_value_metrics``; + 本类只是把职责转发给已有的 ``_lpp_helper``(DRY, 不复制字段映射逻辑)。 + """ + return self._get_lpp_helper().get_value_metrics(stock, date) + def _build_fundamental_row( self, jq_code: str, date_str: str, ) -> Dict[str, Any]: diff --git a/tests/portfolio/test_provider_batch.py b/tests/portfolio/test_provider_batch.py new file mode 100644 index 0000000..6e1ce83 --- /dev/null +++ b/tests/portfolio/test_provider_batch.py @@ -0,0 +1,351 @@ +"""G5 Phase 1: ``LocalUnifiedProvider.get_closes_panel`` 批量行情接口测试。 + +TDD: mock dbbardata (tmp sqlite + 小样本 fixture), 零 VPS / 零网络。 + +覆盖: +- 正确返宽表 (index=date, cols=symbol) +- 缺失 symbol → NaN 列 +- 日期范围生效 +- interval 参数 (d / 15m) +- 空 symbols → 空 df +- 重复行 dedup +- **批量 vs 逐只 get_price 结果一致** (核心回归断言) +- 纯 6 位代码 / jq 风格代码都接受 +- 混合 datetime 格式 (dbbardata 真实数据特性) +""" +from __future__ import annotations + +import sqlite3 +from typing import Any, List + +import numpy as np +import pandas as pd +import pytest + +from sanguo_portfolio.providers.local_unified_provider import LocalUnifiedProvider + + +# ======================== fixture ======================== +def _make_dbbardata_db(tmp_path, rows: List[tuple] | None = None) -> str: + """造 dbbardata 表 + 可选 rows。rows 格式: + (symbol, exchange, datetime, interval, volume, turnover, open_interest, + open_price, high_price, low_price, close_price) + """ + db = tmp_path / "t.db" + c = sqlite3.connect(str(db)) + c.execute( + "CREATE TABLE dbbardata(symbol TEXT, exchange TEXT, datetime TEXT, " + "interval TEXT, volume REAL, turnover REAL, open_interest REAL, " + "open_price REAL, high_price REAL, low_price REAL, close_price REAL)" + ) + if rows: + c.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows) + c.commit() + c.close() + return str(db) + + +def _row(symbol: str, exchange: str, dt: str, close: float, interval: str = "d") -> tuple: + return (symbol, exchange, dt, interval, 1000.0, 1e6, 0, + close, close + 5, close - 5, close) + + +@pytest.fixture +def batch_provider(tmp_path): + """3 只股票 × 3 天日线 fixture。 + + - 600519.XSHG (SSE): close 1000, 1010, 1020 + - 000001.XSHE (SZSE): close 10, 11, 12 + - 600811.XSHG (SSE): close 5, 6, 7 (用于测多股) + 缺失股票 999999.XSHG 无任何行 (NaN 列验证) + """ + rows = [] + for sym, exc, closes in [ + ("600519", "SSE", [1000.0, 1010.0, 1020.0]), + ("000001", "SZSE", [10.0, 11.0, 12.0]), + ("600811", "SSE", [5.0, 6.0, 7.0]), + ]: + for i, dt in enumerate(["2024-06-18 00:00:00", + "2024-06-19 00:00:00", + "2024-06-20 00:00:00"]): + rows.append(_row(sym, exc, dt, closes[i])) + db = _make_dbbardata_db(tmp_path, rows) + return LocalUnifiedProvider({"db_path": db, "data_dir": str(tmp_path)}) + + +# ======================== 1. 宽表格式 ======================== +class TestWideTableFormat: + def test_returns_dataframe_with_datetime_index(self, batch_provider): + df = batch_provider.get_closes_panel( + ["600519.XSHG"], start="2024-06-18", end="2024-06-20" + ) + assert isinstance(df, pd.DataFrame) + assert isinstance(df.index, pd.DatetimeIndex) + + def test_columns_match_input_symbols_order(self, batch_provider): + symbols = ["600519.XSHG", "000001.XSHE", "600811.XSHG"] + df = batch_provider.get_closes_panel(symbols, start="2024-06-18", end="2024-06-20") + assert list(df.columns) == symbols + + def test_values_are_close_prices(self, batch_provider): + df = batch_provider.get_closes_panel( + ["600519.XSHG"], start="2024-06-18", end="2024-06-20" + ) + # 3 天, close 1000/1010/1020 + assert len(df) == 3 + assert abs(df.loc[pd.Timestamp("2024-06-18"), "600519.XSHG"] - 1000.0) < 1e-6 + assert abs(df.loc[pd.Timestamp("2024-06-20"), "600519.XSHG"] - 1020.0) < 1e-6 + + def test_index_sorted_ascending(self, batch_provider): + df = batch_provider.get_closes_panel( + ["600519.XSHG"], start="2024-06-18", end="2024-06-20" + ) + assert list(df.index) == sorted(df.index) + + def test_index_name_is_none(self, batch_provider): + # 与 get_price 一致(index.name=None) + df = batch_provider.get_closes_panel( + ["600519.XSHG"], start="2024-06-18", end="2024-06-20" + ) + assert df.index.name is None + + +# ======================== 2. 缺失 symbol → NaN ======================== +class TestMissingSymbol: + def test_missing_symbol_returns_nan_column(self, batch_provider): + # 999999.XSHG 无任何行 + df = batch_provider.get_closes_panel( + ["600519.XSHG", "999999.XSHG"], start="2024-06-18", end="2024-06-20" + ) + assert "999999.XSHG" in df.columns + assert df["999999.XSHG"].isna().all() + + def test_all_missing_returns_empty_with_columns(self, batch_provider): + df = batch_provider.get_closes_panel( + ["999999.XSHG", "888888.XSHG"], start="2024-06-18", end="2024-06-20" + ) + # 全 NaN,空 index,但列保留(对齐策略 select) + assert list(df.columns) == ["999999.XSHG", "888888.XSHG"] + assert df.empty + assert df["999999.XSHG"].isna().all() + + def test_partial_missing_keeps_present_columns(self, batch_provider): + df = batch_provider.get_closes_panel( + ["600519.XSHG", "MISSING.XSHG"], start="2024-06-18", end="2024-06-20" + ) + # 600519 有 close 数据, MISSING 全 NaN + assert not df["600519.XSHG"].isna().all() + assert df["MISSING.XSHG"].isna().all() + + +# ======================== 3. 日期范围 + interval 参数 ======================== +class TestDateRangeAndInterval: + def test_start_date_excludes_earlier(self, batch_provider): + # start=06-19 → 排除 06-18 + df = batch_provider.get_closes_panel( + ["600519.XSHG"], start="2024-06-19", end="2024-06-20" + ) + assert len(df) == 2 + assert pd.Timestamp("2024-06-18") not in df.index + + def test_end_date_excludes_later(self, batch_provider): + df = batch_provider.get_closes_panel( + ["600519.XSHG"], start="2024-06-18", end="2024-06-19" + ) + assert len(df) == 2 + assert pd.Timestamp("2024-06-20") not in df.index + + def test_interval_d_filters_15m_rows(self, tmp_path): + # 同时插 d 和 15m 行, interval=d 只取日线 + rows = [ + _row("600519", "SSE", "2024-06-18 00:00:00", 1000.0, interval="d"), + _row("600519", "SSE", "2024-06-18 09:45:00", 1005.0, interval="15m"), + ] + db = _make_dbbardata_db(tmp_path, rows) + p = LocalUnifiedProvider({"db_path": db, "data_dir": str(tmp_path)}) + df = p.get_closes_panel(["600519.XSHG"], start="2024-06-18", end="2024-06-19", interval="d") + assert len(df) == 1 + assert abs(df.iloc[0]["600519.XSHG"] - 1000.0) < 1e-6 + + def test_interval_15m_returns_intraday(self, tmp_path): + rows = [ + _row("600519", "SSE", "2024-06-18 09:45:00", 1005.0, interval="15m"), + _row("600519", "SSE", "2024-06-18 10:00:00", 1008.0, interval="15m"), + ] + db = _make_dbbardata_db(tmp_path, rows) + p = LocalUnifiedProvider({"db_path": db, "data_dir": str(tmp_path)}) + df = p.get_closes_panel( + ["600519.XSHG"], start="2024-06-18", end="2024-06-18", interval="15m" + ) + assert len(df) == 2 + + +# ======================== 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) == [] + + def test_duplicate_rows_deduped_keep_last(self, tmp_path): + # 同 (symbol, datetime) 两行 close 不同 → dedup keep last + rows = [ + _row("600519", "SSE", "2024-06-18 00:00:00", 1000.0), + _row("600519", "SSE", "2024-06-18 00:00:00", 1050.0), # 后写覆盖 + _row("600519", "SSE", "2024-06-19 00:00:00", 1010.0), + ] + db = _make_dbbardata_db(tmp_path, rows) + p = LocalUnifiedProvider({"db_path": db, "data_dir": str(tmp_path)}) + df = p.get_closes_panel(["600519.XSHG"], start="2024-06-18", end="2024-06-19") + assert len(df) == 2 # dedup 后 2 行 + # 06-18 取 last = 1050 + assert abs(df.loc[pd.Timestamp("2024-06-18"), "600519.XSHG"] - 1050.0) < 1e-6 + + +# ======================== 5. 输入代码格式 ======================== +class TestInputCodeFormats: + def test_pure_digit_codes_accepted(self, batch_provider): + # 纯 6 位代码 (无 .XSHG 后缀) — jq_to_dbbardata 按首位推断 exchange + df = batch_provider.get_closes_panel( + ["600519", "000001"], start="2024-06-18", end="2024-06-20" + ) + assert "600519" in df.columns + assert "000001" in df.columns + assert abs(df.loc[pd.Timestamp("2024-06-18"), "600519"] - 1000.0) < 1e-6 + assert abs(df.loc[pd.Timestamp("2024-06-18"), "000001"] - 10.0) < 1e-6 + + def test_mixed_formats_keep_input_as_column(self, batch_provider): + # 混用 jq + 纯数字 — 列名按输入顺序保留 + df = batch_provider.get_closes_panel( + ["600519.XSHG", "000001"], start="2024-06-18", end="2024-06-20" + ) + assert list(df.columns) == ["600519.XSHG", "000001"] + + +# ======================== 6. 混合 datetime 格式 ======================== +class TestMixedDatetimeFormat: + """dbbardata datetime 列混合格式(有 "2024-09-26" 也有 "2024-09-26 00:00:00")。 + + pandas 2.3 严格模式要 format="mixed"(与 get_price 一致)。 + """ + + def test_mixed_datetime_does_not_crash(self, tmp_path): + rows = [ + _row("600519", "SSE", "2024-09-25", 1000.0), # 纯日期 + _row("600519", "SSE", "2024-09-26 00:00:00", 1010.0), # 带时间 + _row("600519", "SSE", "2024-09-27", 1020.0), # 纯日期 + ] + db = _make_dbbardata_db(tmp_path, rows) + p = LocalUnifiedProvider({"db_path": db, "data_dir": str(tmp_path)}) + df = p.get_closes_panel(["600519.XSHG"], start="2024-09-25", end="2024-09-27") + assert len(df) == 3 + assert abs(df.loc[pd.Timestamp("2024-09-26"), "600519.XSHG"] - 1010.0) < 1e-6 + + +# ======================== 7. 核心回归: 批量 vs 逐只 get_price ======================== +class TestBatchMatchesPerSymbolGetPrice: + """批量 ``get_closes_panel`` 必须与逐只 ``get_price(fq='raw')`` 取的 close 一致。 + + 这是 G5 Phase 1 的核心契约: 同样的 dbbardata raw close,只改 IO 路径(1 SQL vs N SQL)。 + 数值结果应完全相同(策略替换零回归)。 + """ + + def test_batch_equals_per_symbol_get_price(self, batch_provider): + symbols = ["600519.XSHG", "000001.XSHE", "600811.XSHG"] + start, end = "2024-06-18", "2024-06-20" + + # 批量 + batch = batch_provider.get_closes_panel(symbols, start=start, end=end) + + # 逐只 get_price(fq='raw' 取原始 close) + for sym in symbols: + single = batch_provider.get_price( + sym, start_date=start, end_date=end, fq="raw" + ) + assert "close" in single.columns, f"{sym}: get_price 返缺 close 列" + # 索引对齐后逐值比较 + batch_col = batch[sym].dropna() + single_col = single["close"] + # 索引应一致(都升序 + 同 datetime) + assert list(batch_col.index) == list(single_col.index), ( + f"{sym}: index mismatch\n batch={list(batch_col.index)}\n single={list(single_col.index)}" + ) + # 值完全相同(浮点精确相等 — 同列同源数据) + np.testing.assert_array_equal( + batch_col.values, single_col.values, + err_msg=f"{sym}: values mismatch" + ) + + def test_batch_matches_with_missing_symbol(self, batch_provider): + # 含缺失股票的批量 vs 逐只 — 缺失股票逐只返空, 批量列全 NaN + symbols = ["600519.XSHG", "MISSING.XSHG"] + batch = batch_provider.get_closes_panel(symbols, start="2024-06-18", end="2024-06-20") + # 600519 一致 + single_519 = batch_provider.get_price( + "600519.XSHG", start_date="2024-06-18", end_date="2024-06-20", fq="raw" + ) + np.testing.assert_array_equal( + batch["600519.XSHG"].values, single_519["close"].values + ) + # MISSING 全 NaN + assert batch["MISSING.XSHG"].isna().all() + + def test_batch_matches_pure_digit_codes(self, batch_provider): + # 纯 6 位代码也要能 vs get_price 对齐 + symbols = ["600519", "000001"] + batch = batch_provider.get_closes_panel(symbols, start="2024-06-18", end="2024-06-20") + for sym in symbols: + single = batch_provider.get_price( + sym, start_date="2024-06-18", end_date="2024-06-20", fq="raw" + ) + np.testing.assert_array_equal( + batch[sym].dropna().values, single["close"].values + ) + + +# ======================== 8. 单条 SQL (性能契约) ======================== +class TestSingleSqlCommand: + """G5 Phase 1 性能目标: 1 条 SQL 替代 N 条。 + + 通过 monkeypatch ``pd.read_sql`` 计数,确保只调一次(不论 symbols 多少)。 + """ + + def test_issues_one_sql_regardless_of_symbol_count(self, batch_provider, monkeypatch): + call_count = {"n": 0} + original_read_sql = pd.read_sql + + def counting_read_sql(*args, **kwargs): + call_count["n"] += 1 + return original_read_sql(*args, **kwargs) + + monkeypatch.setattr("sanguo_portfolio.providers.local_unified_provider.pd.read_sql", counting_read_sql) + + # 3 只股票 — 仍只 1 次 SQL + batch_provider.get_closes_panel( + ["600519.XSHG", "000001.XSHE", "600811.XSHG"], + start="2024-06-18", end="2024-06-20", + ) + assert call_count["n"] == 1, f"期望 1 条 SQL,实际 {call_count['n']}" + + def test_issues_one_sql_with_ten_symbols(self, tmp_path, monkeypatch): + # 造 10 只股票 fixture + rows = [] + for i in range(10): + sym = f"600{i:03d}" # 600000, 600001, ... + rows.append(_row(sym, "SSE", "2024-06-18 00:00:00", 100.0 + i)) + db = _make_dbbardata_db(tmp_path, rows) + p = LocalUnifiedProvider({"db_path": db, "data_dir": str(tmp_path)}) + + call_count = {"n": 0} + original_read_sql = pd.read_sql + + def counting_read_sql(*args, **kwargs): + call_count["n"] += 1 + return original_read_sql(*args, **kwargs) + + monkeypatch.setattr("sanguo_portfolio.providers.local_unified_provider.pd.read_sql", counting_read_sql) + + symbols = [f"600{i:03d}.XSHG" for i in range(10)] + p.get_closes_panel(symbols, start="2024-06-18", end="2024-06-18") + assert call_count["n"] == 1