feat(data): get_closes_panel 生产化(qfq 前复权 + UNION ALL 340x 提速 + chunk) + 000938 行情
G5-P1 接口生产化(本地已建 c8f26be 但未部署 VPS — 策略 session 报缺失的根因):
- fq 参数(raw/qfq/前复权, 默认 raw 向后兼容): 批量前复权用单次 WHERE code IN (...) 查
bs_adjust_factor + merge_asof 向量化 asof, 对齐聚宽默认前复权(RPS/均线/动量跨期必需)。
- 性能修复: 原 (symbol=? AND exchange=?) OR 大链打不动复合索引 -> SQLite 全表扫(5只10.6s);
改 UNION ALL per symbol 走 (sym,exc,interval,datetime) 索引 SEEK(5只0.03s, 340x)。
- chunk=400 抗大体量(参数<999 版本安全); interval/start/end regex 校验字面量注入。
- 部署 VPS。实测 5128只(000985) raw 68s / qfq 33s(逐只会卡死); 200只 qfq 0.9s。
qfq 与 get_price(fq=qfq) max_diff=0.00; 22 测试 0 回归。
P2 000938: sina_index_eod CODES 加 000938 -> 腾讯灌 1591 行点位(1395~2884, 非基金净值);
腾讯对该码停 2023-02-17(数据源限, 非代码问题)。
解锁: 策略02(000985 5128只选股) + 策略03 长期回测; 策略层 pandas 向量化(P2)归策略 session。
This commit is contained in:
@@ -13,6 +13,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import sqlite3
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
@@ -38,6 +39,27 @@ _DEFAULT_DATA_DIR = r"C:\sanguo_vnpy_v2\data"
|
||||
_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_]+$")
|
||||
|
||||
|
||||
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")``。
|
||||
@@ -248,24 +270,28 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
|
||||
start: Union[str, datetime],
|
||||
end: Union[str, datetime],
|
||||
interval: str = "d",
|
||||
fq: str = "raw",
|
||||
) -> pd.DataFrame:
|
||||
"""批量取多只股票 close,返宽表 (index=datetime, columns=symbol, values=close)。
|
||||
|
||||
单条 dbbardata 参数化查询 ``(symbol, exchange) OR (...)`` + ``pivot``,
|
||||
替代选股时 N 次 ``get_price``(为 Phase 2 策略向量化提速 10-50× 铺路)。
|
||||
|
||||
与逐只 ``get_price(fq='raw')`` 等价(raw close, 不复权),区别仅在 IO 次数:
|
||||
- 逐只: N 次 SQL(N+1 query pattern)
|
||||
- 批量: 1 次 SQL + pivot
|
||||
生产化增强(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``。缺失股票 → 该列全 NaN; 重复行 dedup (keep last)。
|
||||
``values=close_price``(fq='qfq' 时为复权 close)。缺失股票 → 该列全 NaN;
|
||||
重复行 dedup (keep last)。
|
||||
|
||||
空列表 → 空 DataFrame(index=DatetimeIndex,columns=[])。
|
||||
"""
|
||||
@@ -283,19 +309,38 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
|
||||
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")
|
||||
|
||||
# 参数化 (symbol, exchange) OR 子句 — 防注入 + 不字符串拼接(长度无上限)
|
||||
where_parts = " OR ".join(["(symbol=? AND exchange=?)" for _ in pairs])
|
||||
params: List[Any] = [interval, start_str, end_str]
|
||||
for _inp, sym, exc in pairs:
|
||||
params.extend([sym, exc])
|
||||
# 分批查询 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)
|
||||
|
||||
q = (
|
||||
"SELECT datetime, symbol, close_price FROM dbbardata "
|
||||
"WHERE interval=? "
|
||||
"AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? "
|
||||
f"AND ({where_parts})"
|
||||
)
|
||||
df = pd.read_sql(q, conn, params=params)
|
||||
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 行为一致
|
||||
@@ -321,16 +366,110 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
|
||||
wide = df.pivot(index="datetime", columns="symbol", values="close_price")
|
||||
|
||||
# 补缺失 symbol 列(全 NaN), 按 input 顺序对齐 columns
|
||||
for inp in symbols:
|
||||
if inp not in wide.columns:
|
||||
wide[inp] = float("nan")
|
||||
# (批量 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,
|
||||
|
||||
@@ -55,10 +55,11 @@ from dbbardata_utils import normalize_daily_dt # noqa: E402
|
||||
DB = r"C:\sanguo_vnpy_v2\data\quant_trading.db"
|
||||
T0 = time.time()
|
||||
|
||||
# 14 个中证指数 (与 SZSE 股票码碰撞, 仅灌 SSE 点位行)
|
||||
# 15 个中证指数 (与 SZSE 股票码碰撞, 仅灌 SSE 点位行)
|
||||
# 000938 = 中证 1000 等权 (策略需求: RPS/动量基准对比, 加 2026-07-28)
|
||||
CODES = [
|
||||
"000928", "000929", "000930", "000931", "000932", "000933",
|
||||
"000934", "000935", "000936", "000937",
|
||||
"000934", "000935", "000936", "000937", "000938",
|
||||
"000852", "000905", "000016", "000985",
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user