Files
sanguo_vnpy_v2/tests/portfolio/test_provider.py
T

275 lines
12 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.
"""SanguoMiniQmtProvider 单元测试(mock xtquant)。
bullet-trade 装了(0.2.0+),MiniQMTProvider 基类可继承。
xtquant 没 装,通过 mock_xtquant fixture 注入 sys.modules。
"""
from __future__ import annotations
import math
import pandas as pd
import pytest
from sanguo_portfolio import SanguoMiniQmtProvider
pytestmark = pytest.mark.requires_bullet_trade
class TestSanguoMiniQmtProviderInstantiation:
def test_can_instantiate_with_mock_xtquant(self, mock_xtquant):
# Arrange + Act
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
# Assert
assert provider.name == "sanguo_miniqmt"
# 应继承 MiniQMTProvider
from bullet_trade.data.providers.miniqmt import MiniQMTProvider
assert isinstance(provider, MiniQMTProvider)
class TestGetFundamentalsDf:
def test_returns_dataframe_with_required_columns(self, mock_xtquant):
# Arrange
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
stocks = ["600519.XSHG", "601318.XSHG"]
# Act
df = provider.get_fundamentals_df(stocks, date="2024-09-30")
# Assert
assert isinstance(df, pd.DataFrame)
assert len(df) == 2
# 核心列都在
for col in [
"code", "market_cap", "circulating_market_cap",
"pe_ratio", "pb_ratio", "ps_ratio", "pcf_ratio",
"roe", "roa", "eps",
"total_liability", "total_sheet_owner_equities", "retained_profit",
"roic",
]:
assert col in df.columns, f"missing col: {col}"
# index 是 jq-style code
assert "600519.XSHG" in df.index
def test_empty_stocks_returns_empty_df(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df([], date="2024-09-30")
assert isinstance(df, pd.DataFrame)
assert len(df) == 0
# 空表也要有列定义,方便上层 select
assert "code" in df.columns
def test_market_cap_in_yi_unit(self, mock_xtquant):
"""close × total_capital / 1e8 = 亿元。茅台 1600 × 12.56e8 / 1e8 = 20096 亿。"""
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-09-30")
mc = float(df.iloc[0]["market_cap"])
# 茅台市值应在 20000 亿左右(允许 close 1600±10)
assert 19000 < mc < 22000, f"market_cap 异常: {mc}"
def test_pe_ratio_finite_for_profitable_stock(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-09-30")
pe = float(df.iloc[0]["pe_ratio"])
assert math.isfinite(pe)
assert pe > 0
def test_roic_computed(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-09-30")
roic = float(df.iloc[0]["roic"])
# 茅台 mock 数据:oper=1.2e10, tax=25%, eqy=2.2e11, debt=0, cash=1.7e11
# NOPAT = 1.2e10 * 0.75 = 9e9
# IC = 2.2e11 + 0 - 1.7e11 = 5e10
# ROIC = 9e9 / 5e10 = 0.18
assert 0.05 < roic < 0.5, f"ROIC 异常: {roic}"
def test_financial_data_failure_returns_empty_df(self, mock_xtquant):
"""xtdata.get_financial_data 抛异常时返空表(不崩)。"""
mock_xtquant["xtdata"].get_financial_data.side_effect = Exception("QMT offline")
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-09-30")
assert df.empty
class TestGetFundamentalsQueryDictMode:
def test_dict_with_stocks_returns_dataframe(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
query = {"stocks": ["600519.XSHG", "601318.XSHG"], "date": "2024-09-30"}
df = provider.get_fundamentals(query)
assert len(df) == 2
def test_dict_with_filter_callable_applied(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
query = {
"stocks": ["600519.XSHG", "601318.XSHG"],
"date": "2024-09-30",
"filter": lambda d: d["roe"] > 0.3, # 只保留茅台(归一后 roe=0.30)
}
df = provider.get_fundamentals(query)
assert len(df) == 1
assert df.iloc[0]["code"] == "600519.XSHG"
def test_dict_with_order_by_applied(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
query = {
"stocks": ["600519.XSHG", "601318.XSHG"],
"date": "2024-09-30",
"order_by": [("market_cap", "desc")],
}
df = provider.get_fundamentals(query)
assert df.iloc[0]["code"] == "600519.XSHG" # 茅台市值 > 平安
def test_dict_with_limit_applied(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
query = {
"stocks": ["600519.XSHG", "601318.XSHG"],
"date": "2024-09-30",
"limit": 1,
}
df = provider.get_fundamentals(query)
assert len(df) == 1
class TestSetDataProviderInjection:
def test_set_data_provider_accepts_sanguo_provider(self, mock_xtquant):
"""set_data_provider 注入 SanguoMiniQmtProvider 实例。"""
from bullet_trade.data.api import get_data_provider, set_data_provider
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
# Act
set_data_provider(provider)
# Assert
active = get_data_provider()
assert active is provider
assert active.name == "sanguo_miniqmt"
# ======================== 本地库委托(2026-08-19 生产缺口修复) ========================
# VPS 实盘/影子 8-18/8-19 连续两天空转根因:SanguoMiniQmtProvider 缺
# get_closes_panel/get_constituent_ex + get_fundamentals_df 不认 fields +
# get_index_stocks 忽略历史日期(前后端 session 巡检实锤,见 memory
# data-session-todo-miniqmt-provider-gaps)。修法:内部持有 LocalUnifiedProvider
# 读本地 dbbardata/constituent_unified(与回测同口径)。
def _make_delegate_db(tmp_path) -> str:
"""dbbardata + constituent_unified 小样本库(委托 LocalUnifiedProvider 用)。"""
import sqlite3
db = tmp_path / "delegate.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)"
)
for i, dt in enumerate(["2024-06-18 00:00:00", "2024-06-19 00:00:00",
"2024-06-20 00:00:00"]):
c.execute(
"INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)",
("600519", "SSE", dt, "d", 1000.0, 1e6, 0,
1000.0 + i, 1005.0 + i, 995.0 + i, 1000.0 + i),
)
c.execute(
"CREATE TABLE constituent_unified(index_code TEXT, code TEXT, "
"code_name TEXT, source TEXT, in_current INT, was_removed INT)"
)
c.executemany(
"INSERT INTO constituent_unified VALUES(?,?,?,?,?,?)",
[
("000985", "600519", "贵州茅台", "csindex", 1, 0),
("000985", "000858", "五粮液", "csindex", 0, 1), # 被踢也在并集
],
)
c.commit()
c.close()
return str(db)
@pytest.fixture
def delegate_provider(mock_xtquant, tmp_path):
"""配好本地库路径的 SanguoMiniQmtProvider(db_path/data_dir 透传统一 provider)。"""
db = _make_delegate_db(tmp_path)
return SanguoMiniQmtProvider({
"db_path": db, "data_dir": str(tmp_path), "auto_download": False,
})
class TestGetClosesPanelDelegation:
"""get_closes_panel/get_closes_panel_ex 委托本地 dbbardata(同回测口径)。"""
def test_returns_wide_table_from_local_db(self, delegate_provider):
panel = delegate_provider.get_closes_panel(
["600519.XSHG"], "2024-06-18", "2024-06-20", fq="raw",
)
assert isinstance(panel, pd.DataFrame)
assert list(panel.columns) == ["600519.XSHG"]
assert len(panel) == 3
assert abs(panel["600519.XSHG"].iloc[0] - 1000.0) < 1e-6
assert abs(panel["600519.XSHG"].iloc[-1] - 1002.0) < 1e-6
def test_ex_alias_returns_same_result(self, delegate_provider):
old = delegate_provider.get_closes_panel(
["600519.XSHG"], "2024-06-18", "2024-06-20", fq="raw",
)
ex = delegate_provider.get_closes_panel_ex(
["600519.XSHG"], "2024-06-18", "2024-06-20", fq="raw",
)
pd.testing.assert_frame_equal(old, ex)
def test_missing_symbol_returns_nan_column(self, delegate_provider):
panel = delegate_provider.get_closes_panel(
["600519.XSHG", "999999.XSHG"], "2024-06-18", "2024-06-20",
)
assert list(panel.columns) == ["600519.XSHG", "999999.XSHG"]
assert panel["999999.XSHG"].isna().all()
def test_pure_digit_codes_accepted(self, delegate_provider):
panel = delegate_provider.get_closes_panel(
["600519"], "2024-06-18", "2024-06-20",
)
assert abs(panel["600519"].iloc[-1] - 1002.0) < 1e-6
class TestGetIndexStocksDelegation:
"""get_index_stocks/get_constituent_ex 委托 constituent_unified(支持历史日期口径)。"""
def test_reads_constituent_unified_union(self, delegate_provider):
stocks = delegate_provider.get_index_stocks("000985.XSHG", "2024-06-19")
# 并集语义:在册 + 被踢(was_removed)都返回
assert set(stocks) == {"600519.XSHG", "000858.XSHE"}
def test_constituent_ex_delegates_same(self, delegate_provider):
old = delegate_provider.get_index_stocks("000985.XSHG", "2024-06-19")
ex = delegate_provider.get_constituent_ex("000985.XSHG", "2024-06-19")
assert old == ex
def test_unknown_index_falls_back_to_xt_latest(self, delegate_provider, mock_xtquant, caplog):
"""表里没有的指数 → WARNING + 回退 miniQMT 最新成分(宁可降级不空转)。"""
mock_xtquant["xtdata"].get_index_weight.return_value = {"600519.SH": 0.5}
with caplog.at_level("WARNING", logger="sanguo_portfolio.providers.sanguo_fundamentals"):
stocks = delegate_provider.get_index_stocks("399303.XSHE", "2024-06-19")
assert stocks == ["600519.XSHG"]
assert any("回退" in r.message for r in caplog.records)
class TestGetFundamentalsDfFields:
"""get_fundamentals_df 加 fields 契约(对齐 unified:keep = code + 请求列)。"""
def test_fields_filters_columns(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df(
["600519.XSHG"], date="2024-09-30", fields=["market_cap", "eps"],
)
assert list(df.columns) == ["code", "market_cap", "eps"]
assert 19000 < float(df.iloc[0]["market_cap"]) < 22000
def test_fields_none_keeps_all_columns(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-09-30")
assert "roe" in df.columns # fields=None 全列(向后兼容)
def test_ex_alias_accepts_fields(self, mock_xtquant):
provider = SanguoMiniQmtProvider({"mode": "backtest", "auto_download": False})
df = provider.get_fundamentals_df_ex(
["600519.XSHG"], date="2024-09-30", fields=["eps"],
)
assert list(df.columns) == ["code", "eps"]