Files
sanguo_vnpy_v2/sanguo_portfolio/providers/local_unified_provider.py
T
claude_dev 417907e2ea perf(data): get_limit_status_batch 查询优化(symbol IN+ROW_NUMBER→UNION ALL+90天范围)
P0 阻塞:VPS 实测 800 只 44s(get_closes_panel 14.75s),03 每日调→全周期 8h+ 不实用。
根因(反馈"逐只查"不准,实证):bars 用 symbol IN + ROW_NUMBER 窗口扫全历史(无日期下界)
= get_closes_panel 注释实证的反模式(symbol IN/OR 全表扫 vs UNION ALL 走复合索引 340×)。
优化:bars 改 UNION ALL per (sym,exc) 纯 SELECT + 90 天下界 + pandas 取最近 2 根,
对齐 get_closes_panel 模式。_isst_batch 本已批量(parquet 向量化,无需改)。口径全不变
(3 limit 测试 pass)。真实提速幅度待 VPS 实测(查询模式已对齐实证 340× 的 get_closes_panel)。
2026-07-29 22:34:46 +08:00

1008 lines
43 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.
"""LocalUnifiedProvider: 读方案A 权威数据层, 零 online, 治幸存者偏差(spec §6)。
数据源(全本地 VPS ``C:\\sanguo_vnpy_v2\\data\\``):
- 日线: ``dbbardata('d')`` raw + ``bs_adjust_factor`` 算前复权(§14.7)
- 成份股: ``constituent_unified`` 并集(治偏差,无 date 时点)
- 估值 pe/pb/ps/pcf: ``valuation_baostock/<year>.parquet``(baostock 权威)
- 市值/股本: ``static/valuation`` akshare parquet(baostock 无市值列)
- 三表: ``static/{balance,income,cashflow}`` akshare parquet(委托 LocalParquetProvider)
零 online: 不 import baostock 调 online。Mac 测试用 sqlite+parquet fixture。
"""
from __future__ import annotations
import logging
import os
import re
import sqlite3
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from typing import Any, Dict, List, Optional, Union
import numpy as np
import pandas as pd
# bullet-trade 可能未装,容错 import DataProvider(照 local_parquet_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"
logger = logging.getLogger(__name__)
# VPS 默认路径(Windows); Mac 测试通过 config["db_path"] / config["data_dir"] 覆盖
_DEFAULT_DB = r"C:\sanguo_vnpy_v2\data\quant_trading.db"
_DEFAULT_DATA_DIR = r"C:\sanguo_vnpy_v2\data"
# jq 后缀 ↔ dbbardata exchange
_JQ_SUFFIX_TO_EXC = {"XSHG": "SSE", "XSHE": "SZSE", "SH": "SSE", "SZ": "SZSE"}
_EXC_TO_JQ_SUFFIX = {"SSE": "XSHG", "SZSE": "XSHE"}
# 日期 / interval 字面量校验(get_closes_panel UNION ALL 注入用, 防 SQL 注入 + 避参数超限)
_DATE_RE = re.compile(r"^\d{4}-\d{2}-\d{2}$")
_INTERVAL_RE = re.compile(r"^[A-Za-z0-9_]+$")
# get_fundamentals_df 并发阈值: 超过则 ThreadPool 并发逐只(本地文件 I/O, 安全)
# 小列表(all_weather 多为单只/[stock])走顺序, 避线程池开销; 策略 02 全市场(5128)走并发
_FUND_POOL_THRESHOLD = 64
def _safe_date_literal(s: str) -> str:
"""``"2022-01-01"`` → ``"'2022-01-01'"`` (SQL 安全字面量, 用于 UNION ALL 注入)。
严格 ``YYYY-MM-DD`` regex 校验; 非法格式 raise ValueError(防御性, 防注入)。
"""
if not isinstance(s, str) or not _DATE_RE.match(s):
raise ValueError(f"Invalid date (expect YYYY-MM-DD): {s!r}")
return f"'{s}'" # regex 已限定为数字+连字符, 注入安全
def _safe_interval_literal(s: str) -> str:
"""``"d"`` → ``"'d'"`` (SQL 安全字面量)。仅允许字母数字下划线。"""
if not isinstance(s, str) or not _INTERVAL_RE.match(s) or len(s) > 16:
raise ValueError(f"Invalid interval: {s!r}")
return f"'{s}'"
def jq_to_dbbardata(jq_code: str) -> tuple[str, str]:
"""``600519.XSHG`` → ``("600519", "SSE")``。
纯 6 位按 6 开头 = sh / 0,3 开头 = sz 推断。
"""
s = (jq_code or "").strip()
if "." not in s:
if len(s) == 6:
return s, ("SSE" if s.startswith("6") else "SZSE")
return s, "SSE"
code, suffix = s.split(".", 1)
return code, _JQ_SUFFIX_TO_EXC.get(suffix.upper(), "SSE")
def dbbardata_to_jq(symbol: str, exchange: str) -> str:
"""``("600519", "SSE")`` → ``"600519.XSHG"``。"""
jq_suffix = _EXC_TO_JQ_SUFFIX.get(str(exchange).upper(), "XSHG")
return f"{symbol}.{jq_suffix}"
def _jq_to_bs_code(jq_code: str) -> str:
"""``600519.XSHG`` → ``"sh.600519"``(``bs_adjust_factor.code`` 格式)。"""
sym, exc = jq_to_dbbardata(jq_code)
prefix = "sh" if exc == "SSE" else "sz"
return f"{prefix}.{sym}"
def _build_qfq_factor(
bs_code: str,
conn: sqlite3.Connection,
dates: pd.Series,
) -> pd.Series:
"""构造每个 date 的前复权因子(asof 语义)。
规则:
- 找 ``<= d`` 的最大 dividOperateDate 的 foreAdjustFactor
- 全部事件 > d(早于所有事件)→ 用最早的 factor
- 全部事件 <= d(晚于所有事件)→ 用最新的 factor
- 无事件 → 全 1.0
``qfq[t] = raw[t] * factor[t]``
"""
rows = conn.execute(
"SELECT dividOperateDate, foreAdjustFactor FROM bs_adjust_factor "
"WHERE code=? ORDER BY dividOperateDate",
(bs_code,),
).fetchall()
dates_ts = pd.to_datetime(dates)
if not rows:
return pd.Series([1.0] * len(dates_ts), index=dates_ts)
ev_dates = pd.to_datetime([r[0] for r in rows])
factors = [float(r[1]) for r in rows]
out: List[float] = []
for d in dates_ts:
mask = ev_dates <= d
if mask.any():
# <= d 的最大事件 = 最后一个 True
idx = int(np.where(mask)[0][-1])
else:
# 全部 > d → 用最早(第一个)
idx = 0
out.append(factors[idx])
return pd.Series(out, index=dates_ts)
class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
"""占位类(Task0 骨架; Task1-4 填充方法)。
读方案A 权威数据层, 零 online, 治幸存者偏差(spec §6 使用层)。
"""
name: str = "sanguo_local_unified"
requires_live_data: bool = False
def __init__(self, config: Optional[Dict[str, Any]] = None) -> None:
cfg = config or {}
self.db_path: str = cfg.get("db_path", _DEFAULT_DB)
self.data_dir: str = cfg.get("data_dir", _DEFAULT_DATA_DIR)
self._conn: Optional[sqlite3.Connection] = None
self._val_bs_cache: Dict[int, pd.DataFrame] = {} # year -> valuation_baostock
self._lpp_helper: Any = None
def _connect(self) -> sqlite3.Connection:
"""惰性连接 dbbardata sqlite(单连接复用)。"""
if self._conn is None:
self._conn = sqlite3.connect(self.db_path, timeout=30)
self._conn.execute("PRAGMA busy_timeout = 30000")
return self._conn
@staticmethod
def _to_date_str(value: Optional[Union[str, datetime]]) -> Optional[str]:
if value is None:
return None
if isinstance(value, str):
return value[:10]
try:
return value.strftime("%Y-%m-%d")
except AttributeError:
return str(value)[:10]
# ==================== get_price ====================
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 = "raw",
count: Optional[int] = None,
panel: bool = True,
fill_paused: bool = True,
**kwargs: Any,
) -> pd.DataFrame:
"""读 ``dbbardata('d')`` raw 日线,按需 ``bs_adjust_factor`` 算前复权。
策略契约(all_weather 实证):
- ``panel=False`` 返长表含 ``time`` + ``code`` 列(供 pivot)
- ``fields`` 里缺失列(如 ``high_limit``)补 NaN(降级)
- ``frequency`` 非 day/1d/d → 返空 DataFrame(1m 数据层无)
"""
freq = str(frequency or "").lower()
if freq not in ("daily", "day", "1d", "d"):
return pd.DataFrame()
secs: List[str] = [security] if isinstance(security, str) else list(security or [])
if not secs:
return pd.DataFrame()
conn = self._connect()
start_str = self._to_date_str(start_date) or "1990-01-01"
end_str = self._to_date_str(end_date) or datetime.now().strftime("%Y-%m-%d")
frames: Dict[str, pd.DataFrame] = {}
for jq_code in secs:
sym, exc = jq_to_dbbardata(jq_code)
# substr(datetime,1,10) 取日期部分比 — datetime 列混合格式(有只日期有带时间),
# 纯字符串比 "2024-09-25" < "2024-09-25 00:00:00" 会漏边界行; 比日期(YYYY-MM-DD)规避
q = (
"SELECT datetime, open_price, high_price, low_price, close_price, "
"volume, turnover FROM dbbardata WHERE symbol=? AND exchange=? "
"AND interval='d' AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? "
"ORDER BY datetime"
)
df = pd.read_sql(q, conn, params=(sym, exc, start_str, end_str))
if df.empty:
frames[jq_code] = df
continue
# dbbardata datetime 混合格式(有的 "2024-09-26" 有的 "2024-09-26 00:00:00",
# 不同 schtask/迁移写入);pandas 2.3 严格模式要 format="mixed"
df["datetime"] = pd.to_datetime(df["datetime"], format="mixed")
df = df.set_index("datetime")
df.index.name = None
if count:
df = df.tail(count)
# 前复权
if fq in ("qfq", "pre", "前复权") and not df.empty:
factor = _build_qfq_factor(_jq_to_bs_code(jq_code), conn, pd.Series(df.index))
for col in ("open_price", "high_price", "low_price", "close_price"):
df[col] = df[col].values * factor.values
# jq 风格字段重命名
df = df.rename(columns={
"open_price": "open", "high_price": "high",
"low_price": "low", "close_price": "close",
})
# 缺失字段补默认: paused=False(避免 bool(NaN)=True 被 bullet_trade
# get_current_data 误判停牌→订单 cancel); high_limit/low_limit 按 close±10% 估
# (与 get_current_tick 同口径, 精确涨跌停/ST/创业科创规则 v2); 其他补 NaN
if fields:
for f in fields:
if f in df.columns:
continue
if f == "paused":
df[f] = False
elif f == "high_limit" and "close" in df.columns:
df[f] = (df["close"] * 1.1).round(2)
elif f == "low_limit" and "close" in df.columns:
df[f] = (df["close"] * 0.9).round(2)
else:
df[f] = float("nan")
df = df[[f for f in fields if f in df.columns]]
frames[jq_code] = df
if not frames or all(f.empty for f in frames.values()):
return pd.DataFrame()
if not panel:
parts: List[pd.DataFrame] = []
for jq_code, df in frames.items():
if df.empty:
continue
d = df.reset_index()
# index.name=None 时 reset_index 出 'index' 列; 统一改名 'time'
if "index" in d.columns and "time" not in d.columns:
d = d.rename(columns={"index": "time"})
elif "datetime" in d.columns:
d = d.rename(columns={"datetime": "time"})
d.insert(0, "code", jq_code)
parts.append(d)
return pd.concat(parts, ignore_index=True) if parts else pd.DataFrame()
if len(frames) == 1:
return next(iter(frames.values()))
return pd.concat(frames, axis=1)
# ==================== get_closes_panel (批量行情,G5 Phase1) ====================
def get_closes_panel(
self,
symbols: List[str],
start: Union[str, datetime],
end: Union[str, datetime],
interval: str = "d",
fq: str = "raw",
) -> pd.DataFrame:
"""批量取多只股票 close,返宽表 (index=datetime, columns=symbol, values=close)。
生产化增强(vs commit c8f26be):
- **分批 chunk=400**: 单 OR 大查询 → 多批 concat (5128 只撞老版 SQLite 999 上限)
- **``fq='qfq'`` 批量前复权**: 一次性 ``WHERE code IN (...)`` 查所有 symbol 的
``bs_adjust_factor`` + ``merge_asof`` 向量化(非逐只 ``_build_qfq_factor`` N 次查询),
语义与 ``get_price(fq='qfq')`` / ``_build_qfq_factor`` 严格一致(asof backward +
早于首事件 → 用最早事件 factor)。
Args:
symbols: jq 风格代码列表 (``["600519.XSHG", "000001.XSHE"]``) 或纯 6 位
start/end: 日期 ``YYYY-MM-DD`` (或 datetime,取 date 部分)
interval: ``"d"``=日线 (dbbardata ``interval`` 字段; ``"15m"`` 等 v2 扩展)
fq: ``"raw"``(默认, 不复权) | ``"qfq"``/``"pre"``/``"前复权"`` (前复权)
镜像 ``get_price`` 的 fq 取值; 默认 "raw" 向后兼容。
Returns:
DataFrame, ``index=datetime`` (升序), ``columns=symbols``(按输入顺序),
``values=close_price``(fq='qfq' 时为复权 close)。缺失股票 → 该列全 NaN;
重复行 dedup (keep last)。
空列表 → 空 DataFrame(index=DatetimeIndex,columns=[])。
"""
if not symbols:
return pd.DataFrame(index=pd.DatetimeIndex([]))
# jq 代码 → (db_symbol, db_exchange); 保留 input_code 作列名
# (symbol 不唯一: 932000 中证2000 / 北交所 920xxx 等, 必须按 exchange 配对)
pairs: List[tuple[str, str, str]] = [] # (input_code, db_symbol, db_exchange)
for s in symbols:
sym, exc = jq_to_dbbardata(str(s))
pairs.append((str(s), sym, exc))
conn = self._connect()
start_str = self._to_date_str(start) or "1990-01-01"
end_str = self._to_date_str(end) or datetime.now().strftime("%Y-%m-%d")
# 分批查询 dbbardata: 每 400 symbol 一批, 用 UNION ALL per symbol。
#
# 为何 UNION ALL 而非 OR chain: SQLite 对 ``(sym=? AND exc=?) OR ...`` 大 OR 链
# 无法用 ``dbbardata_symbol_exchange_interval_datetime`` 复合索引(实证 5 sym
# OR 用 idx_dbbardata_interval 全表扫 10s vs UNION ALL 0.03s, 340× 提速)。
# 每个 UNION ALL 子查询独立使用复合索引(SEEK by sym+exc+interval+dt 范围)。
#
# 参数: 2 per subquery (sym, exc) × 400 = 800 < 999 (老版 SQLite 上限);
# interval / start / end 作 *已校验字面量*注入(YYYY-MM-DD regex + alnum)避免每
# subquery 5 参数(会超 999 上限)。日期格式校验在下方 _safe_lit 完成。
start_lit = _safe_date_literal(start_str)
end_lit = _safe_date_literal(end_str)
interval_lit = _safe_interval_literal(interval)
CHUNK_SIZE = 400
chunk_frames: List[pd.DataFrame] = []
for i in range(0, len(pairs), CHUNK_SIZE):
chunk = pairs[i:i + CHUNK_SIZE]
# 每个 subquery 只 2 个参数 (sym, exc); UNION ALL 拼接
sub_template = (
"SELECT datetime, symbol, close_price FROM dbbardata "
f"WHERE symbol=? AND exchange=? AND interval={interval_lit} "
f"AND substr(datetime,1,10)>={start_lit} "
f"AND substr(datetime,1,10)<={end_lit}"
)
q = " UNION ALL ".join([sub_template] * len(chunk))
params: List[Any] = []
for _inp, sym, exc in chunk:
params.extend([sym, exc])
chunk_frames.append(pd.read_sql(q, conn, params=params))
df = pd.concat(chunk_frames, ignore_index=True) if chunk_frames else pd.DataFrame()
if df.empty:
# 全缺失: 返宽表骨架(全 NaN 列, 空 DatetimeIndex) — 与逐只 get_price 行为一致
return pd.DataFrame(
{inp: pd.Series(dtype=float) for inp in symbols},
index=pd.DatetimeIndex([]),
)
# 混合格式 datetime(与 get_price 一致: dbbardata 列混 "2024-09-26" 与
# "2024-09-26 00:00:00",pandas 2.3 严格模式要 format="mixed")
df["datetime"] = pd.to_datetime(df["datetime"], format="mixed")
# db_symbol → input_code 还原(列名回输入 jq code, 与策略其他接口一致)
sym_to_input: Dict[str, str] = {}
for inp, sym, _exc in pairs:
sym_to_input.setdefault(sym, inp)
df["symbol"] = df["symbol"].map(sym_to_input).fillna(df["symbol"])
# dedup: 同 (datetime, symbol) 重复行取最后一条(增量合并容错)
df = df.drop_duplicates(subset=["datetime", "symbol"], keep="last")
# pivot 宽表
wide = df.pivot(index="datetime", columns="symbol", values="close_price")
# 补缺失 symbol 列(全 NaN), 按 input 顺序对齐 columns
# (批量 concat 一次性补, 避免 5128 次 frame.insert 的 PerformanceWarning 碎片化)
missing_cols = [inp for inp in symbols if inp not in wide.columns]
if missing_cols:
wide = pd.concat(
[wide, pd.DataFrame(float("nan"), index=wide.index, columns=missing_cols)],
axis=1,
)
wide = wide[symbols]
# 前复权(批量, 向量化): 与 get_price(fq='qfq') 数值一致, 单次 WHERE code IN + merge_asof
if fq in ("qfq", "pre", "前复权"):
wide = self._apply_qfq_batch(wide, pairs, conn)
# 升序 + 清 index name(与 get_price 一致)
wide = wide.sort_index()
wide.index.name = None
return wide
def _apply_qfq_batch(
self,
wide: pd.DataFrame,
pairs: List[tuple[str, str, str]],
conn: sqlite3.Connection,
) -> pd.DataFrame:
"""批量前复权(merge_asof 向量化, 替代逐只 ``_build_qfq_factor`` N 次查询)。
语义与 ``_build_qfq_factor`` 严格一致:
- 每个 date 用 ``<= d`` 的最大 dividOperateDate 的 foreAdjustFactor(asof backward)
- 早于首事件 → 用最早事件 factor(对应 ``_build_qfq_factor`` 全部 > d → idx=0)
- 无事件 symbol → factor=1.0 → 列不变
实现(高效关键, 避免 N 次 query):
1. 单次(或 400-chunked) ``WHERE code IN (...)`` 查所有 symbol 的全部事件
2. 按 input_code groupby, 每组 ``pd.merge_asof(direction='backward')`` onto
宽表 date index — pandas 内部 searchsorted 向量化, 远快于 Python 循环
3. ``qfq[t] = raw[t] * factor[t]`` 逐列乘
"""
# input_code → bs_code 映射(只对宽表实际存在的列查, dropna 后全 NaN 列不浪费 query)
bs_to_input: Dict[str, str] = {}
for inp, sym, exc in pairs:
if inp in wide.columns:
prefix = "sh" if exc == "SSE" else "sz"
bs_to_input.setdefault(f"{prefix}.{sym}", inp)
if not bs_to_input:
return wide
# 批量查 bs_adjust_factor(chunk=400, 参数 < 999)
all_bs_codes = list(bs_to_input.keys())
factor_rows: List[tuple] = []
for i in range(0, len(all_bs_codes), 400):
chunk = all_bs_codes[i:i + 400]
placeholders = ",".join(["?"] * len(chunk))
q = (
"SELECT code, dividOperateDate, foreAdjustFactor "
f"FROM bs_adjust_factor WHERE code IN ({placeholders})"
)
factor_rows.extend(conn.execute(q, chunk).fetchall())
if not factor_rows:
return wide # 无任何事件 → factor 全 1.0 → 等于 raw
events = pd.DataFrame(
factor_rows, columns=["bs_code", "event_date", "factor"]
)
events["event_date"] = pd.to_datetime(events["event_date"])
events["input_code"] = events["bs_code"].map(bs_to_input)
events = events.dropna(subset=["input_code"])
if events.empty:
return wide
events = events.sort_values(["input_code", "event_date"]).reset_index(drop=True)
# merge_asof 要求 left 升序; 这里先 sort_index, 末尾调用方再 sort_index 是 no-op
wide = wide.sort_index()
date_grid = pd.DataFrame({"date": wide.index})
# 按 input_code 分组, 每组一次 merge_asof
for inp_code, grp in events.groupby("input_code"):
if inp_code not in wide.columns:
continue
g = grp[["event_date", "factor"]].sort_values("event_date")
# 同 event_date 多事件: drop_duplicates keep last(后事件覆盖前事件)
g = g.drop_duplicates(subset=["event_date"], keep="last")
merged = pd.merge_asof(
date_grid,
g,
left_on="date",
right_on="event_date",
direction="backward",
)
factor_series = merged["factor"]
if factor_series.isna().any():
# 早于首事件 → 用最早事件 factor(对齐 _build_qfq_factor idx=0)
factor_series = factor_series.fillna(g["factor"].iloc[0])
# 防御: 若 inp_code 在 wide.columns 有重复(caller 传 dup symbols),
# wide[inp_code] 返 DataFrame 而非 Series, 用 iloc[:, 0] 取首个一致列
existing = wide[inp_code]
if isinstance(existing, pd.DataFrame):
existing = existing.iloc[:, 0]
wide[inp_code] = existing.values * factor_series.values
return wide
# ==================== get_index_stocks (constituent_unified 并集) ====================
def get_index_stocks(
self,
index_symbol: str,
date: Optional[Union[str, datetime]] = None,
) -> List[str]:
"""读 ``constituent_unified`` 并集(``in_current=1 OR was_removed=1``)。
⚠️ 并集模型: 表无 date 列,治"纯当前"幸存者偏差(含已退市/被踢),
但有轻微前视(使用说明标注)。``date`` 参数忽略(无时点数据)。
"""
idx = (
index_symbol.split(".")[0]
if "." in str(index_symbol) else str(index_symbol)
)
conn = self._connect()
try:
rows = conn.execute(
"SELECT code FROM constituent_unified WHERE index_code=? "
"AND (in_current=1 OR was_removed=1)",
(idx,),
).fetchall()
except sqlite3.Error as exc:
logger.warning("constituent_unified 查询失败 %s: %s", idx, exc)
return []
out: List[str] = []
for (code,) in rows:
code_str = str(code).strip()
if len(code_str) != 6 or not code_str.isdigit():
continue
exc_name = "SSE" if code_str.startswith("6") else "SZSE"
out.append(dbbardata_to_jq(code_str, exc_name))
return out
def get_constituent(
self,
index: str,
date: Optional[Union[str, datetime]] = None,
) -> List[str]:
"""spec §6 语义别名 = ``get_index_stocks``。"""
return self.get_index_stocks(index, date)
# ==================== get_fundamentals_df ====================
def _get_lpp_helper(self) -> Any:
"""惰性创建 LocalParquetProvider 委托读 static akshare(三表/市值)。DRY。"""
if self._lpp_helper is None:
from .local_parquet_provider import LocalParquetProvider
self._lpp_helper = LocalParquetProvider({"data_dir": self.data_dir})
return self._lpp_helper
def _read_valuation_baostock(self, year: int) -> pd.DataFrame:
"""读 ``valuation_baostock/<year>.parquet``(baostock 权威: pe/pb/ps/pcf)。"""
if year in self._val_bs_cache:
return self._val_bs_cache[year]
p = os.path.join(self.data_dir, "valuation_baostock", f"{year}.parquet")
if not os.path.exists(p):
self._val_bs_cache[year] = pd.DataFrame()
return pd.DataFrame()
try:
df = pd.read_parquet(p)
self._val_bs_cache[year] = df
return df
except Exception as exc:
logger.warning("读 valuation_baostock/%s 失败: %s", year, exc)
self._val_bs_cache[year] = pd.DataFrame()
return pd.DataFrame()
# fields → 源表依赖(短路: 未请求字段其源表不读, 省 parquet 读 + calc_roic)
# akshare_val: static/valuation(市值 + akshare pe/pb fallback)
# bs_val: valuation_baostock(baostock 覆盖 pe/pb/ps/pcf)
# income: static/income(eps/yoy/net_profit→roe/roa/roic/margin)
# balance: static/balance(total_liability/权益/资产→roe/roa/roic)
# fa: static/financial_abstract(gross_profit_margin)
_FIELDS_AKSHARE_VAL = frozenset({
"market_cap", "circulating_market_cap",
"pe_ratio", "pb_ratio", "ps_ratio", "pcf_ratio",
})
_FIELDS_BS_VAL = frozenset({"pe_ratio", "pb_ratio", "ps_ratio", "pcf_ratio"})
_FIELDS_INCOME = frozenset({
"eps", "inc_revenue_year_on_year", "inc_operation_profit_year_on_year",
"inc_total_revenue_year_on_year", "net_profit_margin", "roe", "roa", "roic",
})
_FIELDS_BALANCE = frozenset({
"total_liability", "total_sheet_owner_equities", "retained_profit",
"roe", "roa", "roic",
})
_FIELDS_FA = frozenset({"gross_profit_margin"})
@staticmethod
def _fields_to_need(fields: Optional[List[str]]) -> Dict[str, bool]:
"""fields 列表 → 各源表是否需读。``fields=None`` 时调用方不调本方法(走全读)。"""
if not fields:
return {"akshare_val": True, "bs_val": True, "income": True,
"balance": True, "fa": True}
fs = set(fields)
cls = LocalUnifiedProvider
return {
"akshare_val": bool(fs & cls._FIELDS_AKSHARE_VAL),
"bs_val": bool(fs & cls._FIELDS_BS_VAL),
"income": bool(fs & cls._FIELDS_INCOME),
"balance": bool(fs & cls._FIELDS_BALANCE),
"fa": bool(fs & cls._FIELDS_FA),
}
def get_fundamentals_df(
self,
stocks: List[str],
date: Optional[Union[str, datetime]] = None,
fields: Optional[List[str]] = None,
) -> pd.DataFrame:
"""合并多股 fundamentals, 列对齐 ``_FUNDAMENTAL_COLUMNS``。
数据源路由:
- ``pe/pb/ps/pcf`` ← ``valuation_baostock/<year>.parquet`` (baostock 权威, 覆盖 akshare)
- ``market_cap`` / ``circulating_market_cap`` ← ``static/valuation`` akshare
(baostock valuation 无市值列)
- 三表(eps/yoy/total_liability/...) ← ``static/{income,balance}`` akshare
(委托 LocalParquetProvider 读, DRY)
性能(2026-07-28 批量优化, 解锁策略 02 全市场 5128 只选股):
- ``fields=`` 按需短路: 只读请求字段依赖的源表(策略 02 只要 market_cap+eps
→ 跳 balance/financial_abstract/roic, 省一半 parquet 读 + 跳 calc_roic)。
- ``len(stocks) > _FUND_POOL_THRESHOLD`` 时 ThreadPool 并发逐只读(本地文件 I/O,
非 baostock 网络 → 并发安全, 不触"baostock 不并发"铁律)。
``fields=None`` 全列(向后兼容, all_weather/small_cap 现有调用零改动)。
"""
from .local_parquet_provider import _FUNDAMENTAL_COLUMNS, jq_to_file_code
if not stocks:
return pd.DataFrame(columns=_FUNDAMENTAL_COLUMNS)
date_str = self._to_date_str(date) or datetime.now().strftime("%Y-%m-%d")
need = self._fields_to_need(fields) if fields else None
def _one(jq_code: str) -> Dict[str, Any]:
return self._build_fundamental_row(jq_code, date_str, need)
if len(stocks) <= _FUND_POOL_THRESHOLD:
rows: List[Dict[str, Any]] = [_one(s) for s in stocks]
else:
workers = min(8, os.cpu_count() or 4)
with ThreadPoolExecutor(max_workers=workers) as ex:
rows = list(ex.map(_one, stocks))
df = pd.DataFrame(rows, columns=_FUNDAMENTAL_COLUMNS)
if "code" in df.columns:
df = df.set_index("code", drop=False)
if fields:
keep = ["code"] + [f for f in fields if f in df.columns]
df = df[keep]
return df
def get_value_metrics(
self,
stock: str,
date: Union[str, datetime],
) -> Optional[Dict[str, Any]]:
"""委托 LocalParquetProvider 读三表多期 + NOTICE_DATE 过滤(供 ValueSelectionStrategy)。
LocalParquetProvider 实现见 ``local_parquet_provider.get_value_metrics``;
本类只是把职责转发给已有的 ``_lpp_helper``(DRY, 不复制字段映射逻辑)。
"""
return self._get_lpp_helper().get_value_metrics(stock, date)
def get_value_metrics_batch(
self,
stocks: List[str],
date: Union[str, datetime],
) -> Dict[str, Optional[Dict[str, Any]]]:
"""批量 get_value_metrics(ThreadPool 并发逐只; 策略01 价值精选提速)。
逐只多期逻辑(ROE/FCF/流动比率/yoy, NOTICE_DATE<=date 前视过滤)委托
LocalParquetProvider.get_value_metrics 不变; 只把 N 次调用并发化。
多期聚合不宜向量化(300 只够用, 更深优化 YAGNI)。返回 {stock: metrics_or_None}。
"""
if not stocks:
return {}
lpp = self._get_lpp_helper()
def _one(jq_code: str):
try:
return jq_code, lpp.get_value_metrics(jq_code, date)
except Exception as exc:
logger.debug("get_value_metrics(%s) 失败: %s", jq_code, exc)
return jq_code, None
if len(stocks) <= _FUND_POOL_THRESHOLD:
return dict(_one(s) for s in stocks)
workers = min(8, os.cpu_count() or 4)
with ThreadPoolExecutor(max_workers=workers) as ex:
return dict(ex.map(_one, stocks))
def _build_fundamental_row(
self, jq_code: str, date_str: str,
need: Optional[Dict[str, bool]] = None,
) -> Dict[str, Any]:
"""单股 fundamentals 行: baostock 估值覆盖 akshare pe/pb, 市值+三表用 LocalParquetProvider。
need: 源表短路字典(None=全读); 由 get_fundamentals_df(fields=) 构造。
"""
from .local_parquet_provider import (
jq_to_file_code, _to_float, _or_nan, _pct_to_decimal,
)
from ..factors.valuation import to_yi
sym, _ = jq_to_dbbardata(jq_code)
fc = jq_to_file_code(jq_code)
# 先委托 LocalParquetProvider 取整行(akshare 全套, 按 need 短路读哪些表)
lpp = self._get_lpp_helper()
ak_row = lpp._build_fundamental_row(jq_code, date_str, need)
# 复制 akshare 行(市值/eps/三表/yoy/...), 然后用 baostock 覆盖 pe/pb/ps/pcf
row: Dict[str, Any] = dict(ak_row)
row["code"] = jq_code
# baostock 估值覆盖 pe/pb/ps/pcf(按 need['bs_val'] 短路; 未请求则跳)
if need is not None and not need.get("bs_val", True):
return row
year = int(date_str[:4])
vbs = self._read_valuation_baostock(year)
vrow = None
if not vbs.empty:
sub = vbs[
(vbs["symbol"].astype(str) == sym)
& (vbs["date"].astype(str) <= date_str)
]
vrow = sub.iloc[-1] if not sub.empty else None
def gbs(k: str) -> Optional[float]:
return _to_float(vrow.get(k)) if vrow is not None else None
bs_pe = gbs("peTTM")
bs_pb = gbs("pbMRQ")
bs_ps = gbs("psTTM")
bs_pcf = gbs("pcfNcfTTM")
# baostock 有该日数据则覆盖; 否则保留 akshare pe/pb
if bs_pe is not None:
row["pe_ratio"] = float(bs_pe)
if bs_pb is not None:
row["pb_ratio"] = float(bs_pb)
if bs_ps is not None:
row["ps_ratio"] = float(bs_ps)
if bs_pcf is not None:
row["pcf_ratio"] = float(bs_pcf)
return row
# ==================== 辅助方法 ====================
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]:
"""读 dbbardata 蓝筹 600519 distinct datetime 取交易日(全市场交易日一致)。"""
conn = self._connect()
rows = conn.execute(
"SELECT DISTINCT datetime FROM dbbardata WHERE symbol='600519' "
"AND exchange='SSE' AND interval='d' ORDER BY datetime"
).fetchall()
days = [pd.Timestamp(r[0]).to_pydatetime() for r in rows if r[0]]
start_ts = pd.Timestamp(start_date) if start_date else None
end_ts = pd.Timestamp(end_date) if end_date else None
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
def get_security_info(self, security: str) -> Dict[str, Any]:
"""读 dbbardata min/max datetime → start/end_date;display_name 从 constituent_unified。"""
sym, exc = jq_to_dbbardata(security)
conn = self._connect()
row = conn.execute(
"SELECT MIN(datetime), MAX(datetime) FROM dbbardata "
"WHERE symbol=? AND exchange=? AND interval='d'",
(sym, exc),
).fetchone()
start_dt = row[0][:10] if row and row[0] else None
end_dt = row[1][:10] if row and row[1] else None
# display_name 查 constituent_unified
name = security
try:
nrow = conn.execute(
"SELECT code_name FROM constituent_unified WHERE code=? LIMIT 1",
(sym,),
).fetchone()
if nrow and nrow[0]:
name = str(nrow[0])
except sqlite3.Error:
pass
return {
"code": security,
"display_name": name,
"name": name,
"start_date": start_dt,
"end_date": end_dt,
"type": "stock",
}
def get_security_info_batch(
self,
securities: List[str],
date: Optional[Union[str, datetime]] = None,
) -> Dict[str, Dict[str, Any]]:
"""批量 get_security_info: **2 条 SQL 替 N×2 逐只查询**(filters ST/次新通病提速)。
filters.filter_st_stock / filter_new_stock 逐只调 get_security_info, 每只 2 条
SQL(dbbardata min/max + constituent_unified name); 03 每日 10 行业×几百只、02/01
调仓千次 → 万次查询。本方法一次查全:
- dbbardata: ``symbol IN (...) GROUP BY symbol, exchange`` 拿 min/max(走索引)
- constituent_unified: ``code IN (...)`` 拿 code_name
返回 {jq_code: {code, display_name, name, start_date, end_date, type}}, 与逐只逐字一致。
"""
if not securities:
return {}
pairs = [(s, jq_to_dbbardata(s)) for s in securities]
syms = list({sym for sym, _ in (p[1] for p in pairs)})
conn = self._connect()
# 1. dbbardata min/max per (symbol, exchange); chunk 防 >999 参数
minmax: Dict[tuple, tuple] = {}
for i in range(0, len(syms), 400):
chunk = syms[i:i + 400]
ph = ",".join("?" * len(chunk))
for r in conn.execute(
f"SELECT symbol, exchange, MIN(datetime), MAX(datetime) FROM dbbardata "
f"WHERE interval='d' AND symbol IN ({ph}) GROUP BY symbol, exchange",
chunk,
):
minmax[(r[0], r[1])] = (
r[2][:10] if r[2] else None,
r[3][:10] if r[3] else None,
)
# 2. constituent_unified code_name(碰撞股名; 缺则回退 jq_code)
names: Dict[str, str] = {}
for i in range(0, len(syms), 400):
chunk = syms[i:i + 400]
ph = ",".join("?" * len(chunk))
try:
for r in conn.execute(
f"SELECT code, code_name FROM constituent_unified WHERE code IN ({ph})",
chunk,
):
if r[1]:
names[r[0]] = str(r[1])
except sqlite3.Error:
pass
out: Dict[str, Dict[str, Any]] = {}
for jq_code, (sym, exc) in pairs:
start_dt, end_dt = minmax.get((sym, exc), (None, None))
name = names.get(sym, jq_code)
out[jq_code] = {
"code": jq_code,
"display_name": name,
"name": name,
"start_date": start_dt,
"end_date": end_dt,
"type": "stock",
}
return out
@staticmethod
def _limit_pct(sym: str, is_st: bool) -> float:
"""涨跌停幅度(%): ST5 / 北交30 / 科创·创业20 / 主板10。"""
if is_st:
return 5.0
if sym.startswith("920") or sym[:1] in ("4", "8"): # 北交所
return 30.0
if sym.startswith("68") or sym.startswith("30"): # 科创 / 创业
return 20.0
return 10.0 # 主板
def _isst_batch(self, syms, year: int, date_str: str) -> set:
"""valuation_baostock[year] 取各 sym 最新(date<=T)的 isST=1 集合(历史 ST 感知)。"""
vbs = self._read_valuation_baostock(year)
if vbs.empty or "isST" not in vbs.columns:
return set()
try:
sub = vbs[
vbs["symbol"].astype(str).isin(syms)
& (vbs["date"].astype(str) <= date_str)
]
sub = sub.sort_values("date").drop_duplicates("symbol", keep="last")
return set(sub.loc[sub["isST"].astype(int) == 1, "symbol"].astype(str))
except Exception as exc:
logger.debug("_isst_batch 失败: %s", exc)
return set()
def get_limit_status_batch(
self,
codes: List[str],
date: Union[str, datetime],
) -> Dict[str, Optional[Dict[str, bool]]]:
"""批量回测当日涨跌停/停牌状态(修 filter_limitup/limitdown/paused 回测失效)。
dbbardata 无 high_limit 列 → 用 ``high_limit=round(prev_close×(1+幅度),2)`` 精确算
(pctChg 阈值在高价股边界失真, 故不用)。幅度板块感知(主板10/创业·科创20/北交30)
+ 历史 ST5%(valuation_baostock.isST, 非当前名)。停牌=当日 volume==0。
返回 {code: {is_limit_up, is_limit_down, is_paused} | None(无 bar)}。
"""
if not codes:
return {}
date_str = self._to_date_str(date)
if not date_str:
return {c: None for c in codes}
pairs = [(c, jq_to_dbbardata(c)) for c in codes]
syms = list({sym for sym, _ in (p[1] for p in pairs)})
conn = self._connect()
# 1. 最近 2 根日线(T + T-1)per (sym, exc)。
# UNION ALL per (sym,exc) 走复合索引——原 symbol IN + ROW_NUMBER 窗口扫全历史
# (无日期下界)慢 44s;改 UNION ALL + ORDER BY DESC LIMIT 2 同 get_closes_panel 模式
# (该模式实证 vs symbol IN/OR 全表扫 340× 提速)+ 90 天下界覆盖长假/停牌取 T-1。
from datetime import timedelta
start_lim = (datetime.strptime(date_str, "%Y-%m-%d") - timedelta(days=90)).strftime("%Y-%m-%d")
start_lit = _safe_date_literal(start_lim)
end_lit = _safe_date_literal(date_str)
interval_lit = _safe_interval_literal("d")
# 纯 SELECT UNION ALL(子查询带 ORDER BY/LIMIT 触发 SQLite compound 限制);
# 用 90 天下界限定范围(全历史 → 近 90 天), 走复合索引, pandas 端取最近 2 根。
sub = (
"SELECT symbol, exchange, close_price, high_price, low_price, volume, datetime "
f"FROM dbbardata WHERE symbol=? AND exchange=? AND interval={interval_lit} "
f"AND substr(datetime,1,10)>={start_lit} AND substr(datetime,1,10)<={end_lit}"
)
bars: Dict[tuple, list] = {}
sym_exc = list({p[1] for p in pairs}) # 去重 (sym,exc)
for i in range(0, len(sym_exc), 400):
chunk = sym_exc[i:i + 400]
q = " UNION ALL ".join([sub] * len(chunk))
params: List[Any] = []
for sym, exc in chunk:
params.extend([sym, exc])
for r in conn.execute(q, params):
bars.setdefault((r[0], r[1]), []).append(r)
for k in bars: # datetime DESC 取最近 2 根: 当日 T(blist[0]), T-1(blist[1])
bars[k].sort(key=lambda x: x[6], reverse=True)
bars[k] = bars[k][:2]
# 2. 历史 ST 集合 → 5% 幅度
st_set = self._isst_batch(syms, int(date_str[:4]), date_str)
out: Dict[str, Optional[Dict[str, bool]]] = {}
for jq_code, (sym, exc) in pairs:
blist = bars.get((sym, exc))
if not blist:
out[jq_code] = None
continue
t_bar = blist[0] # rn=1 = 当日
close_t = t_bar[2]
vol_t = t_bar[5]
is_paused = (vol_t is None or vol_t == 0)
prev_close = blist[1][2] if len(blist) >= 2 else None
is_up = is_down = False
if prev_close and close_t is not None:
pct = self._limit_pct(sym, sym in st_set)
high_limit = round(prev_close * (1 + pct / 100), 2)
low_limit = round(prev_close * (1 - pct / 100), 2)
is_up = close_t >= high_limit
is_down = close_t <= low_limit
out[jq_code] = {
"is_limit_up": is_up,
"is_limit_down": is_down,
"is_paused": is_paused,
}
return out
def get_current_tick(self, security: str) -> Optional[Dict[str, Any]]:
"""dbbardata 最近 close + 高低涨停 ±10%(简化,ST/创业/科创精确规则 v2)。"""
sym, exc = jq_to_dbbardata(security)
conn = self._connect()
row = conn.execute(
"SELECT close_price FROM dbbardata WHERE symbol=? AND exchange=? "
"AND interval='d' ORDER BY datetime DESC LIMIT 1",
(sym, exc),
).fetchone()
if not row or row[0] is None:
return None
close = float(row[0])
return {
"code": security,
"current_price": close,
"close": close,
"high_limit": round(close * 1.1, 2),
"low_limit": round(close * 0.9, 2),
}
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]]:
"""读 bs_adjust_factor → 除权事件列表。"""
bs_code = _jq_to_bs_code(security)
conn = self._connect()
start_str = self._to_date_str(start_date) or "1900-01-01"
end_str = self._to_date_str(end_date) or "2099-12-31"
try:
rows = conn.execute(
"SELECT dividOperateDate, foreAdjustFactor, backAdjustFactor, adjustFactor "
"FROM bs_adjust_factor WHERE code=? AND dividOperateDate>=? "
"AND dividOperateDate<=? ORDER BY dividOperateDate",
(bs_code, start_str, end_str),
).fetchall()
except sqlite3.Error:
return []
return [
{
"code": security,
"date": r[0],
"foreAdjustFactor": float(r[1]) if r[1] is not None else None,
"backAdjustFactor": float(r[2]) if r[2] is not None else None,
"adjustFactor": float(r[3]) if r[3] is not None else None,
}
for r in rows
]
def get_all_securities(
self, types: Optional[List[str]] = None,
) -> pd.DataFrame:
"""读 dbbardata distinct symbol → DataFrame[code, display_name]。"""
conn = self._connect()
rows = conn.execute(
"SELECT DISTINCT symbol, exchange FROM dbbardata WHERE interval='d'"
).fetchall()
codes = [dbbardata_to_jq(sym, exc) for sym, exc in rows if sym and exc]
return pd.DataFrame({"code": codes, "display_name": codes})