417907e2ea
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)。
1008 lines
43 KiB
Python
1008 lines
43 KiB
Python
"""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})
|