Files
sanguo_vnpy_v2/sanguo_portfolio/providers/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

926 lines
37 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: 回测数据源走 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"]