"""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