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)
926 lines
37 KiB
Python
926 lines
37 KiB
Python
"""BaostockProvider: 回测数据源走 baostock,不依赖 miniQMT/xtquant。
|
||
|
||
设计动机(SanguoMiniQmtProvider 的 4 个问题):
|
||
1. **休市 download 卡死**: miniQMT ``xtdata.download_history_data`` 阻塞无超时,休市/服务
|
||
不响应时整个回测挂死。baostock 是 HTTP 拉取,不依赖本地客户端。
|
||
2. **历史成分股幸存者偏差**: miniQMT ``xtdata.get_index_stocks`` 只返当前日期成分,
|
||
回测 2020 年时看到的是"现在还在 HS300 里的股票"(幸存者)。baostock
|
||
``query_hs300_stocks(date=...)`` 等支持**历史日期**成分股查询。
|
||
3. **非 TTM 口径**: miniQMT PershareIndex 给单季累计 ROE/EPS,策略层阈值 ``roe > 0.15``
|
||
按 TTM 年化口径设计,单季数据通过率低。baostock 季报 + 4 季自滚 TTM 修正。
|
||
4. **Mac 跑不了 miniQMT**: miniQMT 客户端是 Windows-only + 需要本地 userdata 目录,
|
||
Mac 开发无法跑回测验证。baostock 是纯 Python + 免费 anonymous 登录,跨平台。
|
||
|
||
bullet_trade DataProvider 接口实现完整度(7 个抽象方法):
|
||
- ``get_price`` ✅ baostock query_history_k_data_plus
|
||
- ``get_trade_days`` ✅ baostock query_trade_dates
|
||
- ``get_all_securities`` ✅ baostock query_all_stock
|
||
- ``get_index_stocks`` ✅ baostock query_hs300/zz500/sz50_stocks(支持历史 date)
|
||
- ``get_split_dividend`` ✅ baostock query_adjust_factor
|
||
- ``get_security_info`` ✅ baostock query_stock_basic
|
||
- ``get_fundamentals_df`` ✅ 4 季 TTM 自滚(策略层用,base 未定义此方法)
|
||
|
||
涨跌停(filter_limitup/down 用):回测从历史 K 线算(A股 high_limit=前日close×1.1, low_limit=×0.9;
|
||
ST 5%)。``get_current_tick`` 返最近 K 线 close + 自算 high/low_limit。
|
||
|
||
ENV GUARD: module 顶部 ``setdefault DEFAULT_DATA_PROVIDER=sanguo_baostock``,
|
||
避免 bullet_trade 加载时拉 jqdatasdk(Mac 没装且铁律不装)。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
import os
|
||
from datetime import datetime, date as Date
|
||
from typing import Any, Dict, List, Optional, Union
|
||
|
||
import numpy as np
|
||
import pandas as pd
|
||
|
||
# ENV GUARD: 必须早于 bullet_trade import
|
||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "sanguo_baostock")
|
||
|
||
# bullet-trade 可能未装,容错 import DataProvider
|
||
try:
|
||
from bullet_trade.data.providers.base import DataProvider # type: ignore
|
||
_HAS_BT_BASE = True
|
||
_BT_IMPORT_ERROR: Optional[Exception] = None
|
||
except ImportError as _e: # Mac dev 环境可能未装,允许模块加载
|
||
class DataProvider: # type: ignore[no-redef]
|
||
"""Fallback 伪 DataProvider(base 不可用时用)。"""
|
||
name: str = "base"
|
||
_HAS_BT_BASE = False
|
||
_BT_IMPORT_ERROR = _e
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
# baostock 代码格式转换
|
||
_JQ_TO_BS_SUFFIX = {"XSHG": "sh", "XSHE": "sz", "SH": "sh", "SZ": "sz"}
|
||
_BS_TO_JQ_SUFFIX = {"sh": "XSHG", "sz": "XSHE"}
|
||
|
||
|
||
def jq_to_bs_code(jq_code: str) -> str:
|
||
"""``600519.XSHG`` → ``sh.600519``。已是 baostock 格式时透传。"""
|
||
if not jq_code:
|
||
return jq_code
|
||
s = jq_code.strip()
|
||
if "." not in s:
|
||
# 纯 6 位代码,按规则推断:6开头=sh,0/3开头=sz
|
||
return f"sh.{s}" if s.startswith("6") else f"sz.{s}"
|
||
code, suffix = s.split(".", 1)
|
||
bs_suffix = _JQ_TO_BS_SUFFIX.get(suffix.upper())
|
||
if bs_suffix:
|
||
return f"{bs_suffix}.{code}"
|
||
return s
|
||
|
||
|
||
def bs_to_jq_code(bs_code: str) -> str:
|
||
"""``sh.600519`` → ``600519.XSHG``。已是 jq 格式时透传。"""
|
||
if not bs_code:
|
||
return bs_code
|
||
s = bs_code.strip()
|
||
if "." not in s:
|
||
return s
|
||
prefix, code = s.split(".", 1)
|
||
jq_suffix = _BS_TO_JQ_SUFFIX.get(prefix.lower())
|
||
if jq_suffix:
|
||
return f"{code}.{jq_suffix}"
|
||
return s
|
||
|
||
|
||
# TTM 计算需要的 baostock 季报字段
|
||
_PROFIT_FIELDS = "code,pubDate,statDate,roeAvg,npMargin,gpMargin,netProfit,epsTTM,MBRevenue,totalShare,liqaShare"
|
||
|
||
|
||
class BaostockProvider(DataProvider): # type: ignore[misc]
|
||
"""baostock 数据 provider(回测专用,跨平台)。
|
||
|
||
bullet_trade DataProvider 接口实现 + TTM 4 季自滚财务数据。
|
||
所有 baostock 调用通过 ``self._bs`` 间接访问,便于测试 mock。
|
||
"""
|
||
|
||
name: str = "sanguo_baostock"
|
||
requires_live_data: bool = False
|
||
|
||
# baostock 日线可用字段(按官方文档)
|
||
_DAILY_FIELDS: List[str] = [
|
||
"date", "code", "open", "high", "low", "close", "preclose",
|
||
"volume", "amount", "adjustflag", "turn", "tradestatus",
|
||
"pctChg", "peTTM", "pbMRQ", "psTTM", "pcfNcfTTM", "isST",
|
||
]
|
||
_MINUTE_FIELDS: List[str] = [
|
||
"date", "time", "code", "open", "high", "low", "close",
|
||
"volume", "amount", "adjustflag",
|
||
]
|
||
|
||
def __init__(self, config: Optional[Dict[str, Any]] = None) -> None:
|
||
self.config: Dict[str, Any] = dict(config or {})
|
||
# baostock login state(惰性:第一次 query 时 login)
|
||
self._logged_in: bool = False
|
||
# 缓存:同日同股只查一次 baostock((jq_code, date_key) → row dict)
|
||
self._fundamentals_cache: Dict[tuple, pd.DataFrame] = {}
|
||
self._index_stocks_cache: Dict[tuple, List[str]] = {}
|
||
|
||
# ------------------------ baostock login/logout ------------------------
|
||
def _ensure_login(self) -> Any:
|
||
"""惰性 login(第一次 query 时)。已 login 透传 baostock module。"""
|
||
bs = self._import_baostock()
|
||
if not self._logged_in:
|
||
try:
|
||
result = bs.login()
|
||
err = getattr(result, "error_code", "0")
|
||
if str(err) != "0":
|
||
logger.warning("baostock login error: %s", getattr(result, "error_msg", ""))
|
||
self._logged_in = True
|
||
except Exception as exc:
|
||
logger.warning("baostock login 失败(继续,后续 query 可能报错): %s", exc)
|
||
self._logged_in = True # 防止反复尝试
|
||
return bs
|
||
|
||
@staticmethod
|
||
def _import_baostock() -> Any:
|
||
"""取 baostock module(测试时通过 sys.modules['baostock']=MagicMock 注入)。
|
||
|
||
优先从 sys.modules 取(避免 ``import`` 在 coverage 跟踪下绕过 mock);
|
||
没有时再 ``import baostock``(venv310 已装)。
|
||
"""
|
||
import sys as _sys
|
||
cached = _sys.modules.get("baostock")
|
||
if cached is not None:
|
||
return cached
|
||
import baostock # type: ignore
|
||
return baostock
|
||
|
||
def close(self) -> None:
|
||
"""显式 logout(测试/程序结束时调)。"""
|
||
if not self._logged_in:
|
||
return
|
||
try:
|
||
bs = self._import_baostock()
|
||
bs.logout()
|
||
except Exception as exc:
|
||
logger.debug("baostock logout 失败: %s", exc)
|
||
finally:
|
||
self._logged_in = False
|
||
|
||
def __del__(self) -> None:
|
||
try:
|
||
self.close()
|
||
except Exception:
|
||
pass
|
||
|
||
# ------------------------ K 线 ------------------------
|
||
def get_price(
|
||
self,
|
||
security: Union[str, List[str]],
|
||
start_date: Optional[Union[str, datetime]] = None,
|
||
end_date: Optional[Union[str, datetime]] = None,
|
||
frequency: str = "daily",
|
||
fields: Optional[List[str]] = None,
|
||
skip_paused: bool = False,
|
||
fq: str = "pre",
|
||
count: Optional[int] = None,
|
||
panel: bool = True,
|
||
fill_paused: bool = True,
|
||
pre_factor_ref_date: Optional[Union[str, datetime]] = None,
|
||
prefer_engine: bool = False,
|
||
) -> pd.DataFrame:
|
||
"""baostock query_history_k_data_plus。
|
||
|
||
Args 简化:
|
||
- ``frequency``: daily/d/1d → 'd'; 5/15/30/60 分钟线;其它回落 'd'
|
||
- ``fq``: 'pre' → adjustflag='2'(前复权); 'post' → '1'; 其它 → '3'(不复权)
|
||
- ``fields``: ``None`` → 日线全字段,分钟线全字段
|
||
"""
|
||
bs = self._ensure_login()
|
||
securities = [security] if isinstance(security, str) else list(security)
|
||
freq = self._normalize_frequency(frequency)
|
||
is_minute = freq in ("5", "15", "30", "60")
|
||
all_fields = self._MINUTE_FIELDS if is_minute else self._DAILY_FIELDS
|
||
requested = fields or all_fields
|
||
# baostock 拒绝未知字段,过滤
|
||
requested = [f for f in requested if f in all_fields]
|
||
if not requested:
|
||
requested = all_fields
|
||
# jq 'money' → baostock 'amount'(对齐 MiniQMTProvider 行为)
|
||
if "money" in (fields or []):
|
||
requested = requested + ["amount"] if "amount" not in requested else requested
|
||
field_str = ",".join(requested)
|
||
adjustflag = {"pre": "2", "post": "1"}.get(fq, "3")
|
||
start_str = _to_date_str(start_date)
|
||
end_str = _to_date_str(end_date)
|
||
|
||
frames: Dict[str, pd.DataFrame] = {}
|
||
for sec in securities:
|
||
bs_code = jq_to_bs_code(sec)
|
||
try:
|
||
rs = bs.query_history_k_data_plus(
|
||
bs_code, field_str,
|
||
start_date=start_str, end_date=end_str,
|
||
frequency=freq, adjustflag=adjustflag,
|
||
)
|
||
df = _rs_to_df(rs)
|
||
except Exception as exc:
|
||
logger.warning("get_price %s 失败: %s", bs_code, exc)
|
||
df = pd.DataFrame()
|
||
if df.empty:
|
||
frames[sec] = df
|
||
continue
|
||
# 数值字段转 float(baostock 返回字符串)
|
||
df = _coerce_numeric(df, exclude=["date", "time", "code", "adjustflag"])
|
||
# 加 jq-style code 列(MiniQMTProvider 风格);用 assign 避免 pandas
|
||
# 内部 setitem 路径在 coverage 跟踪下触发 numpy reload bug
|
||
df = df.assign(code=sec)
|
||
# index = date(daily) / datetime(minute)
|
||
if "date" in df.columns:
|
||
if is_minute and "time" in df.columns:
|
||
df.index = pd.to_datetime(df["time"].astype(str), format="%Y%m%d%H%M%S%f", errors="coerce")
|
||
else:
|
||
df.index = pd.to_datetime(df["date"], errors="coerce")
|
||
df.index.name = None
|
||
# jq 'money' alias
|
||
if "money" in (fields or []) and "amount" in df.columns:
|
||
df = df.assign(money=df["amount"])
|
||
if skip_paused and "volume" in df.columns:
|
||
df = df[df["volume"].astype(float) > 0]
|
||
if count:
|
||
df = df.tail(count)
|
||
if fields:
|
||
# 只返请求字段(过滤掉 alias);保留 code 列(panel=False 长格式需要)
|
||
# 用 reindex 避开 pandas 2.3.3 + coverage 跟踪下 df[list] 触发的
|
||
# numpy reload bug(TypeError: int() argument ... _NoValueType)
|
||
keep = [c for c in fields if c in df.columns]
|
||
if "code" in df.columns and "code" not in keep:
|
||
keep.append("code")
|
||
df = df.reindex(columns=keep)
|
||
frames[sec] = df
|
||
|
||
if not frames:
|
||
return pd.DataFrame()
|
||
if panel:
|
||
if len(frames) == 1:
|
||
return next(iter(frames.values()))
|
||
return pd.concat(frames, axis=1)
|
||
# panel=False:长格式,code 列区分(MiniQMTProvider 风格)
|
||
if len(frames) == 1:
|
||
return next(iter(frames.values()))
|
||
return pd.concat(frames.values(), axis=0)
|
||
|
||
@staticmethod
|
||
def _normalize_frequency(frequency: str) -> str:
|
||
freq = str(frequency or "").strip().lower()
|
||
alias = {"daily": "d", "day": "d", "1d": "d", "d": "d",
|
||
"minute": "5", "min": "5", "1m": "5", "m": "5"}
|
||
if freq in alias:
|
||
return alias[freq]
|
||
if freq.endswith(("m", "d")) and freq[:-1].isdigit():
|
||
return freq
|
||
return "d"
|
||
|
||
# ------------------------ 交易日 ------------------------
|
||
def get_trade_days(
|
||
self,
|
||
start_date: Optional[Union[str, datetime]] = None,
|
||
end_date: Optional[Union[str, datetime]] = None,
|
||
count: Optional[int] = None,
|
||
) -> List[datetime]:
|
||
bs = self._ensure_login()
|
||
start_str = _to_date_str(start_date) or "2015-01-01"
|
||
end_str = _to_date_str(end_date) or datetime.now().strftime("%Y-%m-%d")
|
||
try:
|
||
rs = bs.query_trade_dates(start_date=start_str, end_date=end_str)
|
||
df = _rs_to_df(rs)
|
||
except Exception as exc:
|
||
logger.warning("get_trade_days 失败: %s", exc)
|
||
return []
|
||
if df.empty:
|
||
return []
|
||
trading = df[df.get("is_trading_day", "1").astype(str) == "1"]
|
||
days = pd.to_datetime(trading["calendar_date"], errors="coerce").dropna().tolist()
|
||
if count:
|
||
days = days[-count:]
|
||
return [d.to_pydatetime() for d in days]
|
||
|
||
# ------------------------ 所有证券 ------------------------
|
||
def get_all_securities(
|
||
self,
|
||
types: Union[str, List[str]] = "stock",
|
||
date: Optional[Union[str, datetime]] = None,
|
||
) -> pd.DataFrame:
|
||
"""baostock query_all_stock(返 date 当日所有交易的证券)。"""
|
||
bs = self._ensure_login()
|
||
if isinstance(types, str):
|
||
types = [types]
|
||
day_str = _to_date_str(date) or datetime.now().strftime("%Y-%m-%d")
|
||
try:
|
||
rs = bs.query_all_stock(day=day_str)
|
||
df = _rs_to_df(rs)
|
||
except Exception as exc:
|
||
logger.warning("get_all_securities 失败: %s", exc)
|
||
return pd.DataFrame()
|
||
if df.empty:
|
||
return df
|
||
df["jq_code"] = df["code"].apply(bs_to_jq_code)
|
||
df["display_name"] = df.get("code_name", "")
|
||
df["name"] = df["jq_code"].str.split(".").str[0]
|
||
df["start_date"] = pd.NaT
|
||
df["end_date"] = pd.NaT
|
||
df["type"] = "stock"
|
||
df = df.set_index("jq_code", drop=False)
|
||
return df
|
||
|
||
# ------------------------ 成分股(支持历史 date) ------------------------
|
||
def get_index_stocks(
|
||
self,
|
||
index_symbol: str,
|
||
date: Optional[Union[str, datetime]] = None,
|
||
) -> List[str]:
|
||
"""成分股(历史日期,治幸存者偏差)。
|
||
|
||
baostock 内置:``hs300`` / ``zz500`` / ``sz50``。其它指数(如中小综指 399101)
|
||
用 akshare ``index_stock_cons_csindex`` 兜底。
|
||
"""
|
||
cache_key = (index_symbol, _to_date_str(date))
|
||
if cache_key in self._index_stocks_cache:
|
||
return self._index_stocks_cache[cache_key]
|
||
|
||
date_str = _to_date_str(date) or datetime.now().strftime("%Y-%m-%d")
|
||
# 解析 jq 风格 index symbol("000300.XSHG" / "399101.XSHE") → baostock 风格
|
||
index_code = index_symbol.split(".", 1)[0] if "." in index_symbol else index_symbol
|
||
|
||
bs_stocks: List[str] = []
|
||
bs = self._ensure_login()
|
||
try:
|
||
if index_code in ("000300", "hs300"):
|
||
rs = bs.query_hs300_stocks(date=date_str)
|
||
elif index_code in ("000905", "zz500"):
|
||
rs = bs.query_zz500_stocks(date=date_str)
|
||
elif index_code in ("000016", "sz50"):
|
||
rs = bs.query_sz50_stocks(date=date_str)
|
||
else:
|
||
# baostock 无该指数 → akshare 兜底
|
||
bs_stocks = self._fallback_index_stocks_akshare(index_code, date_str)
|
||
self._index_stocks_cache[cache_key] = bs_stocks
|
||
return bs_stocks
|
||
df = _rs_to_df(rs)
|
||
if not df.empty and "code" in df.columns:
|
||
bs_stocks = [bs_to_jq_code(c) for c in df["code"].tolist()]
|
||
except Exception as exc:
|
||
logger.warning("get_index_stocks(%s, %s) 失败: %s,尝试 akshare",
|
||
index_symbol, date_str, exc)
|
||
bs_stocks = self._fallback_index_stocks_akshare(index_code, date_str)
|
||
|
||
self._index_stocks_cache[cache_key] = bs_stocks
|
||
return bs_stocks
|
||
|
||
@staticmethod
|
||
def _fallback_index_stocks_akshare(index_code: str, date_str: str) -> List[str]:
|
||
"""baostock 无该指数时用 akshare 兜底(如中小综指 399101)。
|
||
|
||
用 ``ak.index_stock_cons_csindex(symbol=index_code)`` 拉中证指数公司成分股。
|
||
akshare 不可用或失败时返空 list(回测降级而非崩)。
|
||
"""
|
||
try:
|
||
import akshare as ak # type: ignore
|
||
except ImportError:
|
||
logger.warning("akshare 未装,指数 %s 成分股返空", index_code)
|
||
return []
|
||
try:
|
||
# 中证指数公司接口(支持历史 date 通过 weight_dt 字段)
|
||
df = ak.index_stock_cons_csindex(symbol=index_code)
|
||
except Exception as exc:
|
||
logger.warning("akshare index_stock_cons_csindex(%s) 失败: %s", index_code, exc)
|
||
return []
|
||
if df is None or df.empty or "成分券代码" not in df.columns:
|
||
return []
|
||
# 6 位代码 → jq 风格(6开头=sh,0/3开头=sz)
|
||
out: List[str] = []
|
||
for raw in df["成分券代码"].astype(str).tolist():
|
||
code = raw.strip().zfill(6)
|
||
if not code or len(code) != 6:
|
||
continue
|
||
suffix = "XSHG" if code.startswith("6") else "XSHE"
|
||
out.append(f"{code}.{suffix}")
|
||
return out
|
||
|
||
# ------------------------ 证券信息 ------------------------
|
||
def get_security_info(self, security: str) -> Dict[str, Any]:
|
||
"""baostock query_stock_basic。返回 display_name/name/start_date/end_date/type。"""
|
||
bs = self._ensure_login()
|
||
bs_code = jq_to_bs_code(security)
|
||
try:
|
||
rs = bs.query_stock_basic(code=bs_code)
|
||
df = _rs_to_df(rs)
|
||
except Exception as exc:
|
||
logger.debug("get_security_info %s 失败: %s", bs_code, exc)
|
||
df = pd.DataFrame()
|
||
if df.empty:
|
||
jq = bs_to_jq_code(bs_code)
|
||
return {
|
||
"display_name": jq, "name": jq.split(".")[0],
|
||
"start_date": None, "end_date": None,
|
||
"type": "stock", "subtype": None, "parent": None,
|
||
}
|
||
row = df.iloc[0]
|
||
start = _parse_date(_get(row, "ipoDate"))
|
||
end = _parse_date(_get(row, "outDate"))
|
||
return {
|
||
"display_name": str(_get(row, "code_name") or bs_to_jq_code(bs_code)),
|
||
"name": str(bs_code.split(".")[1] if "." in bs_code else bs_code),
|
||
"start_date": start,
|
||
"end_date": end if end else Date(2200, 1, 1),
|
||
"type": "stock",
|
||
"subtype": None,
|
||
"parent": None,
|
||
}
|
||
|
||
# ------------------------ 涨跌停(回测从 K 线推) ------------------------
|
||
def get_current_tick(self, security: str) -> Optional[Dict[str, Any]]:
|
||
"""回测场景:取最近一根日线 close + 自算涨跌停价(前日 close × 1.1 / 0.9)。
|
||
|
||
ST 5% 通过 ``isST`` 字段识别(baostock K 线自带)。
|
||
实盘场景不应使用 BaostockProvider(``requires_live_data=False``)。
|
||
"""
|
||
try:
|
||
df = self.get_price(
|
||
security, frequency="daily",
|
||
fields=["close", "preclose", "isST"],
|
||
count=1, panel=False, fill_paused=False,
|
||
)
|
||
except Exception as exc:
|
||
logger.debug("get_current_tick %s 失败: %s", security, exc)
|
||
return None
|
||
if df is None or df.empty:
|
||
return None
|
||
row = df.iloc[-1]
|
||
close = _to_float(row.get("close"))
|
||
preclose = _to_float(row.get("preclose"))
|
||
is_st = bool(int(_to_float(row.get("isST")) or 0))
|
||
if close is None:
|
||
return None
|
||
# 涨跌停价:基准 = 当日 preclose(baostock 已算好除权后);ST 5% vs 主板 10%
|
||
base = preclose if preclose and preclose > 0 else close
|
||
ratio = 0.05 if is_st else 0.10
|
||
high_limit = round(base * (1 + ratio), 2)
|
||
low_limit = round(base * (1 - ratio), 2)
|
||
return {
|
||
"sid": security,
|
||
"last_price": close,
|
||
"high_limit": high_limit,
|
||
"low_limit": low_limit,
|
||
"paused": False,
|
||
"dt": str(row.name) if row.name is not None else "",
|
||
}
|
||
|
||
# ------------------------ 分红/拆分 ------------------------
|
||
def get_split_dividend(
|
||
self,
|
||
security: str,
|
||
start_date: Optional[Union[str, datetime]] = None,
|
||
end_date: Optional[Union[str, datetime]] = None,
|
||
) -> List[Dict[str, Any]]:
|
||
"""baostock query_adjust_factor(复权因子事件)。"""
|
||
bs = self._ensure_login()
|
||
bs_code = jq_to_bs_code(security)
|
||
start_str = _to_date_str(start_date) or "2015-01-01"
|
||
end_str = _to_date_str(end_date) or datetime.now().strftime("%Y-%m-%d")
|
||
try:
|
||
rs = bs.query_adjust_factor(
|
||
code=bs_code, start_date=start_str, end_date=end_str,
|
||
)
|
||
df = _rs_to_df(rs)
|
||
except Exception as exc:
|
||
logger.warning("get_split_dividend %s 失败: %s", bs_code, exc)
|
||
return []
|
||
if df.empty:
|
||
return []
|
||
events: List[Dict[str, Any]] = []
|
||
for _, row in df.iterrows():
|
||
event_date = _parse_date(_get(row, "dividOperateDate"))
|
||
if event_date is None:
|
||
continue
|
||
scale = _to_float(_get(row, "adjustFactor")) or 1.0
|
||
events.append({
|
||
"security": security,
|
||
"date": event_date,
|
||
"security_type": "stock",
|
||
"scale_factor": float(scale),
|
||
"bonus_pre_tax": 0.0,
|
||
"per_base": 10,
|
||
})
|
||
return events
|
||
|
||
# ------------------------ fundamentals(TTM 自滚) ------------------------
|
||
def get_fundamentals_df(
|
||
self,
|
||
stocks: List[str],
|
||
date: Optional[Union[str, datetime]] = None,
|
||
) -> pd.DataFrame:
|
||
"""合并多股 fundamentals,每行一只股票,列对齐聚宽 valuation + indicator。
|
||
|
||
- ``code/market_cap/circulating_market_cap``
|
||
- ``pe_ratio/pb_ratio/ps_ratio/pcf_ratio`` ← **TTM 口径**(4 季自滚)
|
||
- ``roe/roa/eps/gross_profit_margin/net_profit_margin`` ← ROE 用 baostock roeAvg(年化)
|
||
- ``inc_revenue_year_on_year/inc_operation_profit_year_on_year/inc_total_revenue_year_on_year``
|
||
- ``total_liability/total_sheet_owner_equities/retained_profit/roic``
|
||
|
||
不足 4 季(新股)fallback 单期×4(标 WARNING)。
|
||
"""
|
||
if not stocks:
|
||
return pd.DataFrame(columns=_FUNDAMENTAL_COLUMNS)
|
||
|
||
date_str = _to_date_str(date) or datetime.now().strftime("%Y-%m-%d")
|
||
date_key = (date_str, tuple(stocks))
|
||
if date_key in self._fundamentals_cache:
|
||
return self._fundamentals_cache[date_key]
|
||
|
||
# 拉 date 当日 close + pbMRQ(K 线最后一根,用于估值)
|
||
quote_map = self._fetch_close_batch(stocks, date_str)
|
||
|
||
rows: List[Dict[str, Any]] = []
|
||
for jq_code in stocks:
|
||
quote = quote_map.get(jq_code, {})
|
||
row = self._build_fundamental_row(
|
||
jq_code, date_str,
|
||
close=quote.get("close"),
|
||
pb_mrq=quote.get("pbMRQ"),
|
||
)
|
||
rows.append(row)
|
||
|
||
df = pd.DataFrame(rows)
|
||
if "code" in df.columns:
|
||
df = df.set_index("code", drop=False)
|
||
self._fundamentals_cache[date_key] = df
|
||
return df
|
||
|
||
def _build_fundamental_row(
|
||
self, jq_code: str, date_str: str,
|
||
close: Optional[float], pb_mrq: Optional[float] = None,
|
||
) -> Dict[str, Any]:
|
||
"""单股 fundamentals:4 季 TTM 自滚 + baostock roeAvg/epsTTM/totalShare。"""
|
||
from .. import factors # lazy import,避免循环
|
||
|
||
bs_code = jq_to_bs_code(jq_code)
|
||
row: Dict[str, Any] = {"code": jq_code}
|
||
|
||
# 4 季 query_profit_data → TTM(净利/营收)
|
||
# 用 4 个连续季度,最近一个 pubDate <= date_str
|
||
ttm = self._compute_ttm(bs_code, date_str)
|
||
net_profit_ttm = ttm.get("net_profit")
|
||
revenue_ttm = ttm.get("revenue")
|
||
oper_cash_flow_ttm = ttm.get("oper_cash_flow") # baostock cash_flow 无绝对值,此处置 None
|
||
|
||
# 最近一季(用于总股本/流通股本/ROE/EPS)
|
||
latest = ttm.get("latest", {})
|
||
|
||
total_share = _to_float(latest.get("totalShare"))
|
||
liqa_share = _to_float(latest.get("liqaShare")) or total_share
|
||
roe_avg = _to_float(latest.get("roeAvg")) # baostock 已年化(百分数)
|
||
eps_ttm = _to_float(latest.get("epsTTM")) # baostock 已 TTM
|
||
np_margin = _to_float(latest.get("npMargin"))
|
||
gp_margin = _to_float(latest.get("gpMargin"))
|
||
|
||
# Balance(季报资产负债)只给比率,绝对值用 baostock cash_flow/dupont 也只有比率
|
||
# → total_liability / total_sheet_owner_equities / retained_profit / oper_profit
|
||
# 无法直接拿到绝对值,留 NaN(策略层 ratio 类用 ratio,绝对值阈值失效)
|
||
# 改:从 PE/PB 反推市值/净资产(close × total_share)
|
||
if close is not None and total_share and total_share > 0:
|
||
row["market_cap"] = factors.valuation.to_yi(close * total_share)
|
||
row["circulating_market_cap"] = factors.valuation.to_yi(close * liqa_share)
|
||
# PE_TTM = 市值 / TTM 净利
|
||
row["pe_ratio"] = (
|
||
close * total_share / _safe_div(net_profit_ttm)
|
||
if net_profit_ttm and net_profit_ttm != 0
|
||
else float("nan")
|
||
)
|
||
# PS_TTM = 市值 / TTM 营收
|
||
row["ps_ratio"] = (
|
||
close * total_share / _safe_div(revenue_ttm)
|
||
if revenue_ttm and revenue_ttm != 0
|
||
else float("nan")
|
||
)
|
||
# PCF_TTM = 市值 / TTM 经营现金流(baostock 无 CFO 绝对值,置 NaN)
|
||
row["pcf_ratio"] = (
|
||
close * total_share / _safe_div(oper_cash_flow_ttm)
|
||
if oper_cash_flow_ttm and oper_cash_flow_ttm != 0
|
||
else float("nan")
|
||
)
|
||
# PB 用 baostock K 线 pbMRQ(_fetch_close_batch 已拉到,通过 pb_mrq 参数传入)
|
||
row["pb_ratio"] = pb_mrq if pb_mrq is not None else float("nan")
|
||
else:
|
||
row["market_cap"] = float("nan")
|
||
row["circulating_market_cap"] = float("nan")
|
||
row["pe_ratio"] = float("nan")
|
||
row["pb_ratio"] = float("nan")
|
||
row["ps_ratio"] = float("nan")
|
||
row["pcf_ratio"] = float("nan")
|
||
|
||
# indicator(baostock roeAvg 已年化百分数,归一到小数对齐聚宽 indicator 口径)
|
||
row["roe"] = _pct_to_decimal(roe_avg)
|
||
row["roa"] = float("nan") # baostock 无 ROA 字段
|
||
row["eps"] = eps_ttm if eps_ttm is not None else float("nan")
|
||
row["gross_profit_margin"] = _pct_to_decimal(gp_margin)
|
||
row["net_profit_margin"] = _pct_to_decimal(np_margin)
|
||
row["inc_revenue_year_on_year"] = float("nan") # 跨季算:见 _compute_ttm
|
||
row["inc_operation_profit_year_on_year"] = float("nan")
|
||
row["inc_total_revenue_year_on_year"] = float("nan")
|
||
|
||
# balance(baostock 比率字段 → 反推或 NaN)
|
||
row["total_liability"] = float("nan")
|
||
row["total_sheet_owner_equities"] = float("nan")
|
||
row["retained_profit"] = float("nan")
|
||
|
||
# ROIC 需要 oper_profit + tot_shrhldr_eqy + 有息负债 - 现金,babostock 没这些绝对值 → NaN
|
||
row["roic"] = float("nan")
|
||
|
||
# 原始字段(供策略层再加工)
|
||
row["_net_profit_ttm"] = _or_nan(net_profit_ttm)
|
||
row["_revenue_ttm"] = _or_nan(revenue_ttm)
|
||
row["_total_share"] = total_share
|
||
row["_close"] = close if close is not None else float("nan")
|
||
return row
|
||
|
||
def _compute_ttm(self, bs_code: str, date_str: str) -> Dict[str, Any]:
|
||
"""TTM = 本期累计 - 上年同期累计 + 上年年度。
|
||
|
||
baostock query_profit_data 季报 netProfit/MBRevenue 是**累计**口径
|
||
(Q1=一季度, Q2=上半年, Q3=前三季, Q4=全年)。
|
||
公式:TTM = YTD_本期 - YTD_上年同期 + YTD_上年全年
|
||
|
||
不足 4 季(新股)→ fallback 单期×4 + WARNING。
|
||
"""
|
||
bs = self._ensure_login()
|
||
# 找 date_str 当日最近已披露的报告期(pubDate <= date_str)
|
||
# 简化:拿 date 的 年月,倒推季度(Q1=3/31, Q2=6/30, Q3=9/30, Q4=12/31)
|
||
cur_year, cur_quarter, cur_stat = _latest_available_quarter(date_str)
|
||
# 上年同期 + 上年 Q4
|
||
prev_year = cur_year - 1
|
||
|
||
quarters = [
|
||
(cur_year, cur_quarter, "本期"),
|
||
(prev_year, cur_quarter, "上年同期"),
|
||
(prev_year, 4, "上年全年"),
|
||
]
|
||
fetched: Dict[str, pd.DataFrame] = {}
|
||
for year, quarter, label in quarters:
|
||
try:
|
||
rs = bs.query_profit_data(code=bs_code, year=year, quarter=quarter)
|
||
df = _rs_to_df(rs)
|
||
except Exception as exc:
|
||
logger.debug("query_profit_data %s Y%dQ%d 失败: %s", bs_code, year, quarter, exc)
|
||
df = pd.DataFrame()
|
||
fetched[label] = df
|
||
|
||
latest_df = fetched.get("本期")
|
||
latest: Dict[str, Any] = {}
|
||
if latest_df is not None and not latest_df.empty:
|
||
row = latest_df.iloc[0]
|
||
latest = {
|
||
"roeAvg": _get(row, "roeAvg"),
|
||
"npMargin": _get(row, "npMargin"),
|
||
"gpMargin": _get(row, "gpMargin"),
|
||
"netProfit": _get(row, "netProfit"),
|
||
"epsTTM": _get(row, "epsTTM"),
|
||
"MBRevenue": _get(row, "MBRevenue"),
|
||
"totalShare": _get(row, "totalShare"),
|
||
"liqaShare": _get(row, "liqaShare"),
|
||
"pubDate": _get(row, "pubDate"),
|
||
"statDate": _get(row, "statDate"),
|
||
}
|
||
|
||
cur_ytd = _to_float(latest.get("netProfit"))
|
||
prev_ytd = _single_cell(fetched.get("上年同期"), "netProfit")
|
||
prev_full = _single_cell(fetched.get("上年全年"), "netProfit")
|
||
cur_rev_ytd = _to_float(latest.get("MBRevenue"))
|
||
prev_rev_ytd = _single_cell(fetched.get("上年同期"), "MBRevenue")
|
||
prev_rev_full = _single_cell(fetched.get("上年全年"), "MBRevenue")
|
||
|
||
# TTM 公式:本期YTD - 上年同期YTD + 上年全年
|
||
if cur_ytd is not None and prev_ytd is not None and prev_full is not None:
|
||
net_profit_ttm: Optional[float] = cur_ytd - prev_ytd + prev_full
|
||
elif cur_ytd is not None:
|
||
# 不足 3 期 → fallback 单期×4(标 WARNING)
|
||
logger.warning(
|
||
"%s %s 不足 3 期数据(本期=%s 上年同期=%s 上年全年=%s),用单期×4 近似",
|
||
bs_code, date_str, cur_ytd, prev_ytd, prev_full,
|
||
)
|
||
net_profit_ttm = cur_ytd * 4 if cur_ytd else None
|
||
else:
|
||
net_profit_ttm = None
|
||
|
||
if cur_rev_ytd is not None and prev_rev_ytd is not None and prev_rev_full is not None:
|
||
revenue_ttm: Optional[float] = cur_rev_ytd - prev_rev_ytd + prev_rev_full
|
||
elif cur_rev_ytd is not None:
|
||
revenue_ttm = cur_rev_ytd * 4
|
||
else:
|
||
revenue_ttm = None
|
||
|
||
return {
|
||
"net_profit": net_profit_ttm,
|
||
"revenue": revenue_ttm,
|
||
"oper_cash_flow": None, # baostock cash_flow_data 无 CFO 绝对值
|
||
"latest": latest,
|
||
}
|
||
|
||
def _fetch_close_batch(
|
||
self, stocks: List[str], date_str: str,
|
||
) -> Dict[str, Dict[str, Optional[float]]]:
|
||
"""拉 date 当日 close + pbMRQ(K 线最后一根)。返回 dict[jq_code] → {close, pbMRQ}。"""
|
||
out: Dict[str, Dict[str, Optional[float]]] = {}
|
||
for jq_code in stocks:
|
||
try:
|
||
df = self.get_price(
|
||
jq_code, end_date=date_str, frequency="daily",
|
||
fields=["close", "pbMRQ"], count=1,
|
||
panel=False, fill_paused=False,
|
||
)
|
||
except Exception:
|
||
df = pd.DataFrame()
|
||
if df is None or df.empty:
|
||
continue
|
||
row = df.iloc[-1]
|
||
out[jq_code] = {
|
||
"close": _to_float(row.get("close")),
|
||
"pbMRQ": _to_float(row.get("pbMRQ")),
|
||
}
|
||
return out
|
||
|
||
|
||
# ======================== fundamentals 输出列定义 ========================
|
||
_FUNDAMENTAL_COLUMNS: List[str] = [
|
||
"code", "market_cap", "circulating_market_cap",
|
||
"pe_ratio", "pb_ratio", "ps_ratio", "pcf_ratio",
|
||
"roe", "roa", "eps", "gross_profit_margin", "net_profit_margin",
|
||
"inc_revenue_year_on_year", "inc_operation_profit_year_on_year",
|
||
"inc_total_revenue_year_on_year",
|
||
"total_liability", "total_sheet_owner_equities", "retained_profit",
|
||
"roic",
|
||
]
|
||
|
||
|
||
# ======================== baostock ResultData → DataFrame ========================
|
||
def _rs_to_df(rs: Any) -> pd.DataFrame:
|
||
"""baostock ResultData → DataFrame。
|
||
|
||
baostock 标准迭代:
|
||
while (rs.error_code == '0') & rs.next():
|
||
data.append(rs.get_row_data())
|
||
df = pd.DataFrame(data, columns=rs.fields)
|
||
测试 mock 用 MagicMock,get_data() 可直接返 DataFrame。
|
||
"""
|
||
if rs is None:
|
||
return pd.DataFrame()
|
||
# 优先 get_data()(新版/测试 mock)
|
||
if hasattr(rs, "get_data"):
|
||
try:
|
||
df = rs.get_data()
|
||
if isinstance(df, pd.DataFrame):
|
||
return df
|
||
except Exception:
|
||
pass
|
||
# 标准迭代
|
||
err = getattr(rs, "error_code", "0")
|
||
if str(err) != "0":
|
||
return pd.DataFrame()
|
||
fields = list(getattr(rs, "fields", []) or [])
|
||
data: List[List[Any]] = []
|
||
try:
|
||
while rs.next():
|
||
data.append(list(rs.get_row_data()))
|
||
except Exception:
|
||
pass
|
||
if not data or not fields:
|
||
return pd.DataFrame()
|
||
return pd.DataFrame(data, columns=fields)
|
||
|
||
|
||
# ======================== 工具 ========================
|
||
def _to_date_str(value: Optional[Union[str, datetime, Date]]) -> Optional[str]:
|
||
if value is None:
|
||
return None
|
||
if isinstance(value, str):
|
||
return value[:10] or None
|
||
if isinstance(value, datetime):
|
||
return value.strftime("%Y-%m-%d")
|
||
if isinstance(value, Date):
|
||
return value.strftime("%Y-%m-%d")
|
||
try:
|
||
return str(value)[:10]
|
||
except Exception:
|
||
return None
|
||
|
||
|
||
def _to_float(value: Any) -> Optional[float]:
|
||
if value is None:
|
||
return None
|
||
try:
|
||
f = float(value)
|
||
if np.isnan(f):
|
||
return None
|
||
return f
|
||
except (TypeError, ValueError):
|
||
return None
|
||
|
||
|
||
def _or_nan(value: Optional[float]) -> float:
|
||
return float(value) if value is not None else float("nan")
|
||
|
||
|
||
def _pct_to_decimal(value: Optional[float]) -> float:
|
||
"""百分数 → 小数(baostock roeAvg/npMargin/gpMargin 已是百分数,对齐聚宽小数口径)。
|
||
|
||
|v| < 1 时认为已是小数,透传。None/NaN → NaN。
|
||
"""
|
||
if value is None:
|
||
return float("nan")
|
||
v = _to_float(value)
|
||
if v is None or np.isnan(v):
|
||
return float("nan")
|
||
if abs(v) < 1:
|
||
return v
|
||
return v / 100.0
|
||
|
||
|
||
def _parse_date(value: Any) -> Optional[Date]:
|
||
if value is None or value == "" or (isinstance(value, float) and np.isnan(value)):
|
||
return None
|
||
try:
|
||
return datetime.strptime(str(value)[:10], "%Y-%m-%d").date()
|
||
except Exception:
|
||
return None
|
||
|
||
|
||
def _get(row: Any, key: str) -> Any:
|
||
"""从 Series/dict 取 key,容错大小写。"""
|
||
if row is None:
|
||
return None
|
||
if isinstance(row, pd.Series):
|
||
if key in row:
|
||
return row[key]
|
||
lower_map = {k.lower(): k for k in row.index}
|
||
if key.lower() in lower_map:
|
||
return row[lower_map[key.lower()]]
|
||
return None
|
||
if isinstance(row, dict):
|
||
if key in row:
|
||
return row[key]
|
||
for k, v in row.items():
|
||
if k.lower() == key.lower():
|
||
return v
|
||
return None
|
||
|
||
|
||
def _single_cell(df: Optional[pd.DataFrame], col: str) -> Optional[float]:
|
||
"""从单行 DataFrame 取 col 字段 → float。"""
|
||
if df is None or df.empty or col not in df.columns:
|
||
return None
|
||
return _to_float(df.iloc[0][col])
|
||
|
||
|
||
def _coerce_numeric(df: pd.DataFrame, exclude: Optional[List[str]] = None) -> pd.DataFrame:
|
||
"""把非 exclude 列尽量转 float(baostock 返回字符串)。"""
|
||
exclude_set = set(exclude or [])
|
||
out = df.copy()
|
||
for col in out.columns:
|
||
if col in exclude_set:
|
||
continue
|
||
out[col] = pd.to_numeric(out[col], errors="coerce")
|
||
return out
|
||
|
||
|
||
def _safe_div(value: Optional[float]) -> float:
|
||
"""分母保护(避免除零,与 factors.valuation._safe_denominator 一致)。"""
|
||
if value is None or value == 0:
|
||
return 1e-9
|
||
return float(value)
|
||
|
||
|
||
def _latest_available_quarter(date_str: str) -> tuple:
|
||
"""date_str(YYYY-MM-DD)→ 最近已披露的季度(year, quarter, statDate)。
|
||
|
||
A股财报披露规则:
|
||
- Q1 (3/31): 4/30 前披露
|
||
- Q2 (6/30): 8/31 前披露
|
||
- Q3 (9/30): 10/31 前披露
|
||
- Q4 (12/31): 次年 4/30 前披露
|
||
回测 date 当日"已披露"的最新季度:
|
||
- 1/1 ~ 3/31: 上年 Q3 (year-1, Q3)
|
||
- 4/1 ~ 4/30: 上年 Q3(年报可能还没出)→ 保险起见用上年 Q3
|
||
- 5/1 ~ 8/31: 当年 Q1
|
||
- 9/1 ~ 10/31: 当年 Q2
|
||
- 11/1 ~ 12/31: 当年 Q3
|
||
简化:date.month 决定 quarter,date.year 决定 year。
|
||
"""
|
||
d = datetime.strptime(date_str[:10], "%Y-%m-%d")
|
||
m = d.month
|
||
if m <= 4:
|
||
return d.year - 1, 3, f"{d.year - 1}-09-30"
|
||
if m <= 8:
|
||
return d.year, 1, f"{d.year}-03-31"
|
||
if m <= 10:
|
||
return d.year, 2, f"{d.year}-06-30"
|
||
return d.year, 3, f"{d.year}-09-30"
|
||
|
||
|
||
__all__ = ["BaostockProvider", "jq_to_bs_code", "bs_to_jq_code"]
|