b270faf4b9
- 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)
429 lines
17 KiB
Python
429 lines
17 KiB
Python
"""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()
|