c8f26bef80
单条 dbbardata 参数化查询 (symbol,exchange) OR + pivot 返宽表, 替代 N 次 get_price, 为策略向量化提速铺路。raw close 口径一致, 缺失 NaN 列, symbol-exchange 配对防歧义。22 tests 含 batch-vs-逐只回归。本文件另含 get_value_metrics 透传(策略移植)。
352 lines
14 KiB
Python
352 lines
14 KiB
Python
"""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
|