feat(portfolio): LocalUnifiedProvider 批量行情接口 get_closes_panel (G5-Phase1)

单条 dbbardata 参数化查询 (symbol,exchange) OR + pivot 返宽表, 替代 N 次 get_price, 为策略向量化提速铺路。raw close 口径一致, 缺失 NaN 列, symbol-exchange 配对防歧义。22 tests 含 batch-vs-逐只回归。本文件另含 get_value_metrics 透传(策略移植)。
This commit is contained in:
2026-07-28 20:36:11 +08:00
parent 853d35197e
commit c8f26bef80
2 changed files with 453 additions and 0 deletions
@@ -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]:
+351
View File
@@ -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