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)
This commit is contained in:
@@ -3,28 +3,42 @@
|
||||
子模块:
|
||||
- ``factors`` 纯函数估值/ROIC 因子(无外部依赖)
|
||||
- ``filters`` ST/停牌/科创北交/次新/涨跌停 过滤(无外部依赖)
|
||||
- ``providers`` SanguoMiniQmtProvider(继承 bullet-trade MiniQMTProvider)
|
||||
- ``providers`` BaostockProvider(Mac/回测,历史成分股+TTM) /
|
||||
SanguoMiniQmtProvider(VPS 实盘,继承 bullet-trade MiniQMTProvider)
|
||||
- ``strategies`` AllWeatherStrategy(聚宽 post48819 翻译)
|
||||
|
||||
ENV GUARD(bullet-trade 0.9.2 坑):
|
||||
- ``import bullet_trade`` 时 default provider 是 jqdata,会强制 ``import jqdatasdk``,
|
||||
本地/服务器都没装(用户铁律不用 jqdata 付费)。
|
||||
- 在任何 ``import bullet_trade`` **之前**设 ``DEFAULT_DATA_PROVIDER=miniqmt``,
|
||||
default provider 切到 MiniQMTProvider,跳过 jqdatasdk。
|
||||
- 实际数据走 ``set_data_provider(SanguoMiniQmtProvider(...))`` 覆盖。
|
||||
- 在任何 ``import bullet_trade`` **之前**设 ``DEFAULT_DATA_PROVIDER``,
|
||||
default provider 切到非 jqdata provider,跳过 jqdatasdk。
|
||||
- 实际数据走 ``set_data_provider(...)`` 覆盖。
|
||||
"""
|
||||
import os
|
||||
import sys as _sys
|
||||
from unittest.mock import MagicMock as _MagicMock
|
||||
|
||||
# 默认 provider 切 miniqmt,避开 jqdatasdk 硬 import(必须早于任何 bullet_trade import)
|
||||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "miniqmt")
|
||||
# ENV + mock 必须早于任何 bullet_trade import(providers 子模块会触发 bullet_trade):
|
||||
# miniqmt → import xtquant(周六休市 miniQMT 客户端不响应→卡死)
|
||||
# jqdata → import jqdatasdk(用户铁律不装→ModuleNotFoundError)
|
||||
# 方案: ENV 设 jqdata + 预插 mock jqdatasdk(@jq.utils.assert_auth 装饰器 passthrough),
|
||||
# 让 bullet_trade import 走 jqdata 分支拿 mock 不崩不卡;真实 provider 由
|
||||
# set_data_provider 运行时注入覆盖(local/baostock/miniqmt)。
|
||||
# 必须在包 __init__ 顶部(providers import 触发 bullet_trade 之前),runner_backtest 顶部太晚。
|
||||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "jqdata")
|
||||
if "jqdatasdk" not in _sys.modules:
|
||||
_m = _MagicMock()
|
||||
_m.utils.assert_auth = lambda func: func # jqdata.py 用 @jq.utils.assert_auth 装饰器
|
||||
_sys.modules["jqdatasdk"] = _m
|
||||
|
||||
from . import factors, filters
|
||||
from .providers import SanguoMiniQmtProvider
|
||||
from .providers import BaostockProvider, SanguoMiniQmtProvider
|
||||
from .strategies import AllWeatherConfig, AllWeatherStrategy, BrokerFacade
|
||||
|
||||
__all__ = [
|
||||
"factors",
|
||||
"filters",
|
||||
"BaostockProvider",
|
||||
"SanguoMiniQmtProvider",
|
||||
"AllWeatherStrategy",
|
||||
"AllWeatherConfig",
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
"""sanguo_portfolio 数据 provider 层。"""
|
||||
from .baostock_provider import BaostockProvider
|
||||
from .local_parquet_provider import LocalParquetProvider
|
||||
from .sanguo_fundamentals import SanguoMiniQmtProvider
|
||||
|
||||
__all__ = ["SanguoMiniQmtProvider"]
|
||||
__all__ = ["SanguoMiniQmtProvider", "BaostockProvider", "LocalParquetProvider"]
|
||||
|
||||
@@ -0,0 +1,925 @@
|
||||
"""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"]
|
||||
@@ -0,0 +1,571 @@
|
||||
"""LocalParquetProvider: 读 VPS 本地 parquet/csv,零 online 调用。
|
||||
|
||||
数据布局(VPS ``C:\\sanguo_vnpy_v2\\data\\``,用户多源汇总,见 memory vps-local-data-layout):
|
||||
- 日线 K 线: ``qfq/{年}/{code}_daily.parquet`` (date/open/high/low/close/volume)
|
||||
- 三大表: ``static/{balance,income,cashflow}/{code}_{type}.parquet``
|
||||
akshare 东财大写列,通用列 SECUCODE/**REPORT_DATE**/REPORT_TYPE
|
||||
- 每日估值: ``static/valuation/{code}_valuation.parquet`` (中文列 PE(TTM)/市净率/总市值...)
|
||||
- 财务摘要: ``static/financial_abstract/{code}_*.parquet`` (宽表 指标×季度)
|
||||
- 成分股: ``static/index_const/index_const.parquet`` (⚠️ 仅当前快照→幸存者偏差缺口)
|
||||
|
||||
实现 bullet_trade ``DataProvider`` 接口; ``get_fundamentals_df`` 字段对齐
|
||||
``BaostockProvider._FUNDAMENTAL_COLUMNS``(策略 all_weather 依赖)。
|
||||
|
||||
优势(vs BaostockProvider 实时调 baostock HTTP):
|
||||
- 三表是完整绝对值(balance 221列/income 170列), ``total_liability``/``retained_profit`` 填真值
|
||||
(BaostockProvider 比率字段反推受限,多 NaN)
|
||||
- valuation PE(TTM)/PB/PS/PCF 是 akshare 服务端现成值,不用 4 季自滚 TTM
|
||||
- 零 online: 不踩 baostock 限频/黑名单/休市坑(见 memory provider-local-data-only)
|
||||
|
||||
已知缺口(V1 标注,不阻塞 MVP):
|
||||
- 历史成分股: index_const 仅 2026-07-17 最新一期 → 回测历史有幸存者偏差
|
||||
- gross_profit_margin: income 无明确"营业成本"列, V1 NaN, v2 改读 financial_abstract 现成值
|
||||
- roic: 需有息负债拆分, V1 NaN
|
||||
- 单位口径假设: 市值=元(/1e8转亿)、PE/PB=数值、YOY=百分数(/100转小数); 验证时看数值范围校准
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import pandas as pd
|
||||
|
||||
# bullet-trade 可能未装,容错 import DataProvider(照 baostock_provider 模式)
|
||||
try:
|
||||
from bullet_trade.data.providers.base import DataProvider # type: ignore
|
||||
except ImportError: # Mac dev 环境未装,允许模块加载
|
||||
class DataProvider: # type: ignore[no-redef]
|
||||
name: str = "base"
|
||||
|
||||
from ..factors.valuation import to_yi
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# VPS 数据根目录(Windows 路径; Mac 测试时通过 config["data_dir"] 覆盖)
|
||||
_DEFAULT_DATA_DIR = r"C:\sanguo_vnpy_v2\data"
|
||||
|
||||
# jq 代码 ↔ VPS 文件名代码(600519.XSHG ↔ 600519.SH)
|
||||
_JQ_TO_FILE_SUFFIX = {"XSHG": "SH", "XSHE": "SZ", "SH": "SH", "SZ": "SZ"}
|
||||
_FILE_TO_JQ_SUFFIX = {"SH": "XSHG", "SZ": "XSHE"}
|
||||
|
||||
|
||||
def jq_to_file_code(jq_code: str) -> str:
|
||||
"""``600519.XSHG`` → ``600519.SH`` (VPS parquet 文件名格式)。纯数字透传。"""
|
||||
if not jq_code or "." not in jq_code:
|
||||
return jq_code
|
||||
code, suffix = jq_code.split(".", 1)
|
||||
file_suffix = _JQ_TO_FILE_SUFFIX.get(suffix.upper())
|
||||
return f"{code}.{file_suffix}" if file_suffix else jq_code
|
||||
|
||||
|
||||
def file_to_jq_code(file_code: str) -> str:
|
||||
"""``600519.SH`` → ``600519.XSHG``。纯 6 位按 6开头=sh/0,3开头=sz 推断。"""
|
||||
if not file_code:
|
||||
return file_code
|
||||
if "." not in file_code:
|
||||
if len(file_code) == 6:
|
||||
return f"{file_code}.{'XSHG' if file_code.startswith('6') else 'XSHE'}"
|
||||
return file_code
|
||||
code, suffix = file_code.split(".", 1)
|
||||
jq_suffix = _FILE_TO_JQ_SUFFIX.get(suffix.upper())
|
||||
return f"{code}.{jq_suffix}" if jq_suffix else file_code
|
||||
|
||||
|
||||
# jq → VPS K 线文件名(baostock 风格 sh/sz 前缀无点; 与三表 jq 后缀格式不同!)
|
||||
_KLINE_PREFIX = {"XSHG": "sh", "XSHE": "sz", "SH": "sh", "SZ": "sz"}
|
||||
|
||||
|
||||
def jq_to_kline_code(jq_code: str) -> str:
|
||||
"""``600519.XSHG`` → ``sh600519`` (VPS qfq/raw K线文件名)。纯 6 位按 6开头=sh 推断。"""
|
||||
if not jq_code:
|
||||
return jq_code
|
||||
if "." not in jq_code:
|
||||
if len(jq_code) == 6:
|
||||
return ("sh" if jq_code.startswith("6") else "sz") + jq_code
|
||||
return jq_code
|
||||
code, suffix = jq_code.split(".", 1)
|
||||
prefix = _KLINE_PREFIX.get(suffix.upper())
|
||||
return (prefix + code) if prefix else jq_code
|
||||
|
||||
|
||||
# valuation parquet 中文列 → 英文
|
||||
_VAL_COL_MAP = {
|
||||
"数据日期": "date", "当日收盘价": "close", "当日涨跌幅": "pct_chg",
|
||||
"总市值": "total_market_cap", "流通市值": "circ_market_cap",
|
||||
"总股本": "total_share", "流通股本": "circ_share",
|
||||
"PE(TTM)": "pe_ttm", "PE(静)": "pe_static",
|
||||
"市净率": "pb", "PEG值": "peg", "市现率": "pcf", "市销率": "ps",
|
||||
}
|
||||
|
||||
|
||||
def _to_float(v: Any) -> Optional[float]:
|
||||
if v is None:
|
||||
return None
|
||||
if isinstance(v, (int, float)):
|
||||
return float(v)
|
||||
try:
|
||||
s = str(v).strip().replace(",", "").replace("%", "")
|
||||
return float(s) if s else None
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _pct_to_decimal(v: Any) -> float:
|
||||
"""百分数(18.5 表示 18.5%) → 小数(0.185)。None/异常 → NaN。akshare YOY 通常百分数。"""
|
||||
f = _to_float(v)
|
||||
if f is None:
|
||||
return float("nan")
|
||||
return f / 100.0
|
||||
|
||||
|
||||
def _or_nan(v: Any) -> float:
|
||||
f = _to_float(v)
|
||||
return f if f is not None else float("nan")
|
||||
|
||||
|
||||
# 策略 all_weather 依赖的 fundamentals 输出列(对齐 BaostockProvider._FUNDAMENTAL_COLUMNS)
|
||||
_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",
|
||||
]
|
||||
|
||||
|
||||
class LocalParquetProvider(DataProvider): # type: ignore[misc]
|
||||
"""读 VPS 本地 parquet 的数据 provider(回测专用,零 online)。
|
||||
|
||||
所有方法读 ``data_dir`` 下 parquet 文件,不调任何外部 API。
|
||||
"""
|
||||
|
||||
name: str = "sanguo_local_parquet"
|
||||
requires_live_data: bool = False
|
||||
|
||||
def __init__(self, config: Optional[Dict[str, Any]] = None) -> None:
|
||||
cfg = config or {}
|
||||
self.data_dir: str = cfg.get("data_dir", _DEFAULT_DATA_DIR)
|
||||
# 缓存:同股多次读只一次 IO
|
||||
self._val_cache: Dict[str, pd.DataFrame] = {}
|
||||
self._quarter_cache: Dict[tuple, pd.DataFrame] = {}
|
||||
self._index_const_cache: Optional[pd.DataFrame] = None
|
||||
|
||||
# ==================== 路径辅助 ====================
|
||||
def _valuation_path(self, file_code: str) -> str:
|
||||
return os.path.join(self.data_dir, "static", "valuation", f"{file_code}_valuation.parquet")
|
||||
|
||||
def _static_path(self, table: str, file_code: str) -> str:
|
||||
return os.path.join(self.data_dir, "static", table, f"{file_code}_{table}.parquet")
|
||||
|
||||
@staticmethod
|
||||
def _year_range(start: Optional[pd.Timestamp], end: Optional[pd.Timestamp]) -> range:
|
||||
s = start.year if start is not None else 2010
|
||||
e = end.year if end is not None else datetime.now().year
|
||||
if e < s:
|
||||
s, e = e, s
|
||||
return range(s, e + 1)
|
||||
|
||||
# ==================== get_price ====================
|
||||
def get_price(
|
||||
self,
|
||||
security: Union[str, List[str]],
|
||||
start_date: Union[str, datetime] = None,
|
||||
end_date: Union[str, datetime] = None,
|
||||
frequency: str = "day",
|
||||
fields: Optional[List[str]] = None,
|
||||
skip_paused: bool = True,
|
||||
fq: str = "qfq",
|
||||
count: Optional[int] = None,
|
||||
panel: bool = True,
|
||||
fill_paused: bool = True,
|
||||
) -> pd.DataFrame:
|
||||
"""读本地 qfq/raw 日线 parquet,拼多年 + 过滤日期区间。
|
||||
|
||||
聚宽/bullet_trade 兼容参数:
|
||||
- ``count``: 无 start_date 时取 end_date 前 N 根
|
||||
- ``panel``: True=多股 panel(index=date,外层 code); False=长表(time/code/fields)
|
||||
bullet_trade ``_trend_mean`` 用 panel=False + pivot(index=time,columns=code)
|
||||
- ``fill_paused``: 停牌填充(忽略,直接读原始)
|
||||
"""
|
||||
codes = [security] if isinstance(security, str) else list(security or [])
|
||||
freq_dir = "raw" if fq == "raw" else "qfq"
|
||||
if frequency.startswith("min") or frequency in ("1m", "1min"):
|
||||
freq_dir = "minute_15" # V1: 分钟线只支持 15min 目录
|
||||
|
||||
start = pd.Timestamp(start_date) if start_date else None
|
||||
end = pd.Timestamp(end_date) if end_date else None
|
||||
# count 模式: 无 start_date, 读 end 前 N 根(近 3 年覆盖足够)
|
||||
if count and start is None:
|
||||
end_for_count = end or pd.Timestamp.now()
|
||||
years = range(end_for_count.year - 2, end_for_count.year + 1)
|
||||
else:
|
||||
years = self._year_range(start, end)
|
||||
|
||||
frames: Dict[str, pd.DataFrame] = {}
|
||||
for jq_code in codes:
|
||||
fc = jq_to_kline_code(jq_code)
|
||||
parts: List[pd.DataFrame] = []
|
||||
for y in years:
|
||||
p = os.path.join(self.data_dir, freq_dir, str(y), f"{fc}_daily.parquet")
|
||||
if os.path.exists(p):
|
||||
try:
|
||||
parts.append(pd.read_parquet(p))
|
||||
except Exception as exc:
|
||||
logger.warning("读 K 线失败 %s/%s: %s", y, fc, exc)
|
||||
if not parts:
|
||||
continue
|
||||
df = pd.concat(parts, ignore_index=True)
|
||||
if "date" in df.columns:
|
||||
df["date"] = pd.to_datetime(df["date"])
|
||||
df = df.sort_values("date")
|
||||
if end is not None:
|
||||
df = df[df["date"] <= end]
|
||||
if start is not None:
|
||||
df = df[df["date"] >= start]
|
||||
if count:
|
||||
df = df.tail(count) # 取最近 count 根
|
||||
df = df.set_index("date")
|
||||
if fields:
|
||||
keep = [c for c in fields if c in df.columns]
|
||||
df = df[keep] if keep else df
|
||||
frames[jq_code] = df
|
||||
|
||||
if not frames:
|
||||
return pd.DataFrame()
|
||||
# panel=False: 长表(time/code/fields), 兼容 bullet_trade pivot
|
||||
if not panel:
|
||||
long_parts = []
|
||||
for jq_code, df in frames.items():
|
||||
d = df.reset_index().rename(columns={"date": "time"})
|
||||
d.insert(0, "code", jq_code)
|
||||
long_parts.append(d)
|
||||
return pd.concat(long_parts, ignore_index=True) if long_parts else pd.DataFrame()
|
||||
if len(frames) == 1:
|
||||
return next(iter(frames.values()))
|
||||
try:
|
||||
return pd.concat(frames, axis=1)
|
||||
except Exception as exc:
|
||||
logger.warning("多股 panel concat 失败,返回首只: %s", exc)
|
||||
return next(iter(frames.values()))
|
||||
|
||||
# ==================== 估值/三表 读取 ====================
|
||||
def _read_valuation(self, file_code: str) -> pd.DataFrame:
|
||||
if file_code in self._val_cache:
|
||||
return self._val_cache[file_code]
|
||||
p = self._valuation_path(file_code)
|
||||
if not os.path.exists(p):
|
||||
self._val_cache[file_code] = pd.DataFrame()
|
||||
return pd.DataFrame()
|
||||
try:
|
||||
df = pd.read_parquet(p).rename(columns=_VAL_COL_MAP)
|
||||
if "date" in df.columns:
|
||||
df["date"] = pd.to_datetime(df["date"], errors="coerce")
|
||||
df = df.sort_values("date")
|
||||
self._val_cache[file_code] = df
|
||||
return df
|
||||
except Exception as exc:
|
||||
logger.warning("读 valuation 失败 %s: %s", file_code, exc)
|
||||
self._val_cache[file_code] = pd.DataFrame()
|
||||
return pd.DataFrame()
|
||||
|
||||
def _read_quarter(self, table: str, file_code: str) -> pd.DataFrame:
|
||||
key = (table, file_code)
|
||||
if key in self._quarter_cache:
|
||||
return self._quarter_cache[key]
|
||||
p = self._static_path(table, file_code)
|
||||
if not os.path.exists(p):
|
||||
self._quarter_cache[key] = pd.DataFrame()
|
||||
return pd.DataFrame()
|
||||
try:
|
||||
df = pd.read_parquet(p)
|
||||
if "REPORT_DATE" in df.columns:
|
||||
df["REPORT_DATE"] = pd.to_datetime(df["REPORT_DATE"], errors="coerce")
|
||||
df = df.sort_values("REPORT_DATE")
|
||||
self._quarter_cache[key] = df
|
||||
return df
|
||||
except Exception as exc:
|
||||
logger.warning("读 %s 失败 %s: %s", table, file_code, exc)
|
||||
self._quarter_cache[key] = pd.DataFrame()
|
||||
return pd.DataFrame()
|
||||
|
||||
def _read_financial_abstract(self, file_code: str) -> Optional[pd.DataFrame]:
|
||||
"""读 financial_abstract 宽表(指标×季度, 列: 选项/指标/20260331/20251231/...)。"""
|
||||
p = os.path.join(self.data_dir, "static", "financial_abstract", f"{file_code}_financial_abstract.parquet")
|
||||
if not os.path.exists(p):
|
||||
return None
|
||||
try:
|
||||
return pd.read_parquet(p)
|
||||
except Exception as exc:
|
||||
logger.warning("读 financial_abstract 失败 %s: %s", file_code, exc)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _latest_indicator(fa_df: Optional[pd.DataFrame], indicator_name: str) -> Optional[float]:
|
||||
"""从 financial_abstract 宽表取指定指标最新季度值(第一个季度列)。"""
|
||||
if fa_df is None or "指标" not in fa_df.columns:
|
||||
return None
|
||||
rows = fa_df[fa_df["指标"] == indicator_name]
|
||||
if rows.empty:
|
||||
return None
|
||||
quarter_cols = [c for c in fa_df.columns if c not in ("选项", "指标")]
|
||||
if not quarter_cols:
|
||||
return None
|
||||
return _to_float(rows.iloc[0][quarter_cols[0]])
|
||||
|
||||
@staticmethod
|
||||
def _latest_row_before(
|
||||
df: pd.DataFrame, date_col: str, date_str: str,
|
||||
) -> Optional[pd.Series]:
|
||||
"""取 ``date_col <= date_str`` 的最后一行(最新已披露)。"""
|
||||
if df is None or df.empty or date_col not in df.columns:
|
||||
return None
|
||||
ts = pd.Timestamp(date_str)
|
||||
sub = df[df[date_col] <= ts]
|
||||
return sub.iloc[-1] if not sub.empty else None
|
||||
|
||||
# ==================== get_fundamentals_df ====================
|
||||
def get_fundamentals_df(
|
||||
self,
|
||||
stocks: List[str],
|
||||
date: Optional[Union[str, datetime]] = None,
|
||||
) -> pd.DataFrame:
|
||||
"""合并多股 fundamentals,列对齐 ``_FUNDAMENTAL_COLUMNS``。
|
||||
|
||||
数据源(本地 parquet):
|
||||
- valuation: market_cap/pe/pb/ps/pcf(akshare 服务端现成值)
|
||||
- income: eps(BASIC_EPS)/inc_*_yoy(OPERATE_*_YOY)/net_profit_margin(NP÷营收 算)
|
||||
- balance: total_liability/total_sheet_owner_equities/retained_profit(绝对值现成)
|
||||
- 算: roe(NP÷权益)/roa(NP÷总资产)
|
||||
"""
|
||||
if not stocks:
|
||||
return pd.DataFrame(columns=_FUNDAMENTAL_COLUMNS)
|
||||
|
||||
date_str = self._to_date_str(date) or datetime.now().strftime("%Y-%m-%d")
|
||||
rows: List[Dict[str, Any]] = [
|
||||
self._build_fundamental_row(jq_code, date_str) for jq_code in stocks
|
||||
]
|
||||
df = pd.DataFrame(rows, columns=_FUNDAMENTAL_COLUMNS)
|
||||
if "code" in df.columns:
|
||||
df = df.set_index("code", drop=False)
|
||||
return df
|
||||
|
||||
def _build_fundamental_row(self, jq_code: str, date_str: str) -> Dict[str, Any]:
|
||||
fc = jq_to_file_code(jq_code)
|
||||
row: Dict[str, Any] = {"code": jq_code}
|
||||
|
||||
val = self._latest_row_before(self._read_valuation(fc), "date", date_str)
|
||||
inc = self._latest_row_before(self._read_quarter("income", fc), "REPORT_DATE", date_str)
|
||||
bal = self._latest_row_before(self._read_quarter("balance", fc), "REPORT_DATE", date_str)
|
||||
|
||||
def g(d: Optional[pd.Series], k: str) -> Optional[float]:
|
||||
return _to_float(d.get(k)) if d is not None else None
|
||||
|
||||
# --- 估值字段(akshare valuation: 市值元, PE/PB 数值) ---
|
||||
mkt = g(val, "total_market_cap")
|
||||
circ = g(val, "circ_market_cap")
|
||||
row["market_cap"] = to_yi(mkt) if mkt else float("nan")
|
||||
row["circulating_market_cap"] = to_yi(circ) if circ else float("nan")
|
||||
row["pe_ratio"] = _or_nan(g(val, "pe_ttm"))
|
||||
row["pb_ratio"] = _or_nan(g(val, "pb"))
|
||||
row["ps_ratio"] = _or_nan(g(val, "ps"))
|
||||
row["pcf_ratio"] = _or_nan(g(val, "pcf"))
|
||||
|
||||
# --- 利润表字段 ---
|
||||
row["eps"] = _or_nan(g(inc, "BASIC_EPS"))
|
||||
row["inc_revenue_year_on_year"] = _pct_to_decimal(g(inc, "OPERATE_INCOME_YOY"))
|
||||
row["inc_operation_profit_year_on_year"] = _pct_to_decimal(g(inc, "OPERATE_PROFIT_YOY"))
|
||||
# inc_total_revenue: akshare income 无 total_revenue_YOY 独立列,用 OPERATE_INCOME_YOY 近似
|
||||
row["inc_total_revenue_year_on_year"] = _pct_to_decimal(g(inc, "OPERATE_INCOME_YOY"))
|
||||
|
||||
net_profit = g(inc, "PARENT_NETPROFIT") or g(inc, "NETPROFIT")
|
||||
revenue = g(inc, "OPERATE_INCOME")
|
||||
total_assets = g(bal, "TOTAL_ASSETS")
|
||||
parent_equity = g(bal, "TOTAL_PARENT_EQUITY")
|
||||
|
||||
# net_profit_margin = 归母净利润 / 营收(小数)
|
||||
row["net_profit_margin"] = (
|
||||
net_profit / revenue
|
||||
if (net_profit and revenue and revenue != 0)
|
||||
else float("nan")
|
||||
)
|
||||
|
||||
# --- 资产负债表(绝对值,元→亿) ---
|
||||
total_liab = g(bal, "TOTAL_LIABILITIES")
|
||||
row["total_liability"] = to_yi(total_liab) if total_liab else float("nan")
|
||||
row["total_sheet_owner_equities"] = to_yi(parent_equity) if parent_equity else float("nan")
|
||||
retained = (g(bal, "SURPLUS_RESERVE") or 0) + (g(bal, "UNASSIGN_RPOFIT") or 0)
|
||||
row["retained_profit"] = to_yi(retained) if retained else float("nan")
|
||||
|
||||
# --- 算指标(单期非年化TTM;v2 改 financial_abstract 现成年化值) ---
|
||||
row["roe"] = (
|
||||
net_profit / parent_equity
|
||||
if (net_profit and parent_equity and parent_equity != 0)
|
||||
else float("nan")
|
||||
)
|
||||
row["roa"] = (
|
||||
net_profit / total_assets
|
||||
if (net_profit and total_assets and total_assets != 0)
|
||||
else float("nan")
|
||||
)
|
||||
|
||||
# gross_profit_margin: 从 financial_abstract 读现成"毛利率"(百分数→小数)
|
||||
fa = self._read_financial_abstract(fc)
|
||||
row["gross_profit_margin"] = _pct_to_decimal(self._latest_indicator(fa, "毛利率"))
|
||||
# roic: 需有息负债拆分 → V1 NaN
|
||||
# TODO v2: roic = NOPAT / (权益 + 有息负债 - 现金)
|
||||
row["roic"] = float("nan")
|
||||
return row
|
||||
|
||||
@staticmethod
|
||||
def _to_date_str(value: Optional[Union[str, datetime]]) -> Optional[str]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return value.strftime("%Y-%m-%d")
|
||||
|
||||
# ==================== get_security_info ====================
|
||||
def get_security_info(self, security: str) -> Dict[str, Any]:
|
||||
fc = jq_to_file_code(security)
|
||||
val_df = self._read_valuation(fc)
|
||||
if val_df.empty:
|
||||
return {"code": security, "display_name": security, "name": security}
|
||||
last = val_df.iloc[-1]
|
||||
return {
|
||||
"code": security,
|
||||
"display_name": security, # valuation 无名称,用 code
|
||||
"name": security,
|
||||
"start_date": val_df["date"].min().strftime("%Y-%m-%d") if "date" in val_df.columns else None,
|
||||
"end_date": val_df["date"].max().strftime("%Y-%m-%d") if "date" in val_df.columns else None,
|
||||
"type": "stock",
|
||||
}
|
||||
|
||||
# ==================== get_trade_days ====================
|
||||
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]:
|
||||
"""从蓝筹 sh600000 K 线 date 列取交易日(锚定,全市场交易日一致)。
|
||||
|
||||
bullet_trade 引擎调 ``get_trade_days(count=N)`` 取最近 N 天(无 start_date),
|
||||
故兼容 count 参数(其他基类方法也可能传 count)。
|
||||
"""
|
||||
fc = "sh600000"
|
||||
start_ts = pd.Timestamp(start_date) if start_date else None
|
||||
end_ts = pd.Timestamp(end_date) if end_date else None
|
||||
if count and not start_ts:
|
||||
now_y = datetime.now().year
|
||||
years = range(now_y - 2, now_y + 1) # 近 3 年足够覆盖 count 天
|
||||
else:
|
||||
years = self._year_range(start_ts, end_ts)
|
||||
days: List[datetime] = []
|
||||
for y in years:
|
||||
p = os.path.join(self.data_dir, "qfq", str(y), f"{fc}_daily.parquet")
|
||||
if not os.path.exists(p):
|
||||
continue
|
||||
try:
|
||||
df = pd.read_parquet(p, columns=["date"])
|
||||
for d in pd.to_datetime(df["date"]):
|
||||
days.append(d.to_pydatetime())
|
||||
except Exception as exc:
|
||||
logger.warning("读交易日失败 %s: %s", y, exc)
|
||||
if not days:
|
||||
return []
|
||||
days = sorted(set(days))
|
||||
if start_ts:
|
||||
days = [d for d in days if pd.Timestamp(d) >= start_ts]
|
||||
if end_ts:
|
||||
days = [d for d in days if pd.Timestamp(d) <= end_ts]
|
||||
if count:
|
||||
days = days[-count:]
|
||||
return days
|
||||
|
||||
# ==================== get_all_securities ====================
|
||||
def get_all_securities(
|
||||
self, types: Optional[List[str]] = None,
|
||||
) -> pd.DataFrame:
|
||||
"""列 ``static/valuation/`` 下所有股票(文件名 → jq code)。"""
|
||||
val_dir = os.path.join(self.data_dir, "static", "valuation")
|
||||
if not os.path.isdir(val_dir):
|
||||
return pd.DataFrame(columns=["code", "display_name"])
|
||||
codes: List[str] = []
|
||||
for fn in os.listdir(val_dir):
|
||||
if fn.endswith("_valuation.parquet"):
|
||||
codes.append(file_to_jq_code(fn.replace("_valuation.parquet", "")))
|
||||
return pd.DataFrame({"code": codes, "display_name": codes})
|
||||
|
||||
# ==================== get_index_stocks ====================
|
||||
def get_index_stocks(
|
||||
self,
|
||||
index_symbol: str,
|
||||
date: Optional[Union[str, datetime]] = None,
|
||||
) -> List[str]:
|
||||
"""读 ``index_const.parquet`` 过滤指数成分。
|
||||
|
||||
⚠️ 缺口:VPS index_const 仅 2026-07-17 最新一期(当前快照),
|
||||
回测历史日期会用到"现在还在指数里的股票"→ 幸存者偏差(结果虚高)。
|
||||
``date`` 参数目前忽略(无历史数据),待补 csindex 历史成分。
|
||||
"""
|
||||
ic = self._load_index_const()
|
||||
if ic is None or ic.empty:
|
||||
logger.warning("index_const.parquet 无数据,get_index_stocks 返回空")
|
||||
return []
|
||||
idx = index_symbol.split(".")[0] if "." in index_symbol else index_symbol
|
||||
col = "指数代码" if "指数代码" in ic.columns else "index_code"
|
||||
sub = ic[ic[col].astype(str).str.contains(idx, na=False)]
|
||||
code_col = "成分券代码" if "成分券代码" in ic.columns else None
|
||||
if code_col is None:
|
||||
return []
|
||||
return [file_to_jq_code(str(c)) for c in sub[code_col].tolist()]
|
||||
|
||||
def _load_index_const(self) -> Optional[pd.DataFrame]:
|
||||
if self._index_const_cache is not None:
|
||||
return self._index_const_cache
|
||||
p = os.path.join(self.data_dir, "static", "index_const", "index_const.parquet")
|
||||
if not os.path.exists(p):
|
||||
self._index_const_cache = None
|
||||
return None
|
||||
try:
|
||||
self._index_const_cache = pd.read_parquet(p)
|
||||
return self._index_const_cache
|
||||
except Exception as exc:
|
||||
logger.warning("读 index_const 失败: %s", exc)
|
||||
self._index_const_cache = None
|
||||
return None
|
||||
|
||||
# ==================== get_split_dividend (qfq 已复权,占位) ====================
|
||||
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]]:
|
||||
"""除权除息记录。读 qfq 日线已前复权,回测不依赖此方法 → 占位返空 list。
|
||||
|
||||
TODO v2:若需 raw→qfq 自算,从 outstanding_share 变化 + 派息记录派生。
|
||||
"""
|
||||
return []
|
||||
|
||||
# ==================== get_current_tick (回测不用,占位) ====================
|
||||
def get_current_tick(self, security: str) -> Optional[Dict[str, Any]]:
|
||||
"""回测不用实时 tick;从 valuation 最新行推算 close + 涨跌停(filter_limitup 用)。"""
|
||||
val_df = self._read_valuation(jq_to_file_code(security))
|
||||
if val_df.empty:
|
||||
return None
|
||||
last = val_df.iloc[-1]
|
||||
close = _to_float(last.get("close"))
|
||||
pct = _to_float(last.get("pct_chg")) or 0.0
|
||||
# 涨跌停:主板 ±10%(ST/创业/科创 精确规则 v2 补)
|
||||
high_limit = round(close * 1.1, 2) if close else None
|
||||
low_limit = round(close * 0.9, 2) if close else None
|
||||
return {
|
||||
"code": security, "current_price": close, "close": close,
|
||||
"high_limit": high_limit, "low_limit": low_limit,
|
||||
"change_percent": pct,
|
||||
}
|
||||
@@ -1,21 +1,37 @@
|
||||
"""全天候策略回测入口。
|
||||
|
||||
用法(VPS Windows / miniQMT 已连):
|
||||
set DEFAULT_DATA_PROVIDER=miniqmt
|
||||
用法(Mac 默认 baostock;VPS Windows / miniQMT 已连用 miniqmt):
|
||||
# Mac 默认 baostock(跨平台,不依赖 miniQMT 客户端)
|
||||
python -m sanguo_portfolio.runner_backtest \\
|
||||
--start 2020-01-01 --end 2024-12-31 --cash 1000000
|
||||
|
||||
# VPS miniQMT(实盘/精准 xtquant)
|
||||
python -m sanguo_portfolio.runner_backtest --provider miniqmt \\
|
||||
--start 2020-01-01 --end 2024-12-31 --cash 1000000
|
||||
|
||||
JSON 输出(供 SSH 捕获,前端 MVP 用):
|
||||
python -m sanguo_portfolio.runner_backtest --json \\
|
||||
--start 2024-01-01 --end 2024-02-29 --cash 1000000
|
||||
|
||||
Mac 没装 xtquant,这里仅作为入口脚本(测试用 mock,实际跑 rsync 到 VPS)。
|
||||
Mac 跑 baostock 默认链路;miniQMT 链路仍保留(实盘 runner_live 用)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
# ENV GUARD 必须早于任何 bullet_trade import
|
||||
# bullet_trade __init__ 加载时 _create_provider() 读 DEFAULT_DATA_PROVIDER 创建默认 provider:
|
||||
# miniqmt → import xtquant(周六休市 miniQMT 客户端不响应→卡死)
|
||||
# jqdata → import jqdatasdk(用户铁律不装→ModuleNotFoundError)
|
||||
# 方案: ENV 设 jqdata + 预插 mock jqdatasdk, 让 import 走 jqdata 分支拿 mock 不崩不卡;
|
||||
# 真实 provider 由 set_data_provider 运行时注入覆盖(local/baostock/miniqmt)。
|
||||
import os
|
||||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "miniqmt")
|
||||
import sys as _sys
|
||||
from unittest.mock import MagicMock as _MagicMock
|
||||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "jqdata")
|
||||
if "jqdatasdk" not in _sys.modules:
|
||||
_m = _MagicMock()
|
||||
# jqdata.py 用 @jq.utils.assert_auth 装饰器;MagicMock 的 assert_* 前缀被保护→AttributeError
|
||||
_m.utils.assert_auth = lambda func: func # passthrough 装饰器
|
||||
_sys.modules["jqdatasdk"] = _m
|
||||
|
||||
import argparse
|
||||
import json
|
||||
@@ -33,6 +49,10 @@ def parse_args() -> argparse.Namespace:
|
||||
p.add_argument("--benchmark", default="000300.XSHG", help="基准代码")
|
||||
p.add_argument("--max-pool", type=int, default=0, help="限制选股池前N只(0=不限,MVP验证用)")
|
||||
p.add_argument("--frequency", default="day", help="回测频率 day/minute")
|
||||
p.add_argument(
|
||||
"--provider", default="local", choices=["local", "baostock", "miniqmt"],
|
||||
help="数据 provider:baostock(默认,Mac/跨平台,历史成分股+TTM) / miniqmt(VPS 实盘,需 xtquant)",
|
||||
)
|
||||
p.add_argument(
|
||||
"--provider-config", default="{}",
|
||||
help="provider 配置 JSON 字符串,如 '{\"data_dir\":\"D:/xtdata\"}'",
|
||||
@@ -48,10 +68,15 @@ def parse_args() -> argparse.Namespace:
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def build_provider(config_str: str) -> Any:
|
||||
"""构造 SanguoMiniQmtProvider。"""
|
||||
def build_provider(provider_name: str, config_str: str) -> Any:
|
||||
"""构造 provider 实例。
|
||||
|
||||
Args:
|
||||
provider_name: "baostock"(Mac 默认) 或 "miniqmt"(VPS 实盘)
|
||||
config_str: provider 配置 JSON 字符串
|
||||
"""
|
||||
import json
|
||||
from .providers import SanguoMiniQmtProvider
|
||||
from .providers import BaostockProvider, LocalParquetProvider, SanguoMiniQmtProvider
|
||||
|
||||
cfg: Dict[str, Any] = {}
|
||||
if config_str and config_str != "{}":
|
||||
@@ -60,7 +85,15 @@ def build_provider(config_str: str) -> Any:
|
||||
except Exception as exc:
|
||||
logger.warning("provider-config 解析失败,用默认: %s", exc)
|
||||
cfg.setdefault("mode", "backtest")
|
||||
return SanguoMiniQmtProvider(cfg)
|
||||
|
||||
name = (provider_name or "baostock").lower()
|
||||
if name == "miniqmt":
|
||||
return SanguoMiniQmtProvider(cfg)
|
||||
if name == "baostock":
|
||||
return BaostockProvider(cfg)
|
||||
if name == "local":
|
||||
return LocalParquetProvider(cfg)
|
||||
raise ValueError(f"未知 provider: {name}(支持: local / baostock / miniqmt)")
|
||||
|
||||
|
||||
def build_broker_facade(engine: Any) -> Any:
|
||||
@@ -106,7 +139,7 @@ def run_backtest(args: argparse.Namespace) -> Dict[str, Any]:
|
||||
|
||||
from .strategies import AllWeatherStrategy, AllWeatherConfig
|
||||
|
||||
provider = build_provider(args.provider_config)
|
||||
provider = build_provider(args.provider, args.provider_config)
|
||||
set_data_provider(provider)
|
||||
|
||||
# 占位策略:initialize 里把 self(strategy)挂到聚宽风格定时器
|
||||
@@ -225,6 +258,7 @@ def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
cash=float(params.get("initial_cash", 1_000_000.0)),
|
||||
benchmark=params.get("benchmark", "000300.XSHG"),
|
||||
frequency="day",
|
||||
provider=params.get("provider", "local"),
|
||||
provider_config="{}",
|
||||
result_file="", # JSON 模式不写 md
|
||||
max_pool=int(params.get("max_pool", 0)),
|
||||
@@ -382,6 +416,8 @@ def main() -> None:
|
||||
"end_date": args.end,
|
||||
"initial_cash": args.cash,
|
||||
"benchmark": args.benchmark,
|
||||
"provider": args.provider,
|
||||
"max_pool": args.max_pool,
|
||||
})
|
||||
print(json.dumps(result, ensure_ascii=False, default=str))
|
||||
else:
|
||||
|
||||
@@ -506,7 +506,17 @@ def _current_dt(context: Any) -> Any:
|
||||
def _previous_date_str(context: Any) -> Optional[str]:
|
||||
pd_ = getattr(context, "previous_date", None)
|
||||
if pd_ is None:
|
||||
return None
|
||||
# bullet_trade context 无 previous_date 属性, fallback 用 current_dt(当日):
|
||||
# get_fundamentals_df/get_index_stocks 取当日已披露的最新数据(季报/成分)
|
||||
cd = _current_dt(context)
|
||||
if cd is None:
|
||||
return None
|
||||
if isinstance(cd, str):
|
||||
return cd[:10]
|
||||
try:
|
||||
return cd.strftime("%Y-%m-%d")
|
||||
except AttributeError:
|
||||
return str(cd)[:10]
|
||||
if isinstance(pd_, str):
|
||||
return pd_[:10]
|
||||
try:
|
||||
|
||||
+173
-2
@@ -1,9 +1,10 @@
|
||||
"""pytest 配置 + mock xtquant fixtures。
|
||||
"""pytest 配置 + mock xtquant / baostock fixtures。
|
||||
|
||||
约束:
|
||||
- Mac 没 xtquant/miniQMT,所有 ``from xtquant import xtdata`` 必须 mock
|
||||
- baostock 已装在 venv310,但仍提供 mock fixture(单测不依赖网络)
|
||||
- bullet-trade 0.9.2 的 ``import bullet_trade`` 会触发 default provider=jqdata → import jqdatasdk
|
||||
→ 在 ``import bullet_trade`` 前设 ``DEFAULT_DATA_PROVIDER=miniqmt``(本文件最顶部)
|
||||
→ 在 ``import bullet_trade`` 前设 ``DEFAULT_DATA_PROVIDER``(本文件最顶部)
|
||||
- bullet-trade 可能没装完,所有 bullet-trade import 容错 skip
|
||||
"""
|
||||
from __future__ import annotations
|
||||
@@ -11,6 +12,8 @@ from __future__ import annotations
|
||||
import os
|
||||
|
||||
# 必须早于任何 bullet_trade import / sanguo_portfolio(它可能 lazy import bullet_trade)
|
||||
# sanguo_portfolio.providers.baostock_provider 顶部会 setdefault sanguo_baostock;
|
||||
# 这里若未指定则用 miniqmt(向后兼容 SanguoMiniQmtProvider 测试)
|
||||
os.environ.setdefault("DEFAULT_DATA_PROVIDER", "miniqmt")
|
||||
|
||||
import sys
|
||||
@@ -261,3 +264,171 @@ def pytest_collection_modifyitems(config, items):
|
||||
for item in items:
|
||||
if "requires_bullet_trade" in item.keywords:
|
||||
item.add_marker(skip_bt)
|
||||
|
||||
|
||||
# ------------------------ baostock mock ------------------------
|
||||
class _FakeResultData:
|
||||
"""模拟 baostock ResultData(支持 next()/get_row_data()/fields/get_data())。
|
||||
|
||||
用 ``pd.DataFrame`` 构造,迭代器风格访问兼容 baostock 官方文档示例。
|
||||
"""
|
||||
|
||||
def __init__(self, df: pd.DataFrame, error_code: str = "0", error_msg: str = "success"):
|
||||
self._df = df.reset_index(drop=True) if isinstance(df, pd.DataFrame) else pd.DataFrame()
|
||||
self._idx = -1
|
||||
self.error_code = error_code
|
||||
self.error_msg = error_msg
|
||||
self.fields: List[str] = list(self._df.columns)
|
||||
|
||||
def next(self) -> bool:
|
||||
self._idx += 1
|
||||
return self._idx < len(self._df)
|
||||
|
||||
def get_row_data(self) -> List[Any]:
|
||||
if 0 <= self._idx < len(self._df):
|
||||
return [self._df.iloc[self._idx][c] for c in self._df.columns]
|
||||
return []
|
||||
|
||||
def get_data(self) -> pd.DataFrame:
|
||||
return self._df.copy()
|
||||
|
||||
|
||||
def _build_default_kline_df(code: str = "sh.600519") -> pd.DataFrame:
|
||||
"""构造 2 根日线(baostock 字符串风格)。"""
|
||||
return pd.DataFrame({
|
||||
"date": ["2024-09-27", "2024-09-30"],
|
||||
"code": [code, code],
|
||||
"open": ["1580.0", "1610.0"],
|
||||
"high": ["1610.0", "1630.0"],
|
||||
"low": ["1575.0", "1605.0"],
|
||||
"close": ["1600.0", "1620.0"],
|
||||
"preclose": ["1570.0", "1600.0"],
|
||||
"volume": ["1000000", "1200000"],
|
||||
"amount": ["1.6e9", "1.94e9"],
|
||||
"adjustflag": ["2", "2"],
|
||||
"turn": ["0.08", "0.10"],
|
||||
"tradestatus": ["1", "1"],
|
||||
"pctChg": ["1.91", "1.25"],
|
||||
"peTTM": ["25.5", "25.8"],
|
||||
"pbMRQ": ["7.5", "7.6"],
|
||||
"psTTM": ["15.2", "15.4"],
|
||||
"pcfNcfTTM": ["20.1", "20.3"],
|
||||
"isST": ["0", "0"],
|
||||
})
|
||||
|
||||
|
||||
def _build_default_profit_df(net_profit_ytd: float = 5.0e10,
|
||||
revenue_ytd: float = 1.0e11,
|
||||
roe_avg: float = 30.0,
|
||||
eps_ttm: float = 40.0,
|
||||
total_share: float = 1.256e9) -> pd.DataFrame:
|
||||
"""构造 query_profit_data 单季报返回(单行)。
|
||||
|
||||
Args 允许测试覆盖,默认茅台 2024 Q3(累计口径:前三季净利 500 亿,营收 1000 亿)。
|
||||
"""
|
||||
return pd.DataFrame({
|
||||
"code": ["sh.600519"],
|
||||
"pubDate": ["2024-10-15"],
|
||||
"statDate": ["2024-09-30"],
|
||||
"roeAvg": [str(roe_avg)],
|
||||
"npMargin": ["50.0"],
|
||||
"gpMargin": ["91.0"],
|
||||
"netProfit": [str(net_profit_ytd)],
|
||||
"epsTTM": [str(eps_ttm)],
|
||||
"MBRevenue": [str(revenue_ytd)],
|
||||
"totalShare": [str(total_share)],
|
||||
"liqaShare": [str(total_share)],
|
||||
})
|
||||
|
||||
|
||||
def _build_default_hs300_df() -> pd.DataFrame:
|
||||
"""构造 query_hs300_stocks 返回(3 只成分股)。"""
|
||||
return pd.DataFrame({
|
||||
"updateDate": ["2024-09-30"] * 3,
|
||||
"code": ["sh.600519", "sh.601318", "sz.000001"],
|
||||
"code_name": ["贵州茅台", "中国平安", "平安银行"],
|
||||
})
|
||||
|
||||
|
||||
def _build_default_stock_basic_df() -> pd.DataFrame:
|
||||
"""构造 query_stock_basic 返回(茅台)。"""
|
||||
return pd.DataFrame({
|
||||
"code": ["sh.600519"],
|
||||
"code_name": ["贵州茅台"],
|
||||
"ipoDate": ["2001-08-27"],
|
||||
"outDate": [""],
|
||||
"type": ["1"],
|
||||
"status": ["1"],
|
||||
})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_baostock():
|
||||
"""构造 baostock MagicMock,覆盖 login/query_history_k_data_plus/query_hs300_stocks/
|
||||
query_profit_data/query_stock_basic 等。
|
||||
|
||||
yield dict,可在外层覆盖任意 query_xxx 的返回值定制。
|
||||
"""
|
||||
bs = MagicMock(name="baostock")
|
||||
# login/logout
|
||||
bs.login.return_value = MagicMock(error_code="0", error_msg="success")
|
||||
bs.logout.return_value = MagicMock(error_code="0", error_msg="success")
|
||||
|
||||
# K 线
|
||||
bs.query_history_k_data_plus.return_value = _FakeResultData(_build_default_kline_df())
|
||||
|
||||
# 成分股
|
||||
bs.query_hs300_stocks.return_value = _FakeResultData(_build_default_hs300_df())
|
||||
bs.query_zz500_stocks.return_value = _FakeResultData(_build_default_hs300_df())
|
||||
bs.query_sz50_stocks.return_value = _FakeResultData(_build_default_hs300_df())
|
||||
|
||||
# 季报(默认茅台 2024 Q3 累计净利 500 亿,营收 1000 亿)
|
||||
bs.query_profit_data.return_value = _FakeResultData(_build_default_profit_df())
|
||||
|
||||
# 证券基本资料
|
||||
bs.query_stock_basic.return_value = _FakeResultData(_build_default_stock_basic_df())
|
||||
|
||||
# 交易日
|
||||
bs.query_trade_dates.return_value = _FakeResultData(pd.DataFrame({
|
||||
"calendar_date": ["2024-09-27", "2024-09-30"],
|
||||
"is_trading_day": ["1", "1"],
|
||||
}))
|
||||
|
||||
# 全部证券
|
||||
bs.query_all_stock.return_value = _FakeResultData(pd.DataFrame({
|
||||
"code": ["sh.600519", "sh.601318"],
|
||||
"tradeStatus": ["1", "1"],
|
||||
"code_name": ["贵州茅台", "中国平安"],
|
||||
}))
|
||||
|
||||
# 复权因子
|
||||
bs.query_adjust_factor.return_value = _FakeResultData(pd.DataFrame({
|
||||
"code": ["sh.600519"],
|
||||
"dividOperateDate": ["2024-06-19"],
|
||||
"foreAdjustFactor": ["0.99"],
|
||||
"backAdjustFactor": ["1.01"],
|
||||
"adjustFactor": ["1.01"],
|
||||
}))
|
||||
|
||||
module = types.ModuleType("baostock")
|
||||
# 把 MagicMock 当作 baostock 模块(sys.modules)
|
||||
sys.modules["baostock"] = bs
|
||||
|
||||
# 同时直接 patch BaostockProvider._import_baostock(更稳定,
|
||||
# 不受 pytest-cov 改变 import 行为影响)
|
||||
from sanguo_portfolio.providers.baostock_provider import BaostockProvider
|
||||
original_import = BaostockProvider._import_baostock
|
||||
BaostockProvider._import_baostock = staticmethod(lambda: bs) # type: ignore[assignment]
|
||||
|
||||
try:
|
||||
yield {
|
||||
"bs": bs,
|
||||
"kline_df": _build_default_kline_df,
|
||||
"profit_df": _build_default_profit_df,
|
||||
"hs300_df": _build_default_hs300_df,
|
||||
"stock_basic_df": _build_default_stock_basic_df,
|
||||
"FakeResultData": _FakeResultData,
|
||||
}
|
||||
finally:
|
||||
BaostockProvider._import_baostock = original_import # type: ignore[assignment]
|
||||
sys.modules.pop("baostock", None)
|
||||
|
||||
@@ -0,0 +1,428 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user