From e7f9426bcd856ac064772ed8a8ddf25ae09bc81b Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 28 Jul 2026 23:03:32 +0800 Subject: [PATCH] =?UTF-8?q?feat(data):=20get=5Fcloses=5Fpanel=20=E7=94=9F?= =?UTF-8?q?=E4=BA=A7=E5=8C=96(qfq=20=E5=89=8D=E5=A4=8D=E6=9D=83=20+=20UNIO?= =?UTF-8?q?N=20ALL=20340x=20=E6=8F=90=E9=80=9F=20+=20chunk)=20+=20000938?= =?UTF-8?q?=20=E8=A1=8C=E6=83=85?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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。 --- .../providers/local_unified_provider.py | 183 +++++++++++++++--- scripts/data_platform/sina_index_eod.py | 5 +- 2 files changed, 164 insertions(+), 24 deletions(-) diff --git a/sanguo_portfolio/providers/local_unified_provider.py b/sanguo_portfolio/providers/local_unified_provider.py index 321f84e..2b76a06 100644 --- a/sanguo_portfolio/providers/local_unified_provider.py +++ b/sanguo_portfolio/providers/local_unified_provider.py @@ -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, diff --git a/scripts/data_platform/sina_index_eod.py b/scripts/data_platform/sina_index_eod.py index cda5078..c0c2d76 100644 --- a/scripts/data_platform/sina_index_eod.py +++ b/scripts/data_platform/sina_index_eod.py @@ -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", ]