Files
sanguo_vnpy_v2/tests/portfolio/test_provider_batch.py
T
claude_dev c8f26bef80 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 透传(策略移植)。
2026-07-28 20:36:11 +08:00

352 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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