diff --git a/sanguo_portfolio/providers/sanguo_fundamentals.py b/sanguo_portfolio/providers/sanguo_fundamentals.py index 927f5be..8b2f8a1 100644 --- a/sanguo_portfolio/providers/sanguo_fundamentals.py +++ b/sanguo_portfolio/providers/sanguo_fundamentals.py @@ -1,7 +1,12 @@ -"""SanguoMiniQmtProvider:继承 MiniQMTProvider,补齐 ``get_fundamentals``。 +"""SanguoMiniQmtProvider:继承 MiniQMTProvider,补齐 ``get_fundamentals`` + 本地库委托。 -BulletTrade 的 MiniQMTProvider 实现了行情/成分/涨跌停/证券信息,**唯一缺口**是 -base.py:159 的 ``get_fundamentals``(默认抛 NotImplementedError)。本子类填这个缺口。 +BulletTrade 的 MiniQMTProvider 实现了行情/成分/涨跌停/证券信息,**缺口**由本子类填: +- ``get_fundamentals``(base 默认抛 NotImplementedError) +- **本地库委托(2026-08-19 生产事故修复)**: ``get_closes_panel``/``get_constituent_ex`` + 在 base 与 xtdata SDK 都不存在(bullet_trade api 回退链终断 → AttributeError), + ``get_index_stocks`` 只返 miniQMT 最新成份且忽略 date(与回测口径分歧+前视)。 + 修法:内部持有 LocalUnifiedProvider 读本地 dbbardata/constituent_unified, + 与回测同源同口径;表缺指数时 WARNING 回退 base 最新成份(宁可降级不空转)。 数据源映射(miniQMT 实证 2026-07-18,见 docs/portfolio_backtest_result.md): - ``xtdata.get_financial_data(stock_list)`` → dict[stock][table_name] → DataFrame @@ -49,6 +54,8 @@ except ImportError as _e: # Mac dev 环境可能未装,允许模块加载 _HAS_BT_BASE = False _BT_IMPORT_ERROR = _e +from .local_unified_provider import LocalUnifiedProvider + # 聚宽 valuation/indicator/balance 列名 → 我们合并 DataFrame 的列名 # (统一用聚宽列名,方便策略层直接 pandas 筛选) @@ -94,6 +101,66 @@ class SanguoMiniQmtProvider(MiniQMTProvider): # type: ignore[misc] f"{_BT_IMPORT_ERROR}" ) super().__init__(config or {}) + # 本地库委托(2026-08-19):config 的 db_path/data_dir 键透传, + # 缺省走 VPS 生产库默认路径(与回测 LocalUnifiedProvider 同源同口径) + self._unified = LocalUnifiedProvider(config or {}) + + # ------------------------ 本地库委托(panel/成份股,2026-08-19) ------------------------ + def get_closes_panel( + self, + symbols: List[str], + start: Union[str, datetime], + end: Union[str, datetime], + interval: str = "d", + fq: str = "raw", + ) -> pd.DataFrame: + """批量 close 宽表,委托本地 dbbardata(与回测 LocalUnifiedProvider 同口径)。 + + 签名/语义对齐 ``LocalUnifiedProvider.get_closes_panel``(宽表 index=datetime + 升序、columns=输入顺序、缺失股票 NaN 列、strict 契约)。 + """ + return self._unified.get_closes_panel( + symbols, start, end, interval=interval, fq=fq, + ) + + def get_closes_panel_ex( + self, + symbols: List[str], + start: Union[str, datetime], + end: Union[str, datetime], + interval: str = "d", + fq: str = "raw", + ) -> pd.DataFrame: + """TET 别名:与 ``get_closes_panel`` 同一实现(_ex 策略副本调用名)。""" + return self.get_closes_panel(symbols, start, end, interval=interval, fq=fq) + + def get_index_stocks( + self, + index_symbol: str, + date: Optional[Union[str, datetime]] = None, + ) -> List[str]: + """成份股:委托 constituent_unified 并集(与回测同口径,含被踢成份)。 + + base 实现只返 miniQMT 最新权重成份且忽略 date(前视+与回测口径分歧); + 本地表无该指数时 WARNING 回退 base 最新成份——宁可降级不空转 + (2026-08-19 四账户空转两日事故教训)。 + """ + stocks = self._unified.get_constituent_ex(index_symbol, date) + if stocks: + return stocks + logger.warning( + "constituent_unified 无 %s 成份记录(库缺该指数),回退 miniQMT 最新成份" + "(与回测口径分歧,历史 date 忽略)", index_symbol, + ) + return super().get_index_stocks(index_symbol, date=None) + + def get_constituent_ex( + self, + index: str, + date: Optional[Union[str, datetime]] = None, + ) -> List[str]: + """TET 别名:与 ``get_index_stocks`` 同一实现(_ex 策略副本调用名)。""" + return self.get_index_stocks(index, date) # ------------------------ 主入口 ------------------------ def get_fundamentals( @@ -142,6 +209,7 @@ class SanguoMiniQmtProvider(MiniQMTProvider): # type: ignore[misc] self, stocks: List[str], date: Optional[Union[str, datetime]] = None, + fields: Optional[List[str]] = None, ) -> pd.DataFrame: """合并 PershareIndex + Balance + Income + CashFlow + Capital + close。 @@ -151,6 +219,10 @@ class SanguoMiniQmtProvider(MiniQMTProvider): # type: ignore[misc] inc_operation_profit_year_on_year/inc_total_revenue_year_on_year + total_liability/total_sheet_owner_equities/retained_profit + roic(自算)+归母净利润/营收/经营现金流(供策略再算其它因子) + + ``fields`` 契约(2026-08-19 对齐 unified,small_cap 等按需取列): + ``fields=None`` 全列(向后兼容);给定时 ``keep = ["code"] + 请求列``, + 未知列静默丢弃(与 ``LocalUnifiedProvider`` 同口径)。 """ if not stocks: return pd.DataFrame(columns=list(JQ_COLUMN_ALIASES.values())) @@ -187,8 +259,20 @@ class SanguoMiniQmtProvider(MiniQMTProvider): # type: ignore[misc] df = pd.DataFrame(rows) if "code" in df.columns: df = df.set_index("code", drop=False) + if fields: + keep = ["code"] + [f for f in fields if f in df.columns] + df = df[keep] return df + def get_fundamentals_df_ex( + self, + stocks: List[str], + date: Optional[Union[str, datetime]] = None, + fields: Optional[List[str]] = None, + ) -> pd.DataFrame: + """TET 别名:与 ``get_fundamentals_df`` 同一实现(_ex 策略副本调用名)。""" + return self.get_fundamentals_df(stocks, date=date, fields=fields) + def _download_financial_safe( self, qmt_stocks: List[str], tables: List[str], timeout: float = 120.0 ) -> None: diff --git a/tests/portfolio/test_provider.py b/tests/portfolio/test_provider.py index 0867c17..ddc159c 100644 --- a/tests/portfolio/test_provider.py +++ b/tests/portfolio/test_provider.py @@ -141,3 +141,134 @@ class TestSetDataProviderInjection: 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"]