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:
@@ -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]:
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user