Files
sanguo_vnpy_v2/tests/portfolio/test_baostock_provider.py
T
claude_dev b270faf4b9 feat(portfolio): 本地数据 provider 层(baostock/local_parquet)
- BaostockProvider: 读 VPS daily_baostock_full(本地,不调online,守 provider-local-data-only 铁律)
- LocalParquetProvider: 读 parquet 兜底,回测117交易日0.4s/月出JSON
- all_weather 策略 + runner_backtest 适配
- 数据源融合使用层(单 Provider 内部路由,见 data-fusion spec §6)
2026-07-22 10:35:23 +08:00

429 lines
17 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.
"""BaostockProvider 单元测试(mock baostock)。
baostock 装在 venv310,单测仍用 mock(避免依赖网络 + 快速 + 可重复)。
覆盖:
- ``jq_to_bs_code`` / ``bs_to_jq_code`` 代码格式转换
- ``get_price`` mock K 线,断言 DataFrame 格式 + 字段
- ``get_index_stocks`` mock baostock 返成分股,断言 jq 格式转换 + 历史日期透传
- ``get_fundamentals_df`` TTM 4 季自滚(构造累计财务 mock,断言 TTM 公式正确)
- ``get_security_info`` mock query_stock_basic 返 display_name/start_date
- ``get_current_tick`` 从 K 线推涨跌停(preclose × 1.1)
"""
from __future__ import annotations
import math
from typing import Any, Dict, List
import pandas as pd
import pytest
from sanguo_portfolio import BaostockProvider
from sanguo_portfolio.providers.baostock_provider import (
bs_to_jq_code, jq_to_bs_code, _latest_available_quarter,
)
# ======================== 代码格式转换 ========================
class TestCodeFormat:
def test_jq_to_bs_code_sh(self):
# Arrange + Act + Assert
assert jq_to_bs_code("600519.XSHG") == "sh.600519"
def test_jq_to_bs_code_sz(self):
assert jq_to_bs_code("000001.XSHE") == "sz.000001"
def test_jq_to_bs_code_pure_digit_sh(self):
# 6 开头 → sh
assert jq_to_bs_code("600519") == "sh.600519"
def test_jq_to_bs_code_pure_digit_sz(self):
# 0/3 开头 → sz
assert jq_to_bs_code("000001") == "sz.000001"
def test_bs_to_jq_code_sh(self):
assert bs_to_jq_code("sh.600519") == "600519.XSHG"
def test_bs_to_jq_code_sz(self):
assert bs_to_jq_code("sz.000001") == "000001.XSHE"
def test_round_trip_jq_to_bs_to_jq(self):
# Arrange
original = "600519.XSHG"
# Act
rt = bs_to_jq_code(jq_to_bs_code(original))
# Assert
assert rt == original
def test_jq_to_bs_code_already_bs(self):
# 已是 baostock 风格 → 透传
assert jq_to_bs_code("sh.600519") == "sh.600519"
# ======================== 季度推算 ========================
class TestLatestAvailableQuarter:
def test_jan_to_april_returns_prev_year_q3(self):
# 1/1 ~ 4/30 → 上年 Q3
assert _latest_available_quarter("2024-01-15") == (2023, 3, "2023-09-30")
assert _latest_available_quarter("2024-04-30") == (2023, 3, "2023-09-30")
def test_may_to_aug_returns_current_q1(self):
# 5/1 ~ 8/31 → 当年 Q1
assert _latest_available_quarter("2024-05-01") == (2024, 1, "2024-03-31")
assert _latest_available_quarter("2024-08-31") == (2024, 1, "2024-03-31")
def test_sep_to_oct_returns_current_q2(self):
assert _latest_available_quarter("2024-09-15") == (2024, 2, "2024-06-30")
assert _latest_available_quarter("2024-10-31") == (2024, 2, "2024-06-30")
def test_nov_to_dec_returns_current_q3(self):
assert _latest_available_quarter("2024-11-01") == (2024, 3, "2024-09-30")
assert _latest_available_quarter("2024-12-31") == (2024, 3, "2024-09-30")
# ======================== get_price ========================
class TestGetPrice:
def test_returns_dataframe_with_close(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
df = provider.get_price(
"600519.XSHG", end_date="2024-09-30", frequency="daily",
fields=["close"], count=2, panel=False,
)
# Assert
assert isinstance(df, pd.DataFrame)
assert len(df) == 2
assert "close" in df.columns
# 数值已转 float(baostock 返回字符串)
assert df["close"].dtype.kind == "f"
# close 1600 / 1620(mock 数据)
assert df["close"].iloc[-1] == pytest.approx(1620.0)
def test_code_column_is_jq_style(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
df = provider.get_price(
"600519.XSHG", end_date="2024-09-30",
fields=["close"], count=1, panel=False,
)
# Assert
assert df.iloc[0]["code"] == "600519.XSHG"
def test_multiple_stocks_panel_false_returns_long_format(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
df = provider.get_price(
["600519.XSHG", "601318.XSHG"], end_date="2024-09-30",
fields=["close"], count=1, panel=False,
)
# Assert
assert isinstance(df, pd.DataFrame)
codes = set(df["code"].unique())
assert codes == {"600519.XSHG", "601318.XSHG"}
def test_baostock_query_failure_returns_empty(self, mock_baostock):
# Arrange
mock_baostock["bs"].query_history_k_data_plus.side_effect = Exception("network")
provider = BaostockProvider({})
# Act
df = provider.get_price("600519.XSHG", end_date="2024-09-30", count=1)
# Assert
assert isinstance(df, pd.DataFrame)
assert df.empty
def test_count_takes_last_n_rows(self, mock_baostock):
# Arrange:mock K 线默认 2 根
provider = BaostockProvider({})
# Act
df = provider.get_price(
"600519.XSHG", end_date="2024-09-30",
fields=["close"], count=1, panel=False,
)
# Assert
assert len(df) == 1
# tail(1) 取最后一根
assert df["close"].iloc[0] == pytest.approx(1620.0)
# ======================== get_index_stocks(历史日期) ========================
class TestGetIndexStocks:
def test_hs300_returns_jq_codes(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
stocks = provider.get_index_stocks("000300.XSHG", "2024-09-30")
# Assert
assert len(stocks) == 3
# jq 格式:600519.XSHG / 601318.XSHG / 000001.XSHE
assert "600519.XSHG" in stocks
assert "000001.XSHE" in stocks
def test_hs300_passes_date_to_baostock(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
bs_mock = mock_baostock["bs"]
# Act
provider.get_index_stocks("000300.XSHG", "2020-06-15")
# Assert
bs_mock.query_hs300_stocks.assert_called_once()
args, kwargs = bs_mock.query_hs300_stocks.call_args
# baostock 接口:date="" (positional) 或 date= kwargs
passed_date = args[0] if args else kwargs.get("date")
assert passed_date == "2020-06-15"
def test_zz500_routes_to_zz500_query(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
provider.get_index_stocks("000905.XSHG", "2024-01-01")
# Assert
mock_baostock["bs"].query_zz500_stocks.assert_called_once()
def test_sz50_routes_to_sz50_query(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
provider.get_index_stocks("000016.XSHG", "2024-01-01")
# Assert
mock_baostock["bs"].query_sz50_stocks.assert_called_once()
def test_cache_same_date_same_index(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
provider.get_index_stocks("000300.XSHG", "2024-09-30")
provider.get_index_stocks("000300.XSHG", "2024-09-30")
# Assert:第二次走缓存,query_hs300_stocks 只调一次
assert mock_baostock["bs"].query_hs300_stocks.call_count == 1
# ======================== get_fundamentals_df(TTM 自滚) ========================
class TestFundamentalsTTM:
def test_empty_stocks_returns_empty_df(self, mock_baostock):
provider = BaostockProvider({})
df = provider.get_fundamentals_df([], date="2024-09-30")
assert isinstance(df, pd.DataFrame)
assert len(df) == 0
assert "code" in df.columns
def test_ttm_net_profit_formula_correct(self, mock_baostock):
"""TTM 公式:本期YTD - 上年同期YTD + 上年全年。
构造 3 期 mock:
- 本期(2024 Q3):YTD netProfit = 500 亿(前三季累计)
- 上年同期(2023 Q3):YTD netProfit = 400 亿
- 上年全年(2023 Q4):YTD netProfit = 600 亿
TTM = 500 - 400 + 600 = 700 亿
date="2024-11-15" → _latest_available_quarter 推算最近已披露季度 = 2024 Q3
(11月1日~12月31日区间,三季报披露窗口 10/31 已结束,Q3 数据可用)
"""
# Arrange:date 2024-11-15 → 推算最近季度 = 2024 Q3
# _latest_available_quarter("2024-11-15") = (2024, 3) ✓
provider = BaostockProvider({})
def profit_side_effect(code, year=None, quarter=None):
from tests.portfolio.conftest import _FakeResultData, _build_default_profit_df
# 构造不同 (year, quarter) 返不同累计值
if (year, quarter) == (2024, 3):
df = _build_default_profit_df(net_profit_ytd=500e8, revenue_ytd=1000e8)
elif (year, quarter) == (2023, 3):
df = _build_default_profit_df(net_profit_ytd=400e8, revenue_ytd=900e8)
elif (year, quarter) == (2023, 4):
df = _build_default_profit_df(net_profit_ytd=600e8, revenue_ytd=1200e8)
else:
df = pd.DataFrame()
return _FakeResultData(df)
mock_baostock["bs"].query_profit_data.side_effect = profit_side_effect
# close 单价 1600,股本 1.256e9
# 市值 = 1600 * 1.256e9 = 2.0096e12 = 20096 亿元
# PE_TTM = 市值 / TTM净利 = 2.0096e12 / 700e8 = 28.7
# Act
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-11-15")
# Assert
assert len(df) == 1
row = df.iloc[0]
# TTM 净利正确:从 _net_profit_ttm 列读取
ttm = float(row["_net_profit_ttm"])
assert ttm == pytest.approx(700e8, rel=1e-6), f"TTM 净利={ttm},期望 700 亿"
# PE 反算合理
pe = float(row["pe_ratio"])
assert pe == pytest.approx(28.7, rel=0.05)
def test_ttm_revenue_formula_correct(self, mock_baostock):
"""TTM 营收 = 本期YTD - 上年同期YTD + 上年全年。
构造:本期=1000亿 / 上年同期=900亿 / 上年全年=1200亿 → TTM=1300亿
"""
provider = BaostockProvider({})
def profit_side_effect(code, year=None, quarter=None):
from tests.portfolio.conftest import _FakeResultData, _build_default_profit_df
mapping = {
(2024, 3): (500e8, 1000e8),
(2023, 3): (400e8, 900e8),
(2023, 4): (600e8, 1200e8),
}
np_, rev_ = mapping.get((year, quarter), (None, None))
if np_ is None:
return _FakeResultData(pd.DataFrame())
return _FakeResultData(_build_default_profit_df(net_profit_ytd=np_, revenue_ytd=rev_))
mock_baostock["bs"].query_profit_data.side_effect = profit_side_effect
# Act
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-11-15")
# Assert
ttm_rev = float(df.iloc[0]["_revenue_ttm"])
assert ttm_rev == pytest.approx(1300e8, rel=1e-6), f"TTM 营收={ttm_rev},期望 1300 亿"
def test_fallback_single_quarter_x4_when_insufficient_history(self, mock_baostock, caplog):
"""不足 3 期历史(新股)→ fallback 单期×4 + WARNING。
构造:只本期返数据(上年同期/上年全年空表),则 TTM = 500 亿 × 4 = 2000 亿。
"""
import logging
provider = BaostockProvider({})
def profit_side_effect(code, year=None, quarter=None):
from tests.portfolio.conftest import _FakeResultData, _build_default_profit_df
# 只有本期(2024 Q3)有数据
if (year, quarter) == (2024, 3):
return _FakeResultData(_build_default_profit_df(net_profit_ytd=500e8, revenue_ytd=1000e8))
return _FakeResultData(pd.DataFrame()) # 空表
mock_baostock["bs"].query_profit_data.side_effect = profit_side_effect
# Act
with caplog.at_level(logging.WARNING):
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-11-15")
# Assert
ttm = float(df.iloc[0]["_net_profit_ttm"])
assert ttm == pytest.approx(2000e8, rel=1e-6), f"fallback TTM={ttm},期望 500亿×4=2000亿"
# WARNING 日志确认
assert any("不足 3 期" in r.message for r in caplog.records)
def test_market_cap_in_yi_unit(self, mock_baostock):
"""close × total_share / 1e8 = 亿元。茅台 1600 × 1.256e9 / 1e8 ≈ 20096 亿。"""
provider = BaostockProvider({})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-11-15")
mc = float(df.iloc[0]["market_cap"])
assert 19000 < mc < 22000, f"market_cap 异常: {mc}"
def test_required_columns_present(self, mock_baostock):
provider = BaostockProvider({})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-11-15")
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}"
def test_roe_pct_to_decimal(self, mock_baostock):
"""baostock roeAvg=30.0(百分数) → 归一到 0.30 小数。"""
provider = BaostockProvider({})
df = provider.get_fundamentals_df(["600519.XSHG"], date="2024-11-15")
roe = float(df.iloc[0]["roe"])
assert roe == pytest.approx(0.30, abs=0.01)
# ======================== get_security_info ========================
class TestGetSecurityInfo:
def test_returns_display_name_and_start_date(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
# Act
info = provider.get_security_info("600519.XSHG")
# Assert
assert info["display_name"] == "贵州茅台"
assert info["start_date"] is not None
# start_date 应可解析为 date
from datetime import date
assert isinstance(info["start_date"], date) or hasattr(info["start_date"], "year")
def test_jq_code_passed_to_baostock_as_bs_code(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
bs_mock = mock_baostock["bs"]
# Act
provider.get_security_info("600519.XSHG")
# Assert:baostock 收到的是 sh.600519
args, kwargs = bs_mock.query_stock_basic.call_args
passed_code = args[0] if args else kwargs.get("code")
assert passed_code == "sh.600519"
def test_query_failure_returns_fallback_dict(self, mock_baostock):
# Arrange
mock_baostock["bs"].query_stock_basic.side_effect = Exception("network")
provider = BaostockProvider({})
# Act
info = provider.get_security_info("600519.XSHG")
# Assert:不应抛异常,fallback 返 jq code 作 display_name
assert "display_name" in info
assert info["start_date"] is None
# ======================== get_current_tick(涨跌停) ========================
class TestGetCurrentTick:
def test_high_limit_is_preclose_x_1_1(self, mock_baostock):
"""主板涨跌停:preclose × 1.1 / 0.9。
mock K 线最后一根:close=1620, preclose=1600。
high_limit = 1600 × 1.1 = 1760,low_limit = 1600 × 0.9 = 1440。
"""
provider = BaostockProvider({})
tick = provider.get_current_tick("600519.XSHG")
assert tick is not None
assert tick["last_price"] == pytest.approx(1620.0)
assert tick["high_limit"] == pytest.approx(1760.0, abs=0.01)
assert tick["low_limit"] == pytest.approx(1440.0, abs=0.01)
assert tick["paused"] is False
def test_st_uses_5_percent_limit(self, mock_baostock):
"""isST=1 → 涨跌停 5%。preclose=1600 → high=1680, low=1520。"""
from tests.portfolio.conftest import _FakeResultData, _build_default_kline_df
# Arrange:把 isST 改 1
df = _build_default_kline_df()
df["isST"] = ["1", "1"]
mock_baostock["bs"].query_history_k_data_plus.return_value = _FakeResultData(df)
provider = BaostockProvider({})
tick = provider.get_current_tick("600519.XSHG")
assert tick is not None
assert tick["high_limit"] == pytest.approx(1680.0, abs=0.01)
assert tick["low_limit"] == pytest.approx(1520.0, abs=0.01)
# ======================== Provider metadata ========================
class TestProviderMetadata:
def test_name_is_sanguo_baostock(self):
# Arrange + Act + Assert
assert BaostockProvider.name == "sanguo_baostock"
def test_requires_live_data_false(self):
# 回测 provider,不要求实时行情
assert BaostockProvider.requires_live_data is False
def test_login_logout_lifecycle(self, mock_baostock):
# Arrange
provider = BaostockProvider({})
bs_mock = mock_baostock["bs"]
# Act:login 是惰性,第一次 query 触发
provider.get_security_info("600519.XSHG")
# Assert
bs_mock.login.assert_called_once()
# 再调一次,login 不再触发
provider.get_security_info("601318.XSHG")
assert bs_mock.login.call_count == 1
# close() 触发 logout
provider.close()
bs_mock.logout.assert_called_once()