Files
sanguo_vnpy_v2/tests/portfolio/test_fetchers.py
T
claude_dev f9322a7c48
CI/CD / test (push) Successful in 21s
CI/CD / nas-deploy (push) Successful in 1m3s
CI/CD / nas-verify (push) Successful in 12s
feat(provider): TET窄试点B——P0四方法Fetcher化(_ex新接口,strict fail-fast),老接口零改动 [vps]
2026-08-14 23:44:39 +08:00

245 lines
12 KiB
Python

"""TET Fetcher(窄试点B)测试: 等价性 + strict fail-fast + Optional 契约。
三组覆盖(方案 2026-08-15 用户定案):
1. **TestEquivalence**: 合法参数下 `_ex` vs 老接口输出逐值一致(extract SQL
原样搬迁的回归保障;Phase 2 真策略回测对照前的第一道防线)。
2. **TestStrictFailFast**: 非法参数(老接口静默返空/取默认)→ ValueError;
脏数据(核心列全空)→ DataSchemaError。这是 _ex 的有意新契约。
3. **TestOptionalContract**: high_limit/low_limit/paused 缺失补默认=显式契约
(非兜底;补 NaN 会被 bullet_trade bool(NaN)=True 当停牌全撤单)。
fixture 模式照 test_local_unified_provider.py(tmp_path sqlite 自建表,零网络)。
"""
from __future__ import annotations
import sqlite3
import pandas as pd
import pytest
from pydantic import ValidationError
from sanguo_portfolio.providers.fetchers import DataSchemaError
from sanguo_portfolio.providers.local_unified_provider import LocalUnifiedProvider
@pytest.fixture
def unified_provider(tmp_path):
"""dbbardata(含脏数据股) + bs_adjust_factor + constituent_unified。"""
db = tmp_path / "quant_trading.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)"
)
rows = [
# 600519: 除权日 2024-06-19 close 1000→900 跳水(qfq 用)
("600519", "SSE", "2024-06-18 00:00:00", "d", 1000, 1e6, 0,
1000.0, 1010.0, 990.0, 1000.0),
("600519", "SSE", "2024-06-19 00:00:00", "d", 1000, 1e6, 0,
900.0, 910.0, 890.0, 900.0),
("600519", "SSE", "2024-06-20 00:00:00", "d", 1000, 1e6, 0,
910.0, 920.0, 900.0, 910.0),
# 000001 深市 2 天(panel 多标的/宽表用)
("000001", "SZSE", "2024-06-18 00:00:00", "d", 2e6, 2e7, 0,
10.0, 10.5, 9.8, 10.2),
("000001", "SZSE", "2024-06-19 00:00:00", "d", 2e6, 2e7, 0,
10.3, 10.6, 10.1, 10.5),
# 600999 脏数据: 有行但 close 全 NULL(strict 校验靶)
("600999", "SSE", "2024-06-18 00:00:00", "d", 100, 1e3, 0,
None, None, None, None),
]
c.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
c.execute(
"CREATE TABLE bs_adjust_factor(code TEXT, dividOperateDate TEXT, "
"foreAdjustFactor REAL, backAdjustFactor REAL, adjustFactor REAL)"
)
c.execute(
"INSERT INTO bs_adjust_factor VALUES('sh.600519','2024-06-19',0.9,0,0)"
)
c.execute(
"CREATE TABLE constituent_unified(index_code TEXT, code TEXT, "
"in_current INTEGER, was_removed INTEGER, code_name TEXT)"
)
c.executemany(
"INSERT INTO constituent_unified VALUES(?,?,?,?,?)",
[
("000300", "600519", 1, 0, "贵州茅台"),
("000300", "000001", 0, 1, "平安银行"), # 被踢(治偏差)
("000300", "bad-code", 1, 0, "非法码"), # transform 清洗靶
],
)
c.commit()
c.close()
# data_dir 指向空目录: fundamentals 无 static parquet → 全 NaN 行(两接口同路径)
(tmp_path / "static").mkdir(exist_ok=True)
return LocalUnifiedProvider({"db_path": str(db), "data_dir": str(tmp_path)})
# ======================== 1. 等价性: _ex vs 老接口 ========================
class TestEquivalence:
"""合法参数下逐值一致(assert_frame_equal 默认容差 rtol=1e-5)。"""
def test_get_price_raw(self, unified_provider):
old = unified_provider.get_price(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-20")
new = unified_provider.get_price_ex(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-20")
pd.testing.assert_frame_equal(old, new)
def test_get_price_qfq(self, unified_provider):
old = unified_provider.get_price(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-20", fq="qfq")
new = unified_provider.get_price_ex(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-20", fq="qfq")
pd.testing.assert_frame_equal(old, new)
def test_get_price_count_and_fields(self, unified_provider):
kwargs = dict(security="600519.XSHG", count=2,
fields=["close", "high_limit", "paused"], panel=False)
old = unified_provider.get_price(**kwargs)
new = unified_provider.get_price_ex(**kwargs)
pd.testing.assert_frame_equal(old, new)
def test_get_price_multi_panel(self, unified_provider):
secs = ["600519.XSHG", "000001.XSHE"]
old = unified_provider.get_price(
secs, start_date="2024-06-18", end_date="2024-06-20", panel=False)
new = unified_provider.get_price_ex(
secs, start_date="2024-06-18", end_date="2024-06-20", panel=False)
pd.testing.assert_frame_equal(old, new)
def test_get_price_missing_symbol_empty_both(self, unified_provider):
# 库中无数据的标的: 两接口都返空 DataFrame(空=合法缺失)
old = unified_provider.get_price("300999.XSHE", start_date="2024-06-18")
new = unified_provider.get_price_ex("300999.XSHE", start_date="2024-06-18")
assert old.empty and new.empty
def test_get_closes_panel_raw_and_missing_col(self, unified_provider):
syms = ["600519.XSHG", "000001.XSHE", "300999.XSHE"] # 含缺失标的
old = unified_provider.get_closes_panel(
syms, start="2024-06-18", end="2024-06-20")
new = unified_provider.get_closes_panel_ex(
syms, start="2024-06-18", end="2024-06-20")
pd.testing.assert_frame_equal(old, new)
assert new["300999.XSHE"].isna().all() # 缺失标的补全 NaN 列(契约一致)
def test_get_closes_panel_qfq(self, unified_provider):
old = unified_provider.get_closes_panel(
["600519.XSHG"], start="2024-06-18", end="2024-06-20", fq="qfq")
new = unified_provider.get_closes_panel_ex(
["600519.XSHG"], start="2024-06-18", end="2024-06-20", fq="qfq")
pd.testing.assert_frame_equal(old, new)
def test_get_constituent(self, unified_provider):
old = unified_provider.get_constituent("000300.XSHG")
new = unified_provider.get_constituent_ex("000300.XSHG")
assert old == new
assert "600519.XSHG" in new and "000001.XSHE" in new # 并集(含被踢)
assert not any("bad" in s for s in new) # 非法码清洗一致
def test_get_fundamentals_df(self, unified_provider):
# 空 static 目录: 两接口同路径返全 NaN 行,验证组装/index/fields 过滤一致
old = unified_provider.get_fundamentals_df(["600519.XSHG"], date="2024-06-20")
new = unified_provider.get_fundamentals_df_ex(["600519.XSHG"], date="2024-06-20")
pd.testing.assert_frame_equal(old, new)
old2 = unified_provider.get_fundamentals_df(
["600519.XSHG", "000001.XSHE"], date="2024-06-20", fields=["market_cap"])
new2 = unified_provider.get_fundamentals_df_ex(
["600519.XSHG", "000001.XSHE"], date="2024-06-20", fields=["market_cap"])
pd.testing.assert_frame_equal(old2, new2)
# ======================== 2. strict fail-fast ========================
class TestStrictFailFast:
"""非法参数/脏数据直接报错(_ex 有意新契约;老接口静默)。"""
def test_price_bad_frequency(self, unified_provider):
# 老接口对 '1m' 静默返空 DataFrame;_ex 报错(15m 应走 panel_ex)
with pytest.raises(ValidationError, match="frequency"):
unified_provider.get_price_ex("600519.XSHG", frequency="1m")
def test_price_bad_fq(self, unified_provider):
with pytest.raises(ValidationError, match="fq"):
unified_provider.get_price_ex("600519.XSHG", fq="xx")
def test_price_bad_count(self, unified_provider):
with pytest.raises(ValidationError, match="count"):
unified_provider.get_price_ex("600519.XSHG", count=0)
def test_price_bad_date(self, unified_provider):
with pytest.raises(ValidationError, match="日期"):
unified_provider.get_price_ex("600519.XSHG", start_date="2024/06/18")
def test_price_empty_security(self, unified_provider):
with pytest.raises(ValidationError, match="security"):
unified_provider.get_price_ex([])
def test_price_unknown_kwarg_forbidden(self, unified_provider):
# 拼错参数名直接报错(_ex 显式签名 → TypeError;直调 Fetcher → pydantic
# extra=forbid ValidationError;老接口 **kwargs 静默吞)
with pytest.raises((ValidationError, TypeError)):
unified_provider.get_price_ex("600519.XSHG", start_dat="2024-06-18")
def test_price_dirty_data_close_all_null(self, unified_provider):
# 600999 有行但 close 全 NULL: 老接口静默流 NaN,_ex fail-fast
with pytest.raises(DataSchemaError, match="close_price"):
unified_provider.get_price_ex(
"600999.XSHG", start_date="2024-06-18", end_date="2024-06-20")
def test_panel_empty_symbols(self, unified_provider):
with pytest.raises(ValidationError, match="symbols"):
unified_provider.get_closes_panel_ex([], start="2024-06-18", end="2024-06-20")
def test_panel_bad_date(self, unified_provider):
with pytest.raises(ValidationError, match="日期"):
unified_provider.get_closes_panel_ex(
["600519.XSHG"], start="20240618", end="2024-06-20")
def test_panel_bad_interval(self, unified_provider):
with pytest.raises(ValidationError, match="interval"):
unified_provider.get_closes_panel_ex(
["600519.XSHG"], start="2024-06-18", end="2024-06-20", interval="d; DROP TABLE")
def test_panel_dirty_data(self, unified_provider):
with pytest.raises(DataSchemaError):
unified_provider.get_closes_panel_ex(
["600999.XSHG"], start="2024-06-18", end="2024-06-20")
def test_constituent_empty_index(self, unified_provider):
with pytest.raises(ValidationError, match="index"):
unified_provider.get_constituent_ex("")
def test_fundamentals_empty_stocks(self, unified_provider):
with pytest.raises(ValidationError, match="stocks"):
unified_provider.get_fundamentals_df_ex([])
# ======================== 3. Optional 契约(非兜底) ========================
class TestOptionalContract:
"""high_limit/low_limit/paused 缺失补默认=2026-08-15 用户定案的显式契约。"""
def test_high_limit_estimated_not_nan(self, unified_provider):
df = unified_provider.get_price_ex(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-18",
fields=["close", "high_limit", "low_limit"],
)
assert abs(df["high_limit"].iloc[0] - round(1000.0 * 1.1, 2)) < 1e-6
assert abs(df["low_limit"].iloc[0] - round(1000.0 * 0.9, 2)) < 1e-6
def test_paused_false_not_nan(self, unified_provider):
df = unified_provider.get_price_ex(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-18",
fields=["close", "paused"],
)
# bool(NaN)=True 会被当停牌;契约=显式 False
assert not df["paused"].iloc[0]
def test_unknown_field_nan(self, unified_provider):
# 未知可选列(如聚宽特有字段)缺失 → NaN(声明行为,策略自查)
df = unified_provider.get_price_ex(
"600519.XSHG", start_date="2024-06-18", end_date="2024-06-18",
fields=["close", "acc_net_value"],
)
assert df["acc_net_value"].isna().all()