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