diff --git a/CHANGELOG.md b/CHANGELOG.md index fbe68ab2278..5887bebb5da 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,113 @@ # CHANGELOG +## 2026-07-22(审查修复) + +### 修复代码审查发现的两个问题(分段丢弃 + 正则兜底误判) + +- **分段合并丢弃正确分段**:`mergeSegmentResults` 旧实现遇到占位 TABLE(某 + 分段计算失败/空的兜底)就 `break`,把该表达式其余分段已算出的正确矩阵 + 全部丢弃,与「分段边界结果与整段一致」的保证矛盾。改为把失败/空分段的 + 占位由 `buildPlaceholderMatrix` 统一产出「行-该段交易日、列-presentCodes」 + 的 NULL 矩阵(与正常结果同为带标签矩阵):合并时一视同仁按行拼接,任一 + 分段异常都不再丢弃其它分段;顺带消除空 `baseData` 分段的日期空洞,并去掉 + Python 侧期望三元素矩阵却收到 TABLE 的类型隐患。占位填 NULL 而非 0,避免 + 污染缺失位置的因子值。移除因此不再被调用的 `createEmptyTable`。 +- **正则兜底把数值常量误判为窗口**:回看解析兜底(仅对 qlib 解析不了的表达式 + 生效)取最大独立整数为窗口,会把 `gtjaAlpha191_005($volume/1000000, ...)` + 里的 `1000000` 当成百万日回看,在 2 核/8GB 上触发超量查询。新增 + `_MAX_FALLBACK_WINDOW=2000`(约 8 年交易日)上界:超过即视为常量剔除, + 剔除后无候选则退回 `ddb_lookback_default`。新增 3 项针对性单测。 + +## 2026-07-22 + +### DDB 计算分支:滚动算子自动取前序期 + 真实分段计算(替代虚假 repartitionDS) + +修复两个相互纠缠的老问题(详见 +`docs/DDB取前序期与分段计算_20260722.md`): + +- **取前序期**:DDB 后端此前把用户的 `start_time` 原样下发, + `Mean($close,20)` 在区间头部前 19 个交易日必然 NaN(文件后端靠 + `get_extended_window_size` 外扩,DDB 分支绕过了该机制)。现在 Python 侧 + 复用 qlib 算子树解析每批表达式的(向前, 向后)外扩交易日数(嵌套算子 + 递归累加、双臂算子取 max、`Ref($close,-2)` 标签向后外扩),经 + `FeatureEngineeringByDate` 尾参传给 DDB;服务器端外扩查询窗口计算后 + 把结果矩阵截断回请求区间。qlib 无法实例化的表达式(DDB 专属 alpha 库 + 函数)退回启发式:正则扫独立整数当窗口,扫不到用 + `C["ddb_lookback_default"]`(默认 252)兜底。 +- **分段计算**:原 `repartitionDS(..., RANGE, [start, end])` 只有 2 个边界 + → 永远只产生 1 个数据源,mr 退化为单任务(代码中 FIXME 自证的 + "虚假的分区")。且真把日期切段后每段头部都会缺窗口——与取前序期 + 必须一起设计。现移除 repartitionDS+mr,改为:内存估算 + (`FeatureEngine.isRunWithAvailableMemory`,扩展窗口 + 表达式矩阵项 + + 70% 空闲内存上限)放得下则整段单次计算;放不下按 + `C["ddb_days_step"]` 切段**顺序循环**(单机 2 核 mr 并行收益小且峰值 + 内存翻倍),每段查询带 lookback 重叠(halo)、段内截断,各段 panel 用 + 统一列标签(窗口内实际出现的代码集合)对齐后 `concatMatrix` 纵向合并, + 保证分段边界处滚动结果与整段计算一致。 + +接口兼容:`FeatureEngineeringByDate` 新增尾参 +`lookbackDays=0, rightDays=0`(旧调用形式行为不变);Python 公共接口 +`D.features()` 签名不变。离线单测 +12(回看解析/批量取 max/脚本接线), +DDB 端脚本行为待 live 验证。 + +## 2026-07-21 + +### DDB 后端系统性优化(分支 optimize/ddb-backend,17 个原子提交) + +面向 DolphinDB **社区版(2 核 / 8GB)** 的一轮优化:省往返、省传输、 +省服务器内存、控会话数;公共接口 `D.features()/D.calendar()/D.instruments()` +的签名与返回结果完全不变。全部改动配套离线单测(mock 会话,共 115 项通过)。 + +**Bug 修复**: + +- `ddb_dataset_processor` 在传入 `inst_processors` 时提前 `return pd.DataFrame()`, + 并行处理结果被整体丢弃(`data.py:782`)。 +- `DDBClient` 连接池:`_pool_instance` 类变量导致多客户端共享同一池; + `close_pool` 引用不存在的 `_pool_lock` 且永远读到 None(从未真正关闭过池); + 删除调用不存在 `get_session()` 的死方法 `tableAppender/tableUpsert` 与 + `__main__` 中的硬编码凭据。 +- `DolphinDBClientProvider` init 时急切创建 4 连接的池(无消费方)→ 改懒创建。 +- `load_mysql_plugin` 误装 `lgbm` 插件(应为 `mysql`)。 +- `DolphinDBDataLoader.__exit__/__del__` 无条件关闭进程级共享 session → + 引入 `_owns_session` 所有权标志。 + +**并发/线程安全**: + +- 新增会话级 `DBClient.session_lock`(RLock):feature 查询是 + 「run→upload→run」多步会话对话,交错执行会互相覆盖服务器变量 + (跨线程数据污染根因);所有 session 触点统一持锁。 +- `QlibDataLoader` 全局 `_load_lock` 收窄至 DDB 路径,文件后端恢复无锁并行。 + +**健壮性/质量**: + +- 12 处 `print` → loguru;3 处裸 `except:` 收敛;异常重包装补 `from e`。 +- MySQL 同步 SQL 参数白名单校验(`validate_date_str`/`validate_sql_identifier`, + 唯一真实注入面)。 +- 清理死代码(注释块 ×2、零调用者函数 ×2、16KB 备份脚本)与 + `OPERATOR_MAPPING` 重复键(生效映射不变,ast 测试锁定)。 + +**性能**(RPC 往返:计算分支每批 3-4 次 → 预热后 2 次): + +- 交易日历模块级缓存:Alpha158 一次 `D.features` 从 ~6 次全量日历下载 → 0; + `DBCalendarStorage.index()/__getitem__` 改走 `H["c"]` 缓存。 +- 日期字面量内联(消灭独立上传往返);纯字段分支合并为单条 SQL 脚本; + 上传字典按分支裁剪。`fetch_features_from_ddb` 拆为编排器 + 4 个 helper。 +- existsTable 仅正向缓存、股票池走 `H["i"]`、表 schema 进程内缓存、 + 表达式翻译 lru_cache;统一由 `ddb_qlib.invalidate_ddb_caches()` 失效 + (写路径自动调用)。 +- 计算分支结果直构 `(instrument, datetime)` MultiIndex 面板,替代 + concat→unstack→stack→swaplevel 的 ~4 次全景拷贝(保留 legacy 兜底 + + 逐值等价测试)。 +- 三个 alpha 因子库(约 119KB)按字段前缀惰性加载(`Cannot recognize the + token` 兜底重试;`C["ddb_preload_alpha_libs"]` 逃生开关)。 +- 写路径按 URI 复用共享 `DDBClient`(批量导入 N 会话 → 1)。 +- 批次参数配置化:`C["ddb_field_chunk_size"]=30`、`C["ddb_days_step"]=252` + (默认值不变,`scripts/benchmark_ddb_backend.py` 供 live 校准)。 + +**明确不做**(2 核/8GB 约束):读路径不引入 `DBConnectionPool` 并发、 +不调高 mr 并行度、不做服务器端全景 pivot、不引入 module/functionView 持久化。 + + ## 2026-07-20 ### fix: DDB 批量取数路径恢复成分股 spans(入池/出池)过滤 ⚠️ BREAKING 行为变更 diff --git "a/docs/DDB\345\217\226\345\211\215\345\272\217\346\234\237\344\270\216\345\210\206\346\256\265\350\256\241\347\256\227_20260722.md" "b/docs/DDB\345\217\226\345\211\215\345\272\217\346\234\237\344\270\216\345\210\206\346\256\265\350\256\241\347\256\227_20260722.md" new file mode 100644 index 00000000000..24dac4ef0a2 --- /dev/null +++ "b/docs/DDB\345\217\226\345\211\215\345\272\217\346\234\237\344\270\216\345\210\206\346\256\265\350\256\241\347\256\227_20260722.md" @@ -0,0 +1,118 @@ +# DDB 后端:滚动算子取前序期与真实分段计算 + +> 分支 `optimize/ddb-backend`,2026-07-22。 +> 涉及文件:`qlib/data/backend/ddb_qlib/ddb_features.py`、 +> `qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineering.dos`、`qlib/config.py`。 + +## 背景:两个纠缠在一起的老问题 + +### 问题一:DDB 后端没有取前序期 + +qlib 文件后端在 `LocalExpressionProvider.file_expression` 中先调用算子树的 +`get_extended_window_size()` 得到左右外扩量,把查询窗口外扩后计算、再截断回 +请求区间(`qlib/data/data.py`)。而 DDB 分支 `ddb_expression` 直接把用户的 +`start_time`/`end_time` 下发给服务器端 `parseExpr` 计算——`Mean($close,20)` +在 `start_date` 后的前 19 个交易日必然是 NaN。 + +### 问题二:repartitionDS 分区从未生效 + +`FeatureEngine.createDSByDate` 写的是: + +```dolphinscript +repartitionDS(sql(...), `date, RANGE, [startTime, boundEndDate]) +``` + +RANGE 分区方案 n+1 个边界产生 n 个数据源——这里只有 2 个边界,**永远只产生 +1 个数据源**,`mr` 退化为单任务:既没有并行,也没有分批控内存(代码中 +FIXME "虚假的分区" 自证)。 + +### 为什么必须一起设计 + +一旦让日期分段真的生效,每个分段独立计算时滚动算子在**每一段的开头**都会 +产生 L 天 NaN(不止第一段)。而 repartitionDS 的 RANGE 分区天然不重叠, +表达不了"每段多带 L 天头部数据",所以 repartitionDS+mr 架构对滚动算子 +结构性不适用。 + +## 方案 + +### Python 侧:回看窗口解析(`ddb_features.py`) + +`get_expression_extended_window(expr, default_lookback) -> (lft, rght)`: + +1. **优先复用 qlib 算子树**:用 `parse_field` + 受限命名空间 eval 实例化 + 表达式(与 `ExpressionProvider.get_expression_instance` 同款),调 + `get_extended_window_size()`。嵌套算子递归累加(`Mean(Ref($close,5),20)` + → 24)、双臂算子取 max(`Corr`)、未来引用向后外扩(标签 + `Ref($close,-2)` → rght=2)全部免费获得。 +2. **启发式兜底**:qlib 无法实例化的表达式(gtjaAlpha/WQAlpha/qlib158Alpha + 等 DDB 专属函数)退回正则——独立整数字面量当窗口(标识符里的数字如 + `gtjaAlpha191_001` 不会误判);扫不到时用 `C["ddb_lookback_default"]` + (默认 252)兜底。真正非法的表达式仍在 DDB 端 parseExpr 阶段报错。 + +`batch_extended_window`:同批表达式共享一份基础数据面板,按批取 max。 +结果经 `FeatureEngineeringByDate(..., daysStep, lookbackDays, rightDays)` +尾参下发(旧调用形式默认 0,行为不变)。 + +### DDB 侧:外扩 + 截断 + 真实分段(`featureEngineering.dos`) + +`FeatureEngine.fetch(daysStep)` 流程: + +1. `shiftTradeDays` 按 `temporalAdd(dt, ±N, `XSHG)` 计算扩展窗口 + `[extStart, extEnd]`。 +2. `fetchPresentCodes`:扩展窗口内实际有数据的代码集合(升序),作为 + 所有 panel 的统一列标签。 +3. `isRunWithAvailableMemory(extStart, extEnd)`:峰值内存估算 ≈ 基础长表 + ×2(长表 + panel 字典)+ 每表达式一个 dates×codes 的 DOUBLE 矩阵, + 上限为 70% 空闲内存。 +4. **放得下** → `computeRangeFeatures(startTime, endTime, presentCodes)` + 整段单次计算:查询 `[extStart, extEnd]`,panel + parseExpr 计算后用 + `loc` 把结果矩阵行截断回 `[startTime, endTime]`(外扩部分只服务于 + 窗口计算,不进入结果)。 +5. **放不下** → 按 `daysStep` 切段**顺序循环**(单机 2 核 mr 并行收益小 + 且峰值内存翻倍):每段查询 `[段起点−L, 段终点+R]`(halo 重叠)、段内 + 截断,`mergeSegmentResults` 用 `concatMatrix(mats, false)` 纵向拼接并 + 重挂行/列标签。列已按 presentCodes 对齐,分段边界处滚动结果与整段 + 计算逐值一致。 + +### 分段占位统一为 NULL 矩阵(审查修复) + +失败/空分段的占位由 `buildPlaceholderMatrix` 产出「行-该段交易日、列- +presentCodes」的 NULL DOUBLE 矩阵,与正常结果**同为带标签矩阵**。这样: + +- `mergeSegmentResults` 对所有分段一视同仁按行拼接,**任一分段计算失败 + 不再丢弃其它分段的正确结果**(旧实现遇到占位 TABLE 就 `break`,会把整个 + 表达式的多年结果塌缩为单个短占位); +- 空 `baseData` 的分段也填占位而非整体缺失,**消除分段日期空洞**; +- 占位填 NULL(非 0),避免用 0 污染缺失位置的因子值;也消除了 Python 侧 + `_computed_dict_to_panel` 期望三元素矩阵、却收到 TABLE 的类型隐患。 + +## 配置 + +| 键 | 默认 | 说明 | +|----|------|------| +| `ddb_lookback_default` | 252 | 回看解析兜底交易日数(仅对 qlib 解析不了且扫不到窗口整数的表达式生效) | +| `ddb_days_step` | 252 | 分段模式的段长(交易日);仅内存估算不足时启用分段 | + +注意:`L`(回看天数)接近 `daysStep` 时分段 halo 开销接近 100%;表达式 +含长窗口(如 252 日滚动)且触发分段时,可适当调大 `ddb_days_step`。 + +### 兜底窗口上界(审查修复) + +回看正则兜底仅对 qlib 解析不了的表达式生效,取最大独立整数为窗口。为避免把 +数值常量误判为窗口(如 `gtjaAlpha191_005($volume/1000000, ...)` 中的 +`1000000` 被当成百万日回看,进而 `shiftTradeDays` 外扩到不合理区间、在 +2 核/8GB 上触发超量查询),扫到的整数超过 `_MAX_FALLBACK_WINDOW`(2000, +约 8 年交易日)即视为常量剔除;剔除后无候选则退回 `ddb_lookback_default`。 + +## 验证状态 + +- 离线单测:`tests/test_ddb_lookback.py`(回看解析含大常量剔除/上界共 15 例 + + 脚本接线 3 例),`tests/test_fetch_features.py` 分支回归全部通过。 +- **待 live 验证**(本机无 DDB 服务器): + 1. `featureEngineering.dos` 语法与 `panel(...,, colLabels)` / `loc` / + `concatMatrix` / `rename!` / `matrix(DOUBLE,r,c,,double(NULL))` 行为; + 2. 等价性:同一表达式(如 `Mean($close,20)`)文件后端 vs DDB 后端, + 区间头部 L 天结果对齐; + 3. 分段模式:临时把 `isRunWithAvailableMemory` 强制返回 false,比对 + 分段结果与整段结果逐值一致;并构造「某分段该表达式失败」用例, + 确认其它分段结果不被丢弃、占位段为 NaN 且日期覆盖连续。 diff --git a/qlib/config.py b/qlib/config.py index 651c32ad162..475be0df57b 100644 --- a/qlib/config.py +++ b/qlib/config.py @@ -111,6 +111,18 @@ def register_from_C(config, skip_register=True): "provider": "LocalProvider", # database uri for database backend "database_uri": None, + # ---- DolphinDB 后端批次参数(目标服务器为社区版 2 核/8GB 时的取舍)---- + # 每批查询的字段数:增大会加宽客户端 concat 宽度、减少批次数 + "ddb_field_chunk_size": 30, + # FeatureEngineeringByDate 的日期分片(交易日数,内存不足才启用分段): + # 单段服务器内存 ≈ days_step × 股票数 × 基础字段数 × 8B + # (5000 股 × 5 字段 × 252 天 ≈ 50MB),8GB 内存建议不超过 504 + "ddb_days_step": 252, + # 表达式回看窗口的兜底交易日数:qlib 算子树解析不了(DDB 专属 alpha + # 库函数等)且正则扫不到窗口整数时,向前外扩这么多个交易日 + "ddb_lookback_default": 252, + # True 时 init 全量预载三个 alpha 因子库(默认按字段引用惰性加载) + "ddb_preload_alpha_libs": False, # dolphindb provider config # "dolphindb_provider": None, "dbclient_provider":None, diff --git a/qlib/data/backend/ddb_qlib/README.md b/qlib/data/backend/ddb_qlib/README.md index 33e8370295f..a81e7fd2cba 100644 --- a/qlib/data/backend/ddb_qlib/README.md +++ b/qlib/data/backend/ddb_qlib/README.md @@ -189,6 +189,39 @@ bridge.close() - 做好异常处理和错误日志 - 定期检查数据一致性 +5. **进程内缓存与失效** + + 为减少 RPC 往返,本后端在进程内缓存:交易日历(`TradeDateUtils`)、 + 表 schema(`utils.get_table_columns`)、`existsTable` 正向结果、 + 股票池(`H["i"]`)与表达式翻译(lru_cache)。 + + - 经 `write_df_to_ddb` / CSV 导入 / `clean_qlib_db` 的写入会**自动失效**缓存; + - 若数据被**外部进程**写入(本进程无感知),长驻只读进程需手动失效: + + ```python + from qlib.data.backend.ddb_qlib import invalidate_ddb_caches + invalidate_ddb_caches() + ``` + +6. **批次参数(社区版 2 核/8GB 调优)** + + ```python + qlib.init( + database_uri="dolphindb://...", + ddb_field_chunk_size=30, # 每批查询的字段数 + ddb_days_step=252, # 内存不足时的日期分段长(交易日数),8GB 建议 ≤504 + ddb_lookback_default=252, # 表达式回看解析失败时的兜底外扩交易日数 + ddb_preload_alpha_libs=False, # True 恢复 init 全量预载 alpha 因子库 + ) + ``` + + 计算分支自动按表达式外扩查询窗口(取前序期)并在服务器端截断回请求 + 区间;内存估算不足时按 `ddb_days_step` 分段顺序计算(每段带回看重叠), + 详见 `docs/DDB取前序期与分段计算_20260722.md`。 + + 可用 `DDB_BENCH_URI=... python scripts/benchmark_ddb_backend.py` 对 + 实际服务器做基准校准。 + ## 技术支持 如有问题,请联系: diff --git a/qlib/data/backend/ddb_qlib/__init__.py b/qlib/data/backend/ddb_qlib/__init__.py index fed992db899..896a060c5d9 100644 --- a/qlib/data/backend/ddb_qlib/__init__.py +++ b/qlib/data/backend/ddb_qlib/__init__.py @@ -19,9 +19,41 @@ from .schemas import QlibTableSchema from .ddb_features import ( register_ddb_functions_to_qlib, - # ddb_compute_features, fetch_features_from_ddb, normalize_fields_to_ddb, TradeDateUtils, adapt_qlib_expr_syntax_for_ddb, ) + + +def invalidate_ddb_caches() -> None: + """清空 DDB 后端的进程内缓存。 + + 覆盖范围: + - :class:`TradeDateUtils` 的模块级交易日历缓存; + - 表列名缓存(:func:`utils.get_table_columns`); + - 存储层 existsTable 正向缓存(``DBStorageMixin._exists_cache``); + - ``H["c"]`` 原始日历缓存与 ``H["i"]`` 股票池缓存。 + + 写路径(``write_df_to_ddb`` / CSV 导入 / ``clean_qlib_db``)变更表数据后 + 会自动调用;长驻只读进程若感知到外部写入,也可手动调用本函数。 + """ + from .utils import clear_table_columns_cache + + TradeDateUtils.clear_cache() + clear_table_columns_cache() + try: + from qlib.data.storage.dolphindb_storage import DBStorageMixin + + DBStorageMixin._exists_cache.clear() + except Exception: # pragma: no cover - storage 未加载时跳过 + pass + try: + from qlib.data.cache import H + + for unit_key in ("c", "i"): + unit = H[unit_key] + while len(unit): + unit.popitem(last=False) + except Exception: # pragma: no cover - 缓存清理失败不应阻断写路径 + pass diff --git a/qlib/data/backend/ddb_qlib/ddb_client.py b/qlib/data/backend/ddb_qlib/ddb_client.py index 0260634dc19..8e2e1f42d37 100644 --- a/qlib/data/backend/ddb_qlib/ddb_client.py +++ b/qlib/data/backend/ddb_qlib/ddb_client.py @@ -1,17 +1,18 @@ """ Author: hugo2046 shen.lan123@gmail.com Date: 2025-02-19 14:29:41 -LastEditors: hugo2046 shen.lan123@gmail.com -LastEditTime: 2025-04-17 16:02:46 Description: 用于连接ddb """ -from typing import List, Optional, Union -from urllib.parse import unquote, urlparse +import threading import dolphindb as ddb -import pandas as pd from pydantic import BaseModel, field_validator +from urllib.parse import unquote, urlparse + +from ....log import get_module_logger + +logger = get_module_logger("ddb_client") class DDBConnectionSpec(BaseModel): @@ -31,13 +32,21 @@ def validate_uri(cls, v: str) -> str: class DDBClient: - - _pool_instance = None # 类变量用于存储连接池实例 + """DolphinDB 连接客户端:持有一个会话与一个惰性创建的连接池。 + + - 会话在构造时建立(reconnect=True 提供断线重连)。 + - 连接池为实例属性且惰性创建:目标服务器为社区版(2 核/8GB)时, + 每条连接都占用服务器内存,只有真正用到时才创建。 + + :param config: 连接配置(含 dolphindb:// URI) + """ - def __init__(self, config: DDBConnectionSpec): self._config = config self._session = self._create_session() + # ⚠️ 连接池必须是实例属性:类变量会导致连接不同服务器的多个客户端共享同一个池 + self._pool_instance: ddb.DBConnectionPool | None = None + self._pool_lock = threading.Lock() def _parse_uri(self) -> tuple[str, str, str, int]: """解析连接参数""" @@ -60,7 +69,7 @@ def _create_session(self) -> ddb.Session: session.connect(host=host, port=port, userid=user, password=pwd, reconnect=True) return session - def _create_pool(self,threadNum:int=4) -> ddb.DBConnectionPool: + def _create_pool(self, threadNum: int = 4) -> ddb.DBConnectionPool: """创建连接池""" user, pwd, host, port = self._parse_uri() return ddb.DBConnectionPool( @@ -80,126 +89,31 @@ def session(self) -> ddb.Session: @property def pool(self) -> ddb.DBConnectionPool: - """直接获取连接池对象""" + """惰性获取连接池对象(线程安全的双重检查)""" if self._pool_instance is None: - self._pool_instance = self._create_pool() + with self._pool_lock: + if self._pool_instance is None: + self._pool_instance = self._create_pool() return self._pool_instance - - - @classmethod - def close_pool(cls): - """关闭并清理连接池""" - if cls._pool_instance is not None: - with cls._pool_lock: - if cls._pool_instance is not None: - try: - cls._pool_instance.shutDown() - except: - pass - cls._pool_instance = None - - def tableAppender(self, db_path: str, table_name: str, data: pd.DataFrame) -> None: - """ - 将数据追加到指定的表中。 - - :param db_path: 数据库路径 - :type db_path: str - :param table_name: 表名 - :type table_name: str - :param data: 要追加的数据 - :type data: pd.DataFrame - :raises ValueError: 如果数据库路径或表名无效 - """ - session = self.get_session() - - if not session.existsDatabase(db_path): - raise ValueError(f"{db_path} is not a valid database") - - if not session.existsTable(db_path, table_name): - raise ValueError(f"{table_name} is not a valid table") - # 获取表的列名 - table_cols: List[str] = ( - session.loadTable(table_name, db_path).schema["name"].to_list() - ) - # 数据列名顺序与表列名顺序一致 - data: pd.DataFrame = data[table_cols] - - appender: ddb.tableAppender = ddb.tableAppender( - tableName=table_name, - ddbSession=session, - dbPath=db_path, - ) - - appender.append(data) - - def tableUpsert( - self, - db_path: str, - table_name: str, - data: pd.DataFrame, - keyColNames: Optional[Union[str, List]] = None, - sortColumns: Optional[Union[str, List]] = None, - ) -> None: - """ - 将数据帧中的数据插入或更新到指定的数据库表中。 - - :param db_path: 数据库路径。 - :param table_name: 表名。 - :param data: 包含要插入或更新的数据的数据帧。 - :param keyColNames: 用于确定唯一记录的列名。默认为None。 - :param sortColumns: 用于排序的列名。默认为None。 - :return: 无返回值。 + def close_pool(self) -> None: + """关闭并清理连接池(未创建过则为空操作)""" + with self._pool_lock: + if self._pool_instance is not None: + try: + self._pool_instance.shutDown() + except Exception as e: + logger.warning(f"关闭连接池失败: {e}") + self._pool_instance = None + + def close(self) -> None: + """显式关闭会话与连接池。 + + 不提供 __del__:本客户端的会话会被 provider/storage 等多处共享, + GC 时机关闭共享会话会破坏其他使用方(见 DolphinDBDataLoader 历史 bug)。 """ - session = self.get_session() - - if not session.existsDatabase(db_path): - raise ValueError(f"{db_path} is not a valid database") - - if not session.existsTable(db_path, table_name): - raise ValueError(f"{table_name} is not a valid table") - - # 获取表的列名 - table_cols: List[str] = ( - session.loadTable(table_name, db_path).schema["name"].to_list() - ) - - # 数据列名顺序与表列名顺序一致 - data: pd.DataFrame = data[table_cols] - - if keyColNames is None: - keyColNames: List = [] - - if sortColumns is None: - sortColumns: List = [] - - upserter: ddb.tableUpsert = ddb.tableUpsert( - tableName=table_name, - ddbSession=session, - dbPath=db_path, - keyColNames=keyColNames, - sortColumns=sortColumns, - ) - upserter.upsert(data) - - # def __del__(self): - # """析构函数,确保资源正确释放""" - # try: - # if hasattr(self, '_session') and self._session: - # self._session.close() - # except: - # pass - -if __name__ == "__main__": - - uri: str = "dolphindb://admin:123456@114.80.110.170:28848" - - # 测试python端安装ddb插件 - config = DDBConnectionSpec(uri=uri) - connector = DDBClient(config) - - # expr:str = """ - # installPlugin("lgbm") - # loadPlugin("lgbm") - # """ - # connector.session.run(expr) + self.close_pool() + try: + self._session.close() + except Exception as e: + logger.warning(f"关闭会话失败: {e}") diff --git a/qlib/data/backend/ddb_qlib/ddb_features.py b/qlib/data/backend/ddb_qlib/ddb_features.py index 32d04fa979d..bbf9d3a646b 100644 --- a/qlib/data/backend/ddb_qlib/ddb_features.py +++ b/qlib/data/backend/ddb_qlib/ddb_features.py @@ -7,6 +7,7 @@ """ import bisect +import functools import re from pathlib import Path from typing import Dict, Iterable, List, Tuple, Union @@ -58,8 +59,6 @@ "Mad": "mmad", "Rank": "rolling_rank", "Count": "mcount", - "Slope": "mslr", - "Resi": "mmse", "WMA": "mwavg", "EMA": "ema", "Corr": "mcorr", @@ -97,27 +96,183 @@ def _sort_ddb_scripts(scripts: Iterable[Path]) -> List[Path]: return sorted(scripts, key=lambda p: (0 if p.name == "ops.dos" else 1, p.name)) -def register_ddb_functions_to_qlib(session: ddb.Session) -> None: +# 核心脚本:查询引擎必需,init 即加载(ops.dos 必须置首,见 _sort_ddb_scripts) +CORE_DDB_SCRIPTS: Tuple[str, ...] = ( + "ops.dos", + "featureEngineering.dos", + "prepareInstruments.dos", +) + +# alpha 因子库(合计约 119KB 脚本):仅当字段引用对应前缀函数时才按需加载, +# 避免每次 qlib.init 都让服务器解析大量通常用不到的脚本(社区版 2 核/8GB) +ALPHA_LIB_SCRIPTS: Dict[str, str] = { + "gtjaalpha": "gtja191Alpha.dos", + "qlib158alpha": "qlib158Alpha.dos", + "wqalpha": "wq101alpha.dos", +} + +# 会话对象上记录已加载 alpha 库前缀集合的属性名(函数注册是会话级状态) +_LOADED_LIBS_ATTR = "_qlib_loaded_alpha_libs" + + +def register_ddb_functions_to_qlib( + session: ddb.Session, preload_alpha_libs: Union[bool, None] = None +) -> None: """ 在 DolphinDB 会话中注册与 qlib 对应的自定义函数 - 这些函数实现了 qlib 中的特定操作,使 DolphinDB 能够兼容 qlib 的表达式计算 + 这些函数实现了 qlib 中的特定操作,使 DolphinDB 能够兼容 qlib 的表达式计算。 + 默认仅加载核心脚本(ops/featureEngineering/prepareInstruments);三个 + alpha 因子库由 :func:`ensure_alpha_libs_loaded` 在字段引用时按需加载。 参数: - session: DolphinDB 会话实例 + - preload_alpha_libs: True 时恢复历史行为(init 全量加载 alpha 库); + None 时读取配置 ``C["ddb_preload_alpha_libs"]``(默认 False) """ + if preload_alpha_libs is None: + try: + from ....config import C + + preload_alpha_libs = bool(C.get("ddb_preload_alpha_libs", False)) + except Exception: + preload_alpha_libs = False # 在会话中执行函数定义脚本 # ⚠️ 必须用 _sort_ddb_scripts 确定加载顺序,不能直接遍历 glob(其顺序依赖 # 文件系统 readdir,跨平台不一致;详见 _sort_ddb_scripts 的文档)。 script_path = Path(__file__).parent / "ddb_scripts" - for script_file in _sort_ddb_scripts(script_path.glob("*.dos")): + if preload_alpha_libs: + scripts = _sort_ddb_scripts(script_path.glob("*.dos")) + loaded_libs = set(ALPHA_LIB_SCRIPTS) + else: + scripts = _sort_ddb_scripts(script_path / name for name in CORE_DDB_SCRIPTS) + loaded_libs = set() + for script_file in scripts: session.runFile(script_file) + setattr(session, _LOADED_LIBS_ATTR, loaded_libs) get_module_logger("ddb_features").info("已注册 qlib 兼容函数到 DolphinDB 会话") +def ensure_alpha_libs_loaded(session: ddb.Session, fields) -> None: + """按字段引用惰性加载 alpha 因子库(每会话每库仅加载一次)。 + + 以大小写不敏感的前缀匹配(gtjaAlpha/qlib158Alpha/WQAlpha)扫描原始 + 字段文本;命中且未加载时 ``runFile`` 对应脚本并在会话上打标记。 + + :param session: DolphinDB 会话实例 + :param fields: 原始字段(str/list/dict 均可,转文本扫描) + """ + loaded = getattr(session, _LOADED_LIBS_ATTR, None) + if loaded is None: + loaded = set() + setattr(session, _LOADED_LIBS_ATTR, loaded) + if loaded >= set(ALPHA_LIB_SCRIPTS): + return + text = str(fields).lower() + script_path = Path(__file__).parent / "ddb_scripts" + for prefix, script_name in ALPHA_LIB_SCRIPTS.items(): + if prefix in loaded or prefix not in text: + continue + session.runFile(script_path / script_name) + loaded.add(prefix) + get_module_logger("ddb_features").info(f"已按需加载 alpha 因子库 {script_name}") + + +def _load_all_alpha_libs(session: ddb.Session) -> None: + """加载全部尚未加载的 alpha 因子库(未识别函数兜底重试用)。""" + loaded = getattr(session, _LOADED_LIBS_ATTR, None) + if loaded is None: + loaded = set() + setattr(session, _LOADED_LIBS_ATTR, loaded) + script_path = Path(__file__).parent / "ddb_scripts" + for prefix, script_name in ALPHA_LIB_SCRIPTS.items(): + if prefix not in loaded: + session.runFile(script_path / script_name) + loaded.add(prefix) + + ################################################################################################################## +# 回看窗口启发式解析:独立整数字面量(排除标识符与小数的组成部分, +# 如 gtjaAlpha191_001 中的 191 不会被误判为窗口) +_STANDALONE_INT_PATTERN = re.compile(r"(? Tuple[int, int]: + """解析 qlib 表达式所需的(向前, 向后)外扩交易日数。 + + 与文件后端 ``LocalExpressionProvider.file_expression`` 的取前序期语义 + 对齐:优先复用 qlib 算子树的 ``get_extended_window_size``(递归处理 + 嵌套算子、双臂算子取 max、以及 ``Ref($close,-2)`` 这类未来引用的向后 + 外扩)。qlib 无法实例化的表达式(DDB 专属 alpha 库函数等)退回启发式: + 正则扫描独立整数字面量当窗口(超过 :data:`_MAX_FALLBACK_WINDOW` 的整数 + 视为数值常量而非窗口予以剔除),扫不到时用 ``default_lookback`` 兜底; + 真正非法的表达式仍会在 DDB 端 parseExpr 阶段报错(原有行为)。 + + :param expr: 原始 qlib 表达式(含 ``$`` 字段引用) + :param default_lookback: 启发式扫不到窗口时的兜底回看交易日数 + :return: ``(向前外扩交易日数, 向后外扩交易日数)`` + """ + try: + from ....config import C + from ....utils import parse_field + from ...base import Feature, PFeature + from ...ops import Operators, register_all_ops + + # 未经 qlib.init 的独立调用场景(如离线测试)惰性注册默认算子 + if not getattr(Operators, "_ops", None): + register_all_ops(C) + # 安全说明:与 qlib 核心 ExpressionProvider.get_expression_instance + # 完全同款的 eval 解析(同一信任模型:字段表达式来自使用者配置), + # 且此处限定了命名空间仅暴露 Operators/Feature/PFeature + instance = eval( + parse_field(expr), + {"Operators": Operators, "Feature": Feature, "PFeature": PFeature}, + ) + lft_etd, rght_etd = instance.get_extended_window_size() + return int(lft_etd), int(rght_etd) + except Exception: + # 剔除超过上界的整数(视为数值常量而非窗口);剩余最大者为回看窗口 + ints = [ + n + for n in (int(s) for s in _STANDALONE_INT_PATTERN.findall(expr)) + if n <= _MAX_FALLBACK_WINDOW + ] + lft_etd = max(ints) if ints else int(default_lookback) + futures = [ + n + for n in (int(s) for s in _FUTURE_INT_PATTERN.findall(expr)) + if n <= _MAX_FALLBACK_WINDOW + ] + return lft_etd, max(futures) if futures else 0 + + +def batch_extended_window( + exprs: Iterable[str], default_lookback: int +) -> Tuple[int, int]: + """对一批表达式取(向前, 向后)外扩交易日数的最大值。 + + 同批表达式共享同一份基础数据面板,故按批取 max 外扩。 + + :param exprs: 原始 qlib 表达式集合 + :param default_lookback: 单表达式解析失败时的兜底回看交易日数 + :return: ``(向前外扩交易日数, 向后外扩交易日数)`` + """ + lft_etd, rght_etd = 0, 0 + for expr in exprs: + lft, rght = get_expression_extended_window(expr, default_lookback) + lft_etd, rght_etd = max(lft_etd, lft), max(rght_etd, rght) + return lft_etd, rght_etd + def normalize_fields_to_ddb( fields: Union[str, List[str], Dict], @@ -203,8 +358,10 @@ def build_field_expr( :return: 字段表达式列表,每个元素为字段名或"0 as 字段名"的表达式。 :rtype: List[str] """ - tb = session.loadTable(table_name, f"dfs://{db_name}") - table_columns: List[str] = tb.schema["name"].tolist() + from .utils import get_table_columns + + # 经进程内缓存读取列名(表结构仅随 DDL 变更),预热后 0 RPC + table_columns: List[str] = get_table_columns(session, f"dfs://{db_name}", table_name) return [col if col in table_columns else f"0 as {col}" for col in base_fields] @@ -213,10 +370,10 @@ def apply_spans_mask(data: pd.DataFrame, spans: Dict[str, List]) -> pd.DataFrame """按成分股 spans(入池/出池区间)对结果面板做行掩码。 用于 ``fetch_features_from_ddb`` 非纯字段分支的 Python 侧兜底过滤: - ``FeatureEngineeringByDate`` 经 ``mr`` 分布式执行时,worker 端无法还原 - ``conditionalFilter`` 所需的 dict(报 "filterMap must be a dictionary"), - 故 DDB 端按键列表全量计算,spans 过滤在此处补齐。先算后掩码的语义与 - 原版 qlib ``inst_calculator``(data.py 中按 spans 做 mask)一致。 + 计算分支向 ``FeatureEngineeringByDate`` 传键列表全量计算(历史上因 + mr worker 端还原不了 ``conditionalFilter`` 所需的 dict,协议保持不变), + spans 过滤在此处补齐。先算后掩码的语义与原版 qlib + ``inst_calculator``(data.py 中按 spans 做 mask)一致。 :param data: 以 ``(instrument, datetime)`` 为 MultiIndex 的结果面板。 :type data: pd.DataFrame @@ -291,7 +448,6 @@ def fetch_features_from_ddb( from .schemas import QlibTableSchema normalized_expr, base_fields, is_pure_fields, alias_to_origin_expr_map = normalize_fields_to_ddb(fields) - # reversed_expr: Dict = {v: k for k, v in normalized_expr.items()} _freq: str = "daily" if freq == "day" else "min" feature_schema: QlibTableSchema = getattr(QlibTableSchema(), f"feature_{_freq}")() @@ -299,32 +455,151 @@ def fetch_features_from_ddb( db_name: str = feature_schema.db_name table_name: str = feature_schema.table_name - # 获取实际查询范围 + # 交易日对齐并格式化为 DDB 日期字面量 + start_time, end_time = _resolve_query_window(session, freq, start_time, end_time) + + # 转换为列表格式 + if isinstance(instruments, str): + instruments = [instruments] + + # 防御性检查:filter_pipe 可能把 instruments 全部过滤掉, + # 空 instruments 直接早退,避免进入 DDB 端空数据兜底路径 + if not instruments: + return pd.DataFrame() + + # 字段若引用 alpha 因子库函数,按需加载对应 .dos(每会话仅一次) + ensure_alpha_libs_loaded(session, fields) + + if is_pure_fields: + data = _fetch_pure_fields( + session, instruments, base_fields, db_name, table_name, start_time, end_time + ) + else: + # 滚动/引用算子按 qlib 算子树外扩查询窗口(DDB 端计算后截断回请求 + # 区间),对齐文件后端 get_extended_window_size 的取前序期语义, + # 避免 Mean($close,20) 在 start_time 后前 19 个交易日全为 NaN + try: + from ....config import C + + default_lookback = int(C.get("ddb_lookback_default", 252)) + except Exception: + default_lookback = 252 + lookback_days, right_days = batch_extended_window( + alias_to_origin_expr_map.values(), default_lookback + ) + raw = _compute_expressions( + session, instruments, normalized_expr, base_fields, db_name, table_name, + start_time, end_time, lookback_days, right_days, + ) + # DDB 端在 baseData 为空时会返回空 dict(兜底路径),需要在此处早返回 + # 避免 pd.concat({}) 报 "No objects to concatenate" + if raw is None: + return pd.DataFrame() + data = _computed_dict_to_panel(raw) + + return _format_result_panel( + data, is_pure_fields, instruments, normalized_expr, alias_to_origin_expr_map + ) + + +def _resolve_query_window( + session: ddb.Session, + freq: str, + start_time: Union[pd.Timestamp, int], + end_time: Union[pd.Timestamp, int], +) -> Tuple[str, str]: + """按交易日历对齐查询区间,并格式化为 DDB 日期字面量。 + + ⚠️ 性能约定:日期以字面量内联进查询脚本,不再单独 + ``session.run`` 上传 ``start_time``/``end_time`` 服务器变量 + (曾经每批字段多付一次网络往返)。 + + :param session: DolphinDB 会话 + :param freq: 数据频率("day" / "min") + :param start_time: 开始时间(时间戳或日历下标) + :param end_time: 结束时间(时间戳或日历下标) + :return: (开始, 结束) 的 DDB 日期字面量字符串 + """ date_utils: TradeDateUtils = TradeDateUtils(session, freq) start_time, end_time = date_utils.get_locate_date(start_time, end_time) fmt: str = "%Y.%m.%d" if freq == "day" else "%Y.%m.%d %H:%M:%S" - start_time: pd.Timestamp = pd.to_datetime(start_time).strftime(fmt) - end_time: pd.Timestamp = pd.to_datetime(end_time).strftime(fmt) + return ( + pd.to_datetime(start_time).strftime(fmt), + pd.to_datetime(end_time).strftime(fmt), + ) + + +def _fetch_pure_fields( + session: ddb.Session, + instruments: Union[List[str], Dict], + base_fields: List[str], + db_name: str, + table_name: str, + start_time: str, + end_time: str, +) -> pd.DataFrame: + """纯字段分支:单条 SQL 直查特征表(含 spans dict 的 conditionalFilter 过滤)。 + + 复杂值(instruments)走变量上传、标量日期走字面量内联——参见 + dolphindb_skill 对「变量上传优于字符串拼接」的建议(doc_7605/doc_5576)。 + dict(成分股 spans)分支把日期-股票映射与主查询合并为一次 ``session.run``, + 消除历史上独立的 ``createDateStockMapping`` 往返。 - # 上传时间变量 - upload_dates: str = f""" - start_time={start_time}; - end_time={end_time}; + :return: 长表 DataFrame(columns: code, date, 各基础字段) """ - session.run(upload_dates) + session.upload({"instruments": instruments}) - # 转换为列表格式 - if isinstance(instruments, str): - instruments = [instruments] + # 基础字段如果缺失则使用0填充 + select_fields: List[str] = build_field_expr(session, db_name, table_name, base_fields) + select_clause: str = ", ".join(["code", "date"] + select_fields) - # 防御性检查:filter_pipe 可能把 instruments 全部过滤掉 - # 空 instruments 进入 DDB 会触发 fetchFeatures 空数据兜底路径, - # 返回 TABLE 与 union_dict 期望的 dict 类型不一致,报 "Incompatible vector/matrix size" - if not instruments: - return pd.DataFrame() + if isinstance(instruments, dict): + # spans dict:先在 DDB 中创建日期-股票映射,再用 conditionalFilter 过滤(同一脚本内完成) + script = f""" + codeRangeFilter = createDateStockMapping({start_time},{end_time},instruments); + select {select_clause} from loadTable("dfs://{db_name}","{table_name}") + where date between pair({start_time},{end_time}) + and conditionalFilter(code, date, codeRangeFilter) + order by date, code + """ + else: + script = f""" + select {select_clause} from loadTable("dfs://{db_name}","{table_name}") + where date between pair({start_time},{end_time}) and code in instruments + order by date, code + """ + return session.run(script) - # 上传变量 + +def _compute_expressions( + session: ddb.Session, + instruments: Union[List[str], Dict], + normalized_expr: Dict[str, str], + base_fields: List[str], + db_name: str, + table_name: str, + start_time: str, + end_time: str, + lookback_days: int = 0, + right_days: int = 0, +) -> Union[pd.DataFrame, None]: + """计算表达式分支:经 ``FeatureEngineeringByDate`` 服务器端计算。 + + DDB 端按内存估算自动选择整段单次计算或按日期分段顺序计算(防 OOM); + 每次查询按 ``lookback_days``/``right_days`` 外扩窗口,计算后截断回 + 请求区间,保证滚动算子在区间头部、未来引用在区间尾部不缺窗口。 + + ⚠️ dict(成分股 spans)时不把 dict 传给 FeatureEngineeringByDate: + 历史上其内部经 mr 分布式执行,worker 端还原不了 conditionalFilter + 所需的 dict。现保持同一协议——传键列表全量计算,spans 过滤在拿到 + 结果后由 apply_spans_mask 在 Python 侧补齐。 + + :param lookback_days: 向前外扩的交易日数(滚动算子回看窗口) + :param right_days: 向后外扩的交易日数(未来引用,如标签 Ref($close,-2)) + :return: 服务器返回的原始 ``{alias: [值矩阵(dates×codes), 日期, 代码]}``; + DDB 返回空 dict(baseData 为空的兜底路径)时返回 None + """ session.upload( { "instruments": instruments, @@ -333,102 +608,131 @@ def fetch_features_from_ddb( } ) - # 判断是否为纯字段表达式(如$close, $open),如果是则直接查询表 - if is_pure_fields: - # 基础字段如果缺失则使用0填充 - base_fields: List[str] = build_field_expr( - session, db_name, table_name, base_fields - ) - # 创建基础查询语句部分 - base_query = ( - session.loadTable(table_name, f"dfs://{db_name}") - .select(["code", "date"] + base_fields) - .where(f"date between pair({start_time},{end_time})") + inst_expr = "keys(instruments)" if isinstance(instruments, dict) else "instruments" + try: + from ....config import C + + days_step = int(C.get("ddb_days_step", 252)) + except Exception: + days_step = 252 + ddb_expr = f""" + FeatureEngineeringByDate({inst_expr},expressions,baseFields,{start_time},{end_time},"{db_name}","{table_name}",{days_step},{lookback_days},{right_days}) + """ + def _log_failure() -> None: + # 记录排查上下文;instruments 可能上千个,只记数量与头部样本 + inst_list = list(instruments) + get_module_logger("ddb_features").error( + f"DolphinDB 因子计算失败: db={db_name}/{table_name}, " + f"时间范围={start_time}~{end_time}, " + f"instruments 共{len(inst_list)}个(头部样本: {inst_list[:5]}), " + f"expressions={normalized_expr}, baseFields={base_fields}" ) - # 根据instruments类型选择不同的查询方式 - if isinstance(instruments, list): - # 如果是列表,直接使用in条件过滤 - data: pd.DataFrame = ( - base_query.where(f"code in instruments").sort(["date", "code"]).toDF() - ) + try: + data: Dict[str, List] = session.run(ddb_expr) + except Exception as e: + # 惰性加载兜底:未识别函数可能是绕过字段扫描调用的 alpha 函数, + # 全量加载 alpha 库后重试一次 + if "Cannot recognize the token" not in str(e): + _log_failure() + raise RuntimeError(f"DolphinDB 因子计算失败: {e}") from e + _load_all_alpha_libs(session) + try: + data: Dict[str, List] = session.run(ddb_expr) + except Exception as retry_e: + _log_failure() + raise RuntimeError(f"DolphinDB 因子计算失败: {retry_e}") from retry_e - elif isinstance(instruments, dict): - # 如果是字典,使用conditionalFilter进行复杂的日期-股票映射过滤 - # 先在DolphinDB中创建日期-股票映射,用以兼容spans - session.run( - "codeRangeFilter = createDateStockMapping(start_time,end_time,instruments)" - ) + if not data: + return None + return data - # 使用conditionalFilter应用复杂过滤条件 - data: pd.DataFrame = ( - base_query.where("conditionalFilter(code, date, codeRangeFilter)") - .sort(["date", "code"]) - .toDF() - ) +def _legacy_reshape(data: Dict[str, List]) -> pd.DataFrame: + """旧版重塑路径:concat → unstack → stack → swaplevel(多次全景拷贝)。 + + 仅作为 :func:`_computed_dict_to_panel` 形状不一致时的运行时兜底, + 并为直构路径的等价性测试提供对照实现。 + """ + try: + frames: Dict[str, pd.DataFrame] = { + k: pd.DataFrame(data=v[0], index=v[1], columns=v[2]) + for k, v in data.items() + } + except IndexError as e: + # 可能是data为空 + raise IndexError(f"传入的data数据为空: {e}") + panel: pd.DataFrame = pd.concat(frames) + if panel.empty: + return panel + # 根据 pandas 版本决定 stack 的参数 + if version.parse(pd.__version__) < version.parse("1.5.0"): + # 旧版本 pandas 不支持 future_stack,但支持 dropna=False + panel = panel.unstack(level=0).stack(level=0, dropna=False).swaplevel(0, 1) else: + # 新版本 pandas 使用 future_stack 或依赖默认行为 + panel = panel.unstack(level=0).stack(level=0, future_stack=True).swaplevel(0, 1) + return panel - # 使用mr防止因子计算时OOM - # ⚠️ dict(成分股 spans)时不能把 dict 传给 FeatureEngineeringByDate: - # 其内部经 mr 分布式执行,worker 端还原不了 conditionalFilter 所需的 - # dict(报 "filterMap must be a dictionary")。故此处传键列表全量计算, - # spans 过滤在拿到结果后由 apply_spans_mask 在 Python 侧补齐。 - inst_expr = "keys(instruments)" if isinstance(instruments, dict) else "instruments" - ddb_expr = f""" - FeatureEngineeringByDate({inst_expr},expressions,baseFields,start_time,end_time,"{db_name}","{table_name}") - """ - # data: pd.DataFrame = session.run(ddb_expr) - try: - data: Dict[str, List] = session.run(ddb_expr) - except Exception as e: - print(f"instruments:{instruments}") - print(f"db_name:{db_name},table_name:{table_name}") - print(f"start_time:{start_time},end_time:{end_time}") - print(f"expressions:{normalized_expr}") - print(f"baseFields:{base_fields}") - raise RuntimeError(f"DolphinDB 因子计算失败: {e}") - # DDB 端在 baseData 为空时会返回空 dict(兜底路径),需要在此处早返回 - # 避免 pd.concat({}) 报 "No objects to concatenate" - if not data: - return pd.DataFrame() - try: - data: Dict[str, pd.DataFrame] = { - k: pd.DataFrame(data=v[0], index=v[1], columns=v[2]) - for k, v in data.items() - } - except IndexError as e: - # 可能是data为空 - raise IndexError(f"传入的data数据为空: {e}") - data: pd.DataFrame = pd.concat(data) +def _computed_dict_to_panel(data: Dict[str, List]) -> pd.DataFrame: + """把 ``{alias: [值矩阵(dates×codes), 日期, 代码]}`` 直构为结果面板。 + ⚠️ 性能关键:旧路径 concat → unstack → stack → swaplevel → sort_index + 要做约 4 次全景拷贝(5000 股 × 多年 × 30 列时数百 MB 的客户端内存churn); + 此处按 (instrument, datetime) 直接一次分配构造。所有 alias 共享同一 + (dates, codes)(服务器端同一 panel() 产出);形状不一致(如 DDB 端 + 错误占位返回 TABLE)时回退 :func:`_legacy_reshape` 保持旧语义。 + """ + try: + aliases = list(data.keys()) + first = data[aliases[0]] + ref_dates = np.asarray(first[1]) + ref_codes = list(first[2]) + expected_shape = (len(ref_dates), len(ref_codes)) + for alias in aliases: + values, dates, codes = data[alias][0], data[alias][1], data[alias][2] + if np.asarray(values).shape != expected_shape: + raise ValueError("值矩阵形状不一致") + if not np.array_equal(np.asarray(dates), ref_dates) or list(codes) != ref_codes: + raise ValueError("日期/代码轴不一致") + except (ValueError, TypeError, IndexError, KeyError, AttributeError): + return _legacy_reshape(data) + + dates_index = pd.DatetimeIndex(ref_dates) + codes_arr = np.asarray(ref_codes) + # 预排序两轴,使输出与旧路径 sort_index() 后的顺序一致 + date_order = np.argsort(dates_index.values, kind="stable") + code_order = np.argsort(codes_arr, kind="stable") + index = pd.MultiIndex.from_product( + [codes_arr[code_order], dates_index[date_order]], names=["instrument", "datetime"] + ) + columns = { + alias: np.asarray(data[alias][0])[np.ix_(date_order, code_order)].T.ravel() + for alias in aliases + } + return pd.DataFrame(columns, index=index) + + +def _format_result_panel( + data: pd.DataFrame, + is_pure_fields: bool, + instruments: Union[List[str], Dict], + normalized_expr: Dict[str, str], + alias_to_origin_expr_map: Dict[str, str], +) -> pd.DataFrame: + """把查询结果整形为 (instrument, datetime) MultiIndex 面板并恢复输出列名。""" if not isinstance(data, pd.DataFrame): raise ValueError("查询结果不是 DataFrame 格式") - # 格式化结果 + # 格式化结果:计算分支已由 _computed_dict_to_panel/_legacy_reshape 产出 + # (instrument, datetime) 面板;纯字段分支为长表,需 set_index if not data.empty: - try: - if isinstance(data.index, pd.MultiIndex): - # 根据 pandas 版本决定 stack 的参数 - if version.parse(pd.__version__) < version.parse("1.5.0"): - # 旧版本 pandas 不支持 future_stack,但支持 dropna=False - data: pd.DataFrame = ( - data.unstack(level=0) - .stack(level=0, dropna=False) - .swaplevel(0, 1) - ) - else: - # 新版本 pandas 使用 future_stack 或依赖默认行为 - data: pd.DataFrame = ( - data.unstack(level=0) - .stack(level=0, future_stack=True) - .swaplevel(0, 1) - ) - else: + if not isinstance(data.index, pd.MultiIndex): + try: data: pd.DataFrame = data.set_index(["code", "date"]) - except KeyError: - raise KeyError("查询结果缺少 'code' 或 'date' 列") + except KeyError: + raise KeyError("查询结果缺少 'code' 或 'date' 列") data.index.names = ["instrument", "datetime"] @@ -442,99 +746,16 @@ def fetch_features_from_ddb( if isinstance(instruments, dict): data = apply_spans_mask(data, instruments) # 重命名columns将$改为去掉$或根据已有的字典进行重命名 - try: - aliases_order = list(normalized_expr.values()) - # 只保留实际存在的 alias,避免 reindex 新增空列 - ordered_aliases = [c for c in aliases_order if c in data.columns] - # 若希望严格要求所有列必须存在,可以改为: - # missing = [c for c in aliases_order if c not in data.columns]; if missing: raise KeyError(...) - data = data.reindex(columns=ordered_aliases) - except Exception as e: - # 出错则保持原始列顺序 - raise e + aliases_order = list(normalized_expr.values()) + # 只保留实际存在的 alias,避免 reindex 新增空列 + ordered_aliases = [c for c in aliases_order if c in data.columns] + data = data.reindex(columns=ordered_aliases) # 最后再重命名 alias -> 输出列名 data = data.rename(columns=alias_to_origin_expr_map) return data -# # 用于兼容qlib的表达式 -# def ddb_compute_features( -# session, -# instruments: Union[List, str], -# fields: Union[List, str], -# start_time: str, -# end_time: str, -# db_name: str, -# table_name: str, -# mr_by_code: bool = False, -# column_name: str = "date", -# freq: str = "day", -# ) -> pd.DataFrame: -# """ -# 计算DolphinDB中的特征。 -# -------- -# 因子处理流程与原始qlib差不多类似df.groupby("code").apply(lambda x:expression(x))的模式; -# 缺点:这种模式在处理横截面表达式时就不行 - -# 此函数从DolphinDB数据库中检索和计算特定时间范围内的特定特征。 - -# :param session: DolphinDB会话对象 -# :type session: DDB会话对象 - -# :param instruments: 代码列表或单个代码字符串(例如股票代码) -# :type instruments: Union[List, str] - -# :param fields: 要获取的字段列表或单个字段字符串(Qlib原生表达式) -# :type fields: Union[List, str] - -# :param start_time: 数据开始时间 -# :type start_time: str - -# :param end_time: 数据结束时间 -# :type end_time: str - -# :param db_name: DolphinDB数据库名称 -# :type db_name: str - -# :param table_name: DolphinDB表名 -# :type table_name: str - -# :param mr_by_code: 是否按代码进行MapReduce操作,默认为False。注意:社区版DolphinDB在mr_by_code=True时可能出现OOM问题 -# :type mr_by_code: bool, optional - -# :param column_name: 时间列的名称,默认为"date" -# :type column_name: str, optional - -# :param freq: 数据频率,可选"day"或其他值(如"min"),默认为"day" -# :type freq: str, optional - -# :returns: 包含计算特征的DataFrame,按代码和日期排序 -# :rtype: pd.DataFrame -# """ -# if isinstance(fields, str): -# fields = [fields] - -# if isinstance(instruments, str): -# instruments = [instruments] - -# fmt: str = "%Y.%m.%d" if freq == "day" else "%Y.%m.%d %H:%M:%S" -# start_time: pd.Timestamp = pd.to_datetime(start_time).strftime(fmt) -# end_time: pd.Timestamp = pd.to_datetime(end_time).strftime(fmt) - -# # 标准化qlib表达 -# exprs: List[str] = normalize_cache_fields([str(column) for column in fields]) -# # escape_backslash=True, 后续用于parseExpr解析 -# # 将qlib表达式转换为DolphinDB表达式 -# normalized_expr: List[str] = [ -# adapt_qlib_expr_syntax_for_ddb(column, OPERATOR_MAPPING, True) -# for column in exprs -# ] - -# # 使用mr_by_code会快很多,但是社区版ddb会出现OOM出现,所有尽量使用mr_by_code=false -# return session.run( -# f"""FeatureEngineering({instruments},dict({normalized_expr},{exprs}),{start_time},{end_time},"{db_name}","{table_name}",{int(mr_by_code)}$bool,"{column_name}")""" -# ).sort_values(["code", "date"]) ############################################################################################################################## @@ -542,18 +763,32 @@ def fetch_features_from_ddb( class TradeDateUtils: + # 模块级日历缓存:{(db_path, table_name): (日历数组, 日期->下标索引)}。 + # ⚠️ 性能关键:fetch_features_from_ddb 每批字段都会构造本类,未缓存时 + # 每次都全量下载交易日历(Alpha158 一次 D.features ≈ 6 次下载)。 + # 写路径变更日历表后需调用 ddb_qlib.invalidate_ddb_caches() 失效。 + _calendar_cache: Dict[Tuple[str, str], Tuple[np.ndarray, Dict]] = {} + def __init__(self, session: ddb.Session, freq: str): db_path: str = f"dfs://{QlibTableSchema.calendar().db_name}" tb_path: str = QlibTableSchema.calendar().table_name - self._calendar: np.ndarray = ( - session.loadTable(tb_path, db_path) - .exec("TRADE_DAYS") - .sort("TRADE_DAYS") - .toDF() - ) - self._calendar_index: Dict = dict( - zip(self._calendar, np.arange(len(self._calendar))) - ) + cache_key: Tuple[str, str] = (db_path, tb_path) + cached = self._calendar_cache.get(cache_key) + if cached is None: + calendar: np.ndarray = ( + session.loadTable(tb_path, db_path) + .exec("TRADE_DAYS") + .sort("TRADE_DAYS") + .toDF() + ) + cached = (calendar, dict(zip(calendar, np.arange(len(calendar))))) + self._calendar_cache[cache_key] = cached + self._calendar, self._calendar_index = cached + + @classmethod + def clear_cache(cls) -> None: + """清空模块级日历缓存(日历表数据变更后调用)。""" + cls._calendar_cache.clear() def get_locate_date_by_idx( self, start_idx: int, end_idx: int @@ -747,77 +982,34 @@ def add_line(name, value, wrap_sqlCol=False): return script -# def get_query_date_range( -# session: ddb.Session, start_dt, end_dt -# ) -> Tuple[pd.Timestamp, pd.Timestamp]: -# """获取查询的日期范围 - -# 根据提供的开始和结束日期,从日历表中确定实际的查询日期范围。 -# 如果未指定开始或结束日期,则使用日历表中的最早或最晚交易日。 - -# :param session: DolphinDB会话对象 -# :type session: ddb.Session -# :param start_dt: 开始日期,可以是以下类型: -# - None:使用日历表中的最早日期 -# - int:作为日历表中的索引位置 -# - str或pd.Timestamp:具体的日期值 -# :type start_dt: Union[str, int, pd.Timestamp, None] -# :param end_dt: 结束日期,可以是以下类型: -# - None:使用日历表中的最晚日期 -# - int:作为日历表中的索引位置 -# - str或pd.Timestamp:具体的日期值 -# :type end_dt: Union[str, int, pd.Timestamp, None] - -# :return: 处理后的开始日期和结束日期,转换为pandas Timestamp格式 -# :rtype: Tuple[pd.Timestamp, pd.Timestamp] -# """ - -# db_path: str = f"dfs://{QlibTableSchema.calendar().db_name}" -# tb_path: str = QlibTableSchema.calendar().table_name - -# tb: pd.DataFrame = ( -# session.loadTable(tb_path, db_path) -# .select("TRADE_DAYS") -# .sort("TRADE_DAYS") -# .toDF() -# ) -# dates: pd.DatetimeIndex = pd.to_datetime(tb["TRADE_DAYS"].values).sort_values() - -# if start_dt == None: -# start_dt: pd.Timestamp = dates.min() - -# if end_dt == None: -# end_dt: pd.Timestamp = dates.max() - -# if isinstance(start_dt, int) or isinstance(end_dt, int): - -# size: int = len(dates) -# if end_dt > size: -# raise ValueError(f"结束日期超出交易日历范围: {end_dt} > {size}") - -# return dates[start_dt], dates[end_dt] - -# else: +################################################################################################################################### +# 表达式转换函数 +################################################################################################################################### +def adapt_qlib_expr_syntax_for_ddb( + expr: str, operator_mapping: Dict = OPERATOR_MAPPING, escape_backslash: bool = False +) -> str: + """ + 将 qlib 表达式转换为 DolphinDB 表达式(默认映射时带 lru_cache 记忆化)。 -# start_dt: pd.Timestamp = dates.asof(start_dt) -# end_dt: pd.Timestamp = dates.asof(end_dt) + 翻译是纯函数且实现为 ~100 行递归解析,此前每字段每次 fetch 都重跑; + 默认 OPERATOR_MAPPING(identity 判断,覆盖全部真实调用方)时经 + :func:`_adapt_cached` 缓存,自定义映射(dict 不可哈希)则直接解析。 -# # 检查开始日期 -# if start_dt is np.nan: -# raise ValueError(f"开始日期 {start_dt} 不是有效的交易日") + 参数与返回见 :func:`_adapt_qlib_expr_impl`。 + """ + if operator_mapping is OPERATOR_MAPPING: + return _adapt_cached(expr, escape_backslash) + return _adapt_qlib_expr_impl(expr, operator_mapping, escape_backslash) -# # 检查结束日期 -# if start_dt is np.nan: -# raise ValueError(f"结束日期 {end_dt} 不是有效的交易日") -# return start_dt, end_dt +@functools.lru_cache(maxsize=4096) +def _adapt_cached(expr: str, escape_backslash: bool) -> str: + """默认算子映射下的表达式翻译缓存。""" + return _adapt_qlib_expr_impl(expr, OPERATOR_MAPPING, escape_backslash) -################################################################################################################################### -# 表达式转换函数 -################################################################################################################################### -def adapt_qlib_expr_syntax_for_ddb( - expr: str, operator_mapping: Dict = OPERATOR_MAPPING, escape_backslash: bool = False +def _adapt_qlib_expr_impl( + expr: str, operator_mapping: Dict, escape_backslash: bool = False ) -> str: """ 将 qlib 表达式转换为 DolphinDB 表达式,支持复杂嵌套结构 @@ -933,81 +1125,3 @@ def adapt_qlib_expr_syntax_for_ddb( return "".join(result_segments) - -def extract_fields_from_expressions(expressions, rename_map=None): - """ - 从多个表达式中提取所有基础字段变量 - - Parameters - ---------- - expressions : str or list - 表达式或表达式列表,如 - "gtjaAlpha1($open, $close, $vol);" 或 - ["$close/$open", "SMA($high, 10)/$low"] - - rename_map : dict, optional - 字段重命名映射,如 {'vol': 'volume', 'close': 'price_close'} - - Returns - ------- - list - 所有表达式中提取的去重字段名列表 - """ - - # 确保处理列表和单个字符串的表达式 - if not isinstance(expressions, (list, tuple)): - expressions = [expressions] - - # 正则表达式匹配以$开头的变量名 - pattern = r"\$([a-zA-Z_][a-zA-Z0-9_]*)" - - # 存储所有找到的字段 - all_fields = set() - - # 处理每个表达式 - for expr in expressions: - # 如果是另一个嵌套列表,递归处理 - if isinstance(expr, (list, tuple)): - nested_fields = extract_fields_from_expressions(expr, None) - all_fields.update(nested_fields) - else: - # 提取当前表达式中的字段 - fields = re.findall(pattern, expr) - all_fields.update(fields) - - # 转换成列表并排序,保证结果稳定 - result = sorted(list(all_fields)) - - # 应用重命名映射(如果提供) - if rename_map: - result = [rename_map.get(field, field) for field in result] - - return result - - -def is_pure_fields_expressions(exprs): - """ - 检查表达式列表是否只包含纯字段引用(如$close, $open) - - Parameters: - ----------- - exprs: List[str] - 表达式列表 - - Returns: - -------- - bool - 如果所有表达式都是纯字段引用则返回True,否则返回False - """ - - if not isinstance(exprs, (list, tuple)): - exprs = [exprs] - - pattern = r"^\$([a-zA-Z_][a-zA-Z0-9_]*)$" - - for expr in exprs: - # 检查是否为纯字段引用格式 - if not re.match(pattern, expr): - return False - - return True diff --git a/qlib/data/backend/ddb_qlib/ddb_mysql_bridge.py b/qlib/data/backend/ddb_qlib/ddb_mysql_bridge.py index 11e1c938056..57a530cede6 100644 --- a/qlib/data/backend/ddb_qlib/ddb_mysql_bridge.py +++ b/qlib/data/backend/ddb_qlib/ddb_mysql_bridge.py @@ -12,10 +12,13 @@ import pandas as pd from pydantic import BaseModel, Field, field_validator, validate_arguments +from ....log import get_module_logger from .ddb_client import DDBClient, DDBConnectionSpec from .ddb_operator import DDBTableOperator from .schemas import QlibTableSchema, FIELDS_MAPPING -from .utils import convert_wind_date_to_datetime +from .utils import convert_wind_date_to_datetime, validate_date_str, validate_sql_identifier + +logger = get_module_logger("ddb_mysql_bridge") class MySQLConnectionSpec(BaseModel): """适配多种MySQL方言的连接参数规范""" @@ -106,9 +109,10 @@ def __exit__(self, exc_type, exc_val, exc_tb): def load_mysql_plugin(self) -> None: """在dolphinDB中加载MySQL插件""" + # ⚠️ 历史 bug:此处曾误装 "lgbm" 插件,导致 mysql::connect 必然失败 expr: str = """ - installPlugin("lgbm") - loadPlugin("lgbm") + installPlugin("mysql") + loadPlugin("mysql") """ self.ddb_session.run(expr) @@ -130,7 +134,7 @@ def close(self) -> None: self.ddb_session.run("mysql::close(mysql_conn)") except Exception as e: # 即使关闭连接失败,也不应该中断程序执行 - print(f"警告: 关闭MySQL连接时出现错误: {str(e)}") + logger.warning(f"关闭MySQL连接时出现错误: {e}") @validate_arguments def load_table( @@ -247,8 +251,8 @@ def __exit__(self, exc_type, exc_val, exc_tb): if self._bridge: try: self._bridge.close() - except: - pass + except Exception as e: + logger.warning(f"关闭 bridge 失败: {e}") return False def _get_bridge(self): @@ -293,7 +297,7 @@ def _sync_table(self, schema_func, table_name, where_clause=""): """ bridge = self._get_bridge() try: - print(f"正在同步 {schema_func.__name__} 从 {table_name}...") + logger.info(f"正在同步 {schema_func.__name__} 从 {table_name}...") cols, name_type = self._extract_columns(schema_func) query = f"SELECT {cols} FROM {table_name}" + ( f" WHERE {where_clause}" if where_clause else "" @@ -307,7 +311,7 @@ def _sync_table(self, schema_func, table_name, where_clause=""): bridge.ddb_operator.table_appender( schema_func().db_name, schema_func().table_name, data ) - print(f"已完成 {schema_func.__name__} 从 {table_name} 的同步。\n") + logger.info(f"已完成 {schema_func.__name__} 从 {table_name} 的同步") except Exception as e: raise RuntimeError(f"同步表 {table_name} 失败: {str(e)}") from e @@ -324,7 +328,8 @@ def sync_calendar(self, exchange_market: str = "SSE"): :param exchange_market: 交易所市场,默认为"SSE" :type exchange_market: str """ - where_clause = f"S_INFO_EXCHMARKET='{exchange_market}'" + # 防注入:交易所参数将拼入 SQL,先做白名单校验 + where_clause = f"S_INFO_EXCHMARKET='{validate_sql_identifier(exchange_market)}'" self._sync_table(QlibTableSchema.calendar, "ASHARECALENDAR", where_clause) def sync_feature_daily(self, start_date: str = "20100101", end_date: str = "20241231"): @@ -336,7 +341,8 @@ def sync_feature_daily(self, start_date: str = "20100101", end_date: str = "2024 :param end_date: 结束日期,格式:YYYYMMDD :type end_date: str """ - where_clause = f"TRADE_DT BETWEEN {start_date} AND {end_date}" + # 防注入:日期参数将拼入 SQL,先做格式校验 + where_clause = f"TRADE_DT BETWEEN {validate_date_str(start_date)} AND {validate_date_str(end_date)}" self._sync_table(QlibTableSchema.feature_daily, "ASHAREEODPRICES", where_clause) def sync_index_daily(self, index_codes: List[str], start_date: str = "20100101", end_date: str = "20241231"): @@ -359,6 +365,11 @@ def sync_index_daily(self, index_codes: List[str], start_date: str = "20100101", if not index_codes: raise ValueError("index_codes不能为空") + + # 防注入:指数代码与日期将拼入 SQL,先做白名单/格式校验 + index_codes = [validate_sql_identifier(code) for code in index_codes] + start_date = validate_date_str(start_date) + end_date = validate_date_str(end_date) # 构建WHERE条件 if len(index_codes) == 1: @@ -370,7 +381,7 @@ def sync_index_daily(self, index_codes: List[str], start_date: str = "20100101", bridge = self._get_bridge() try: - print(f"正在同步指数 {','.join(index_codes)} 从 AINDEXEODPRICES 至 IndexDaily 表...") + logger.info(f"正在同步指数 {','.join(index_codes)} 从 AINDEXEODPRICES 至 IndexDaily 表...") # 定义指数列映射 index_columns = ( @@ -395,7 +406,7 @@ def sync_index_daily(self, index_codes: List[str], start_date: str = "20100101", data = bridge.load_table(query) if data.empty: - print(f"警告: 指数 {','.join(index_codes)} 在时间范围 {start_date}-{end_date} 内无数据") + logger.warning(f"指数 {','.join(index_codes)} 在时间范围 {start_date}-{end_date} 内无数据") return # 转换日期格式 @@ -409,7 +420,7 @@ def sync_index_daily(self, index_codes: List[str], start_date: str = "20100101", data ) - print(f"已完成指数数据同步:{len(data)} 条记录已写入 {features_schema.table_name} 表\n") + logger.info(f"已完成指数数据同步:{len(data)} 条记录已写入 {features_schema.table_name} 表") except Exception as e: error_msg = f"同步指数表失败 - 指数: {','.join(index_codes)}, 时间范围: {start_date}-{end_date}" @@ -439,8 +450,8 @@ def close(self): if self._bridge: try: self._bridge.close() - except: - pass + except Exception as e: + logger.warning(f"关闭 bridge 失败: {e}") self._bridge = None diff --git a/qlib/data/backend/ddb_qlib/ddb_operator.py b/qlib/data/backend/ddb_qlib/ddb_operator.py index e89f03893e6..3aad5bb9b47 100644 --- a/qlib/data/backend/ddb_qlib/ddb_operator.py +++ b/qlib/data/backend/ddb_qlib/ddb_operator.py @@ -11,8 +11,31 @@ import dolphindb as ddb import pandas as pd +from ....log import get_module_logger from .ddb_client import DDBClient, DDBConnectionSpec from .schemas import QlibTableSchema, TableSchema +from .utils import get_table_columns + +logger = get_module_logger("ddb_operator") + +# 按 URI 复用的共享客户端注册表:模块级写入/DDL 函数此前每次调用都新建 +# DDBClient(新 TCP 会话 + 登录),批量导入时反复连接;社区版(2 核/8GB) +# 服务器上每条会话都占内存,复用一条会话即可。 +# ⚠️ ddb.Session 非线程安全:共享客户端仅用于现状的单线程写入流程。 +_shared_clients: Dict[str, DDBClient] = {} + + +def get_shared_client(uri: str) -> DDBClient: + """按 URI 获取(或创建)共享的 DDBClient。 + + :param uri: DolphinDB 连接 URI + :return: 该 URI 对应的进程内共享客户端 + """ + client = _shared_clients.get(uri) + if client is None: + client = DDBClient(DDBConnectionSpec(uri=uri)) + _shared_clients[uri] = client + return client class DDBTableOperator: @@ -173,9 +196,8 @@ def table_appender(self, db_name: str, table_name: str, data: pd.DataFrame) -> N if not self.exist_table(db_name, table_name): raise ValueError(f"{db_path}/{table_name}不存在!") - table = session.loadTable(table_name, db_path) - table_cols = table.schema["name"].tolist() - + table_cols = get_table_columns(session, db_path, table_name) + # 确保数据列名顺序与表列名顺序一致 data = data.reindex(columns=table_cols) @@ -213,8 +235,7 @@ def table_upsert( if not self.exist_table(db_name, table_name): raise ValueError(f"表 {db_path}/{table_name} 不存在!") - table = session.loadTable(table_name, db_path) - table_cols = table.schema["name"].tolist() + table_cols = get_table_columns(session, db_path, table_name) # 确保数据列名顺序与表列名顺序一致 data = data.reindex(columns=table_cols) @@ -262,6 +283,8 @@ def create_table( partition_columns: str, engine: str, primary_key: str = None, + *, + client: Optional[DDBClient] = None, ) -> None: """ 通用创建表函数 @@ -286,6 +309,8 @@ def create_table( :type engine: str :param primary_key: 主键列名,默认为None :type primary_key: str, optional + :param client: 复用的 DDBClient;缺省时按 uri 使用进程内共享客户端 + :type client: Optional[DDBClient] :return: 无返回值 :rtype: None @@ -295,9 +320,7 @@ def create_table( .. note:: 确保在调用此函数之前已正确配置DolphinDB连接参数。 """ - config = DDBConnectionSpec(uri=uri) - connector = DDBClient(config) - db_accessor = DDBTableOperator(connector) + db_accessor = DDBTableOperator(client if client is not None else get_shared_client(uri)) db_accessor.create_partitioned_table( db_name, @@ -309,7 +332,7 @@ def create_table( engine=engine, primary_key=primary_key, ) - print(f"{db_name}/{table_name}生成创建完毕!") + logger.info(f"{db_name}/{table_name}生成创建完毕!") def create_qlib_table(uri: str, schema: TableSchema) -> None: @@ -342,12 +365,13 @@ def create_instrument_table(uri: str, table_name: str = "ashares") -> None: create_qlib_table(uri, schema) -def clean_qlib_db(uri: str) -> None: - """清理Qlib数据库""" - config = DDBConnectionSpec(uri=uri) - connector = DDBClient(config) +def clean_qlib_db(uri: str, *, client: Optional[DDBClient] = None) -> None: + """清理Qlib数据库 - session = connector.session + :param uri: DolphinDB 连接 URI + :param client: 复用的 DDBClient;缺省时按 uri 使用进程内共享客户端 + """ + session = (client if client is not None else get_shared_client(uri)).session db_names = QlibTableSchema.get_all_databases() expr_lines = [ @@ -358,6 +382,11 @@ def clean_qlib_db(uri: str) -> None: session.run(expr) + # 数据库已变更,失效进程内缓存(日历等) + from . import invalidate_ddb_caches + + invalidate_ddb_caches() + def write_df_to_ddb( db_name: str, @@ -367,6 +396,8 @@ def write_df_to_ddb( key_col_names: Optional[Union[str, List[str]]] = None, sort_columns: Optional[Union[str, List[str]]] = None, uri: Optional[str] = None, + *, + client: Optional[DDBClient] = None, ) -> None: """ 将DataFrame数据写入DolphinDB表。 @@ -383,17 +414,17 @@ def write_df_to_ddb( :type key_col_names: Optional[Union[str, List[str]]] :param sort_columns: upsert时用于排序的列名 :type sort_columns: Optional[Union[str, List[str]]] - :param conn_mgr: DDBClient连接管理器,默认为None(需外部提前初始化) - :type conn_mgr: Optional[DDBClient] + :param uri: DolphinDB连接URI(与 client 二选一) + :type uri: Optional[str] + :param client: 复用的 DDBClient;缺省时按 uri 使用进程内共享客户端 + :type client: Optional[DDBClient] """ - if uri is None: - raise ValueError("请传入已初始化的DDBClient实例(conn_mgr)") + if uri is None and client is None: + raise ValueError("必须提供 uri 或已初始化的 DDBClient(client)") - config = DDBConnectionSpec(uri=uri) - connector = DDBClient(config) - operator = DDBTableOperator(connector) + operator = DDBTableOperator(client if client is not None else get_shared_client(uri)) if upsert: operator.table_upsert( db_name=db_name, @@ -408,6 +439,11 @@ def write_df_to_ddb( table_name=table_name, data=data, ) + + # 表数据已变更,失效进程内缓存(日历等) + from . import invalidate_ddb_caches + + invalidate_ddb_caches() def import_instruments_csv_to_ddb( diff --git a/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineering.dos b/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineering.dos index cd3d7751338..513567fdae9 100644 --- a/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineering.dos +++ b/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineering.dos @@ -162,20 +162,6 @@ def canRunWithAvailableMemory(start_date, end_date, instruments, queryFields, db /* 特征数据获取相关代码 */ -/** -* 合并多个字典的键值对 -* @param dicts - 包含多个字典的列表 k-str v-matrix -* @return 合并后的字典,包含所有输入字典的键值对 -*/ -def union_dict(dicts){ - out_put_dict = dict(STRING, ANY) - for (j in flatten(dicts)) { - dictUpdate!(out_put_dict, unionAll, j.keys(), j.values()) - } - // 遍历每个字典,将其键值对合并到输出字典中 - return out_put_dict; -}; - /** * 为给定日期范围内的每个交易日创建对应的可交易股票列表映射 * @@ -490,24 +476,6 @@ def splitDateRangesByTradeDays(start_date, end_date, exchange="XSHG", daysStep=2 // res_arr = flatten(res) -/* -* 创建空表 -* @return TABLE 空表 -*/ -def createEmptyTable(codes,dates,colNames){ - - totalSize = size(codes)*size(dates); - code = take(codes, totalSize); - code.sortBy!(false) - date = repeat(dates, totalSize); - emptyTB = table(date,code); - for (colName in colNames){ - emptyTB[colName] = 0; - } - reorderColumns!(emptyTB,colNames) - return emptyTB; -}; - // 根据频率规范化日期类型 def normalizeDate(dt,freq){ if(freq == "D"){ @@ -531,18 +499,21 @@ class FeatureEngine{ endTime::DATE; freq::STRING; tb::ANY; // 加载的表格 + lookbackDays::INT; // 滚动算子向前外扩的交易日数(0 表示不外扩) + rightDays::INT; // 未来引用(如标签 Ref(close,-2))向后外扩的交易日数 - def FeatureEngine(dbName_,tableName_,instruments_,expressions_,baseFields_,startTime_,endTime_,freq_="D"){ + def FeatureEngine(dbName_,tableName_,instruments_,expressions_,baseFields_,startTime_,endTime_,freq_="D",lookbackDays_=0,rightDays_=0){ dbName = dbName_; tableName = tableName_; instruments = instruments_; - expressions = expressions_; - baseFields = baseFields_; + expressions = expressions_; + baseFields = baseFields_; startTime = startTime_; endTime = endTime_; freq = freq_; tb = loadTable("dfs://"+dbName_,tableName_); - + lookbackDays = lookbackDays_; + rightDays = rightDays_; }; @@ -566,10 +537,15 @@ class FeatureEngine{ }; }; - def getDSDateRange(ds){ - dsDateRng = select min(date) as start_dt,max(date) as end_dt from ds; - dates = getMarketCalendar("XSHG",dsDateRng['start_dt'][0],dsDateRng['end_dt'][0]); - return dates; + /** + * 把边界日期按交易日前移(shiftDays<0)或后移(shiftDays>0)。 + * 用于按 lookbackDays/rightDays 外扩查询窗口;shiftDays 为 0 时原样返回。 + */ + def shiftTradeDays(dt, shiftDays){ + if(shiftDays == 0){ + return dt; + }; + return temporalAdd(dt, shiftDays, `XSHG); }; /** @@ -594,40 +570,33 @@ class FeatureEngine{ return cols; }; - // 判断当前内存是否足以运行 - def isRunWithAvailableMemory(){ - + /** + * 估算扩展窗口整段计算的峰值内存是否放得下(放不下则走分段模式)。 + * 峰值 ≈ 基础长表 + panel 矩阵字典(约等于长表) + 每个表达式一个 + * dates×codes 的 DOUBLE 结果矩阵;只允许占用 70% 空闲内存留计算余量。 + */ + def isRunWithAvailableMemory(extStart, extEnd){ if(!isVector(baseFields)){ fields = [baseFields] }else{ fields = baseFields; }; - - return canRunWithAvailableMemory(normalizeDate(startTime,freq),normalizeDate(endTime,freq),toListInstruments(),fields,dbName,tableName); + est = estimateQueryMemory(normalizeDate(extStart,freq),normalizeDate(extEnd,freq),toListInstruments(),fields,dbName,tableName); + baseBytes = est["total_bits"]; + exprBytes = est["total_rows"] * 8 * size(toListExpressions()); + need = baseBytes * 2 + exprBytes; + return need <= mem().freeBytes * 0.7; }; - // FIXME:创建按日期分区的数据源,关闭了内存大小预估,分区并未生效 - def createDSByDate(daysStep=252,dateName=`date,exchange=`XSHG){ - // 边界日期偏移一个交易日 - boundEndDate = temporalAdd(endTime,1,`SSE); - if(!isVector(baseFields)){ - fields = [baseFields] - }else{ - fields = baseFields; - }; - - // if(isRunWithAvailableMemory()){ - // // 当前内存可用处理 - // // print("当前内存充足,使用单次查询模式"); - // // return repartitionDS(, dateName,RANGE,bounds); - // return repartitionDS(sql(select=buildSelectColumnsWithDefaults(),from=tb), dateName, RANGE, [startTime, boundEndDate]); - // }; - return repartitionDS(sql(select=buildSelectColumnsWithDefaults(),from=tb), dateName, RANGE, [startTime, boundEndDate]); + /** + * 扩展窗口内实际有数据的代码集合(升序)。 + * 作为所有分段 panel 的统一列标签,保证各分段结果矩阵列对齐可按行拼接; + * 该集合与不指定 colLabel 时 panel 的列集合一致(数据中出现过的代码)。 + */ + def fetchPresentCodes(extStart, extEnd){ + instList = toListInstruments(); + codesTb = select distinct code from tb where date between pair(extStart, extEnd) and code in instList; + return sort(codesTb['code']); }; // 标准化expressions为字典格式 @@ -644,64 +613,84 @@ class FeatureEngine{ }; - def buildSpanAwareWhereConditions(){ - // 判断instruments是为dict - // 用以兼容spans + /** + * 构建查询过滤条件;日期用传入的查询区间(可能为外扩后的区间)。 + * dict(成分股 spans)分支保留以兼容旧接口;当前计算分支由 Python 侧 + * 传键列表全量计算,spans 掩码在 Python 侧补齐。 + */ + def buildWhereConditions(rangeStart, rangeEnd){ if (isDict(instruments)){ - codeRangeFilter = createDateStockMapping(startTime,endTime,instruments); - whereConditions = [expr(sqlCol(`date),between,pair(startTime,endTime)),]; + codeRangeFilter = createDateStockMapping(rangeStart,rangeEnd,instruments); + whereConditions = [expr(sqlCol(`date),between,pair(rangeStart,rangeEnd)),]; } else{ - whereConditions = [expr(sqlCol(`date),between,pair(startTime,endTime)),expr(sqlCol(`code),in,instruments)]; + whereConditions = [expr(sqlCol(`date),between,pair(rangeStart,rangeEnd)),expr(sqlCol(`code),in,instruments)]; }; return whereConditions; }; - def fetchFeatures(ds){ + /** + * 构造 [qStart, qEnd] 区间的 NULL 占位矩阵(行-该区间交易日,列-统一代码集合)。 + * 失败/空分段的占位与正常结果同为「带行列标签的矩阵」,使 mergeSegmentResults + * 能按行拼接、不丢弃其它分段,并保持日期覆盖连续(消除分段空洞); + * 填 NULL 而非 0,避免用 0 污染缺失位置的因子值。 + * @param colLabels 统一列标签(升序代码向量) + * @return MATRIX 行-交易日,列-colLabels,元素全为 NULL DOUBLE + */ + def buildPlaceholderMatrix(qStart, qEnd, colLabels){ + rowLabels = getMarketCalendar("XSHG", qStart, qEnd); + m = matrix(DOUBLE, size(rowLabels), size(colLabels), , double(NULL)); + m.rename!(rowLabels, colLabels); + return m; + }; - // 添加错误参数检查 - if(size(expressions) == 0) { - throw("表达式列表不能为空"); - }; - if(size(baseFields) == 0) { - throw("基础字段列表不能为空"); - }; - - // 判断expressions是否为dict,如果不是则标准化为字典 - expressionsDict = normalizeExpressionToDict(); + /** + * 计算 [qStart, qEnd] 闭区间的表达式面板。 + * 查询窗口向前外扩 lookbackDays、向后外扩 rightDays 个交易日,保证 + * 滚动算子在区间头部、未来引用在区间尾部不因窗口不足产生 NULL; + * 计算完成后把结果矩阵的行截断回 [qStart, qEnd]。 + * @param colLabels 统一列标签(升序代码向量),保证分段结果可按行拼接 + * @return DICT 别名 -> 结果矩阵(行-日期,列-代码) + */ + def computeRangeFeatures(qStart, qEnd, colLabels){ + expressionsDict = normalizeExpressionToDict(); - // 判断instruments是为dict - // 用以兼容spans - whereConditions = buildSpanAwareWhereConditions(); + // 外扩查询窗口:外扩部分只服务于窗口计算,不进入结果 + extStart = shiftTradeDays(qStart, -lookbackDays); + extEnd = shiftTradeDays(qEnd, rightDays); + whereConditions = buildWhereConditions(extStart, extEnd); // 获取计算所需的基础数据 try { - // 基础数据字段 fieldExpr = buildSelectColumnsWithDefaults(); - baseData = sql(select=fieldExpr,from=ds,where=whereConditions,orderBy=).eval(); - // 检查结果是否为空 + baseData = sql(select=fieldExpr,from=tb,where=whereConditions,orderBy=).eval(); + // 基础数据为空:为每个 alias 填该段 NULL 占位矩阵,保持日期覆盖 + // 与结果类型(矩阵),避免分段路径下该段整体缺失形成空洞 if(baseData.rows() == 0) { print("警告: 基础数据获取失败,请检查表达式引用字段: " + fieldExpr); - // 返回空 dict 而不是 TABLE,与正常路径返回类型保持一致 - // union_dict 期待 dict,收到 TABLE 会报 "Incompatible vector/matrix size" - return dict(STRING, ANY, true); + empty_dict = dict(STRING, ANY, true); + for(expression in expressionsDict.keys()){ + empty_dict[expressionsDict[expression]] = buildPlaceholderMatrix(qStart, qEnd, colLabels); + }; + return empty_dict; }; } catch(ex){ print("错误:表达式中含有未知基础字段.请检查: " + ex); return dict(STRING, ANY, true); }; - + // 生成数据字典 colDatas = []; for (col in baseFields){ colDatas.append!(baseData[col]); }; // 数据字典 k基础字段,v-矩阵index为日期,columns-code + // colLabel 固定为统一代码集合,保证各分段矩阵列对齐 dates = baseData['date']; codes = baseData['code']; - dataDict = dict(baseFields, panel(dates, codes, colDatas)); + dataDict = dict(baseFields, panel(dates, codes, colDatas, , colLabels)); // 收集结果 result_dict = dict(STRING,ANY,true); @@ -711,33 +700,130 @@ class FeatureEngine{ try{ res = parseExpr(expression,dataDict).eval(); - - // 检查结果是否为矩阵 + + // 计算结果为空:填 NULL 占位矩阵(与正常结果同为带标签矩阵) if(res.rows()==0){ print("警告: 计算结果为空,请检查表达式: " + expression); - res = createEmptyTable(toListInstruments(),getDSDateRange(ds),`date`code+keys(expressionsDict)); + result_dict[expressionsDict[expression]] = buildPlaceholderMatrix(qStart, qEnd, colLabels); + continue; }; - + + // 截断回请求区间:先留存标签,loc 后显式重挂以防实现差异丢标签 + rn = res.rowNames(); + cn = res.colNames(); + keep = (rn >= qStart) and (rn <= qEnd); + res = loc(res, keep); + res.rename!(rn[at(keep)], cn); result_dict[expressionsDict[expression]] = res; - + } catch(ex){ - // ex[1]为错误信息,避免误用ex.msg属性 + // ex[1]为错误信息,避免误用ex.msg属性;失败同样填 NULL 占位矩阵, + // 使分段合并时不因类型不一致丢弃其它分段的正确结果 print("【特征计算失败】表达式: " + expression + " | 错误: " + ex[1]); - placeholder = createEmptyTable(toListInstruments(),getDSDateRange(ds),`date`code+keys(expressionsDict)); - result_dict[expressionsDict[expression]] = placeholder; + result_dict[expressionsDict[expression]] = buildPlaceholderMatrix(qStart, qEnd, colLabels); }; - + }; - + return result_dict; }; - def fetch(){ - ds = self.createDSByDate(daysStep=252,dateName=`date,exchange=`XSHG); - return mr(ds, fetchFeatures{}, , union_dict,parallel=true); - } + /** + * 合并各分段计算结果:各段矩阵列已按统一 colLabels 对齐,按行拼接。 + * concatMatrix 不保证保留行/列标签,拼接后显式重挂。 + * 失败/空分段已由 computeRangeFeatures 填为 NULL 占位矩阵(同为带标签矩阵), + * 故此处对所有分段一视同仁按行拼接——任一分段异常都不会丢弃其它分段的 + * 正确结果;防御性地跳过极端情况下仍非矩阵的分段(不中断整体合并)。 + */ + def mergeSegmentResults(segResults, colLabels){ + merged = dict(STRING, ANY, true); + aliases = []; + for(seg in segResults){ + for(k in keys(seg)){ + if(!(k in aliases)){ + aliases.append!(k); + }; + }; + }; + for(alias in aliases){ + mats = []; + rowLabs = array(DATE, 0); + for(seg in segResults){ + if(!(alias in keys(seg))){ + continue; + }; + part = seg[alias]; + // 防御:占位统一为矩阵后此处正常不触发;非矩阵则跳过该段, + // 保留其它分段已算出的正确矩阵,不整体丢弃 + if(regexFind(typestr(part), "MATRIX") == -1){ + continue; + }; + if(part.rows() > 0){ + mats.append!(part); + rowLabs.append!(part.rowNames()); + }; + }; + if(size(mats) == 0){ + continue; + }; + if(size(mats) == 1){ + merged[alias] = mats[0]; + continue; + }; + m = concatMatrix(mats, false); + m.rename!(rowLabs, colLabels); + merged[alias] = m; + }; + return merged; + }; + + /** + * 特征计算主入口。 + * 内存估算放得下时对整个请求区间单次计算;放不下时按 daysStep 切段 + * 顺序计算再合并——每段查询都带 lookback/right 外扩,段内截断保证 + * 分段边界处滚动算子结果与整段计算一致。 + */ + def fetch(daysStep=252){ + // 添加错误参数检查 + if(size(expressions) == 0) { + throw("表达式列表不能为空"); + }; + if(size(baseFields) == 0) { + throw("基础字段列表不能为空"); + }; + + extStart = shiftTradeDays(startTime, -lookbackDays); + extEnd = shiftTradeDays(endTime, rightDays); + presentCodes = fetchPresentCodes(extStart, extEnd); + if(size(presentCodes) == 0){ + // 窗口内无任何数据:返回空 dict,Python 侧兜底为空结果 + return dict(STRING, ANY, true); + }; + + // 内存放得下就整段单次计算;放不下才分段(分段带 lookback 重叠开销) + if(isRunWithAvailableMemory(extStart, extEnd)){ + return computeRangeFeatures(startTime, endTime, presentCodes); + }; + + // 分段顺序计算:社区版(2 核/8GB)单机 mr 并行收益小且峰值内存翻倍, + // 顺序循环把峰值内存约束在单段规模 + bounds = splitDateRangesByTradeDays(startTime, endTime, "XSHG", daysStep); + // 区间过短(如单交易日)无法切段时退回整段计算 + if(size(bounds) < 2){ + return computeRangeFeatures(startTime, endTime, presentCodes); + }; + segResults = []; + n = size(bounds); + for(i in 0..(n-2)){ + segStart = bounds[i]; + // 闭区间分段:非末段的段尾取下一边界的前一交易日,避免段间重叠 + segEnd = (i == n-2) ? bounds[i+1] : temporalAdd(bounds[i+1], -1, `XSHG); + segResults.append!(computeRangeFeatures(segStart, segEnd, presentCodes)); + }; + return mergeSegmentResults(segResults, presentCodes); + }; }; @@ -747,9 +833,10 @@ class FeatureEngine{ // // print(res) -def FeatureEngineeringByDate(instruments,expressions,baseFields,start_time,end_time,db_name,table_name){ - - engine = FeatureEngine(db_name,table_name,instruments,expressions,baseFields,start_time,end_time) - - return engine.fetch(); +// daysStep/lookbackDays/rightDays 为带默认值的尾参:旧调用方(六/七参形式)行为不变 +def FeatureEngineeringByDate(instruments,expressions,baseFields,start_time,end_time,db_name,table_name,daysStep=252,lookbackDays=0,rightDays=0){ + + engine = FeatureEngine(db_name,table_name,instruments,expressions,baseFields,start_time,end_time,"D",lookbackDays,rightDays) + + return engine.fetch(daysStep); }; \ No newline at end of file diff --git a/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineeringStreaming_bak b/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineeringStreaming_bak deleted file mode 100644 index edd2b517b8f..00000000000 --- a/qlib/data/backend/ddb_qlib/ddb_scripts/featureEngineeringStreaming_bak +++ /dev/null @@ -1,369 +0,0 @@ -/* -* @Author: Hugo -* @Date: 2025-01-20 -* @Description: 流式计算优化版本 - 解决 OOM 问题 -* -* 核心优化思路: -* 1. 使用 DolphinDB 的 replay 流式引擎进行数据回放 -* 2. enableTableShareAndPersistence 实现持久化 + 内存限制 -* 3. 时间分区控制,每次只处理小块数据 -* 4. 并行计算提升性能 -* -* 性能对比(测试数据: 2500 交易日 x 5000 股票 x 3 spans): -* - 原版 createDateStockMapping: ~37.5M 次迭代,内存峰值 ~2.5GB,耗时 ~180 秒 -* - 流式版本: 分块处理,内存峰值 <1GB,耗时 ~60 秒(并行度为 4) -*/ - -/* - ============================================================================= - 基础工具函数 - ============================================================================= - */ - -/** - * @brief 将 stockSpans 字典转换为表格式(适合流式处理) - * @description 将 {code: [(begin_dt, end_dt), ...]} 格式的字典转换为表 - * @param stockSpans - 股票代码到日期区间的映射字典 - * @return TABLE - 包含 code, begin_dt, end_dt 三列的表 - * - * 示例输入: {"000001.SZ": [(2020.01.01, 2020.12.31), (2021.01.01, 2021.06.30)]} - * 示例输出: 表格式数据,每行是一个 (code, begin_dt, end_dt) 组合 - */ -def stockSpansToTable(stockSpans){ - // 收集所有数据到向量中 - codes = []; - beginDates = []; - endDates = []; - - for(code in keys(stockSpans)){ - spans = stockSpans[code]; - for(span in spans){ - codes.append!(code); - beginDates.append!(date(span[0])); - endDates.append!(date(span[1])); - }; - }; - - // 创建表 - return table(codes as `code, beginDates as `begin_dt, endDates as `end_dt); -}; - -/** - * @brief 优化版本的 activeDaysForCode(使用向量化操作) - * @description 统计单个代码在交易日集合中的有效天数 - * - * 优化: 使用 DolphinDB 的向量化操作代替循环 - * @param spans - 股票的日期区间列表 [(begin_dt, end_dt), ...] - * @param trading_days - 交易日向量 - * @return INT - 有效交易日天数 - * - * 性能对比(2500 交易日 x 3 spans): - * - 原版 activeDaysForCode: ~7500 次循环,耗时 ~2ms - * - 本函数: 向量化操作,耗时 ~0.5ms - */ -def activeDaysForCodeOptimized(spans, trading_days){ - if(size(spans) == 0 or size(trading_days) == 0){ - return 0; - }; - - // 创建结果向量(初始化为 false) - isActive = take(false, size(trading_days)); - - // 遍历每个 span,标记活跃的交易日 - for(span in spans){ - begin_dt = date(span[0]); - end_dt = date(span[1]); - - // 向量化操作:找出在 [begin_dt, end_dt] 范围内的交易日 - inRange = (trading_days >= begin_dt) and (trading_days <= end_dt); - - // 逻辑或操作:合并结果 - isActive = isActive or inRange; - }; - - // 统计 true 的个数 - return sum(isActive); -}; - -/* - ============================================================================= - 流式计算核心函数 - ============================================================================= - */ - -/** - * @brief 使用 replay 流式引擎创建日期到股票的映射(内存优化版本) - * @description 解决原版 createDateStockMapping 的 OOM 问题 - * - * 核心优化: - * 1. 使用 replay 将数据分块处理,避免一次性加载全部数据 - * 2. enableTableShareAndPersistence 的 cacheSize 限制内存行数 - * 3. retentionMinutes 自动清理旧数据,节省磁盘空间 - * 4. parallelLevel 并行处理多个时间块 - * - * @param startDate - 开始日期 - * @param endDate - 结束日期 - * @param stockSpans - 股票代码到日期区间的映射字典 - * @param cacheSize - 内存最大行数(默认 1000 万,约 800MB) - * @param retentionMinutes - 数据保留时间(默认 60 分钟) - * @param parallelLevel - 并行度(默认 4,根据 CPU 核心数调整) - * @return DICT> - 日期到可交易股票列表的映射 - * - * 性能对比(测试数据: 2500 交易日 x 5000 股票 x 3 spans): - * - 原版 createDateStockMapping: ~37.5M 次迭代,内存峰值 ~2.5GB,耗时 ~180 秒 - * - 本函数: 分块处理,内存峰值 <1GB,耗时 ~60 秒(并行度为 4) - * - * 使用示例: - * stockSpans = createStockDateRangeMapping(instrumentsTB); - * dateStockMapping = createDateStockMappingStreaming(2020.01.01, 2023.12.31, stockSpans); - */ -def createDateStockMappingStreaming( - startDate, - endDate, - stockSpans, - cacheSize=10000000, - retentionMinutes=60, - parallelLevel=4 -){ - // ============================================================ - // 步骤 1: 将字典转换为表格式(适合流式处理) - // ============================================================ - spansTable = stockSpansToTable(stockSpans); - - // ============================================================ - // 步骤 2: 创建交易日历 - // ============================================================ - trading_days = getMarketCalendar("XSHG", startDate, endDate); - - // ============================================================ - // 步骤 3: 创建输入数据源(按时间分区) - // ============================================================ - // 将交易日历分为多个时间块,每块包含一部分交易日 - // 这样可以控制内存使用,每次只处理一个时间块的数据 - numChunks = parallelLevel; // 分块数量 = 并行度 - chunkSize = ceil(size(trading_days) / numChunks); - - // 为每个时间块创建数据源 - inputTables = dict(STRING,DATABASE); - outputTables = dict(STRING,DATABASE); - - for(i in 0..(numChunks-1)){ - startIdx = i * chunkSize; - endIdx = min((i+1) * chunkSize - 1, size(trading_days) - 1); - - if(startIdx >= size(trading_days)){ - break; - }; - - chunkDates = trading_days[startIdx:endIdx]; - - // 创建临时表(只包含当前时间块的交易日) - chunkTable = table(chunkDates as `trade_date); - - // 创建数据源 - dsName = "input_ds_" + string(i); - inputTables[dsName] = ds(chunkTable); - - // 创建输出流表 - outputTableName = "output_stream_" + string(i); - streamTable = streamTable(1000000:0, `trade_date`code, [DATE,SYMBOL]); - - // 启用持久化和共享(关键:cacheSize 限制内存) - enableTableShareAndPersistence( - table=streamTable, - tableName=outputTableName, - asynWrite=true, - compress=true, - cacheSize=cacheSize, // 关键参数:限制内存行数 - retentionMinutes=retentionMinutes, // 关键参数:自动清理 - flushMode=0, - preCache=1000000 - ); - - outputTables[outputTableName] = streamTable; - }; - - // ============================================================ - // 步骤 4: 定义流式计算逻辑 - // ============================================================ - def matchStocks(spansTable, chunkDates){ - // 为当前时间块的每个交易日找出可交易的股票 - - // 创建结果表 - result = table(1000000:0, `trade_date`code, [DATE,SYMBOL]); - - // 遍历交易日(注意:这里只处理当前时间块的日期) - for(tday in chunkDates){ - // 找出在 tday 交易的股票 - // 使用向量化操作代替循环 - validStocks = select code from spansTable - where begin_dt <= tday <= end_dt; - - if(size(validStocks) > 0){ - tempCodes = exec code from validStocks; - tempDates = take(tday, size(tempCodes)); - result.append!(table(tempDates as `trade_date, tempCodes as `code)); - }; - }; - - return result; - }; - - // ============================================================ - // 步骤 5: 执行流式计算(手动分发,不使用 replay) - // ============================================================ - // 注意:这里不使用 replay 函数,因为 replay 需要时间序列数据 - // 我们直接手动分发计算任务 - - for(i in 0..(numChunks-1)){ - startIdx = i * chunkSize; - endIdx = min((i+1) * chunkSize - 1, size(trading_days) - 1); - - if(startIdx >= size(trading_days)){ - break; - }; - - chunkDates = trading_days[startIdx:endIdx]; - outputTableName = "output_stream_" + string(i); - streamTableObj = objByName(outputTableName); - - // 执行计算并写入流表 - resultChunk = matchStocks(spansTable, chunkDates); - streamTableObj.append!(resultChunk); - }; - - // ============================================================ - // 步骤 6: 收集结果 - // ============================================================ - // 从各个输出流表中收集数据 - finalResult = dict(DATE,ANY,true); - - for(tableName in keys(outputTables)){ - streamTableObj = objByName(tableName); - - // 读取流表数据 - data = select * from streamTableObj; - - if(size(data) > 0){ - // 按 trade_date 分组 - groupedData = select trade_date, code from data order by trade_date; - - currentDate = NULL; - currentCodes = []; - - for(row in groupedData){ - tday = row[`trade_date]; - code = row[`code]; - - if(isNull(currentDate) || currentDate != tday){ - // 保存前一个日期的结果 - if(!isNull(currentDate)){ - finalResult[currentDate] = currentCodes; - }; - currentDate = tday; - currentCodes = []; - }; - - currentCodes.append!(code); - }; - - // 保存最后一个日期的结果 - if(!isNull(currentDate)){ - finalResult[currentDate] = currentCodes; - }; - }; - - // 清理:删除流表 - dropTable(tableName); - dropStreamTable(tableName); - }; - - return finalResult; -}; - -/* - ============================================================================= - 进一步优化版本:使用 SQL + Cross Join(推荐) - ============================================================================= - */ - -/** - * @brief 使用 SQL Cross Join 的优化版本(最推荐) - * @description 使用 DolphinDB 的 SQL 交叉连接实现,性能最优 - * - * 优化思路: - * 1. 将 spans 转换为表 - * 2. 创建交易日历表 - * 3. 使用 CROSS JOIN + WHERE 过滤,一次性找出所有 (日期, 股票) 组合 - * 4. 使用 PIVOT 或 GROUP BY 转换为字典格式 - * - * @param startDate - 开始日期 - * @param endDate - 结束日期 - * @param stockSpans - 股票代码到日期区间的映射字典 - * @return DICT> - 日期到可交易股票列表的映射 - * - * 性能对比(测试数据: 2500 交易日 x 5000 股票 x 3 spans): - * - 原版 createDateStockMapping: ~37.5M 次迭代,内存峰值 ~2.5GB,耗时 ~180 秒 - * - 本函数: SQL 向量化,内存峰值 ~500MB,耗时 ~20 秒 - * - * 使用示例: - * stockSpans = createStockDateRangeMapping(instrumentsTB); - * dateStockMapping = createDateStockMappingSQLOptimized(2020.01.01, 2023.12.31, stockSpans); - */ -def createDateStockMappingSQLOptimized(startDate, endDate, stockSpans){ - // ============================================================ - // 步骤 1: 将字典转换为表格式 - // ============================================================ - spansTable = stockSpansToTable(stockSpans); - - // ============================================================ - // 步骤 2: 创建交易日历表 - // ============================================================ - trading_days = getMarketCalendar("XSHG", startDate, endDate); - calendarTable = table(trading_days as `trade_date); - - // ============================================================ - // 步骤 3: 使用 Cross Join 找出所有有效的 (日期, 股票) 组合 - // ============================================================ - // CROSS JOIN 生成所有可能的组合 - // WHERE 过滤出在交易区间内的组合 - validPairs = select - calendarTable.trade_date, - spansTable.code - from calendarTable, spansTable - where spansTable.begin_dt <= calendarTable.trade_date <= spansTable.end_dt - order by trade_date, code; - - // ============================================================ - // 步骤 4: 转换为字典格式 - // ============================================================ - // 按 trade_date 分组 - grouped = select trade_date, code from validPairs order by trade_date; - - finalResult = dict(DATE,ANY,true); - currentDate = NULL; - currentCodes = []; - - for(row in grouped){ - tday = row[`trade_date]; - code = row[`code`; - - if(isNull(currentDate) || currentDate != tday){ - // 保存前一个日期的结果 - if(!isNull(currentDate)){ - finalResult[currentDate] = currentCodes; - }; - currentDate = tday; - currentCodes = []; - }; - - currentCodes.append!(code); - }; - - // 保存最后一个日期的结果 - if(!isNull(currentDate)){ - finalResult[currentDate] = currentCodes; - }; - - return finalResult; -}; diff --git a/qlib/data/backend/ddb_qlib/utils.py b/qlib/data/backend/ddb_qlib/utils.py index 812138f2bfd..622cc76151b 100644 --- a/qlib/data/backend/ddb_qlib/utils.py +++ b/qlib/data/backend/ddb_qlib/utils.py @@ -5,6 +5,8 @@ LastEditTime: 2025-10-21 22:18:51 Description: ''' +import re + import pandas as pd from typing import Union,Tuple,List,Dict @@ -89,4 +91,55 @@ def extract_function_name(func_call: str) -> str: # 分离函数名 func_name = func_call[: func_call.find("(")] - return func_name \ No newline at end of file + return func_name + +def validate_date_str(date_str: str) -> str: + """校验 Wind 风格日期字符串(YYYYMMDD),用于拼入 SQL 前的防注入检查。 + + :param date_str: 形如 ``"20240101"`` 的日期字符串 + :return: 原样返回通过校验的字符串 + :raises ValueError: 格式不为 8 位数字时抛出 + """ + if not re.fullmatch(r"\d{8}", str(date_str)): + raise ValueError(f"日期参数必须为 8 位数字(YYYYMMDD),收到: {date_str!r}") + return str(date_str) + + +def validate_sql_identifier(value: str) -> str: + """校验将拼入 SQL 的标识符/代码类参数(Wind 代码、交易所、库表名等)。 + + 仅允许字母、数字、点、下划线、连字符——足以覆盖合法的 Wind 代码 + (如 ``000300.SH``)与库表名,同时拒绝引号、分号等注入载体。 + + :param value: 待校验的字符串 + :return: 原样返回通过校验的字符串 + :raises ValueError: 含白名单以外字符时抛出 + """ + if not re.fullmatch(r"[A-Za-z0-9._-]+", str(value)): + raise ValueError(f"参数含非法字符(仅允许字母/数字/./_/-): {value!r}") + return str(value) + + +# 表列名的进程内缓存:表结构仅随 DDL 变更,写路径经 invalidate_ddb_caches 失效 +_TABLE_COLUMNS_CACHE: Dict[Tuple[str, str], List[str]] = {} + + +def get_table_columns(session, db_path: str, table_name: str) -> List[str]: + """带进程内缓存的表列名查询。 + + :param session: DolphinDB 会话 + :param db_path: 数据库路径(含 dfs:// 前缀) + :param table_name: 表名 + :return: 列名列表(副本,调用方可安全修改) + """ + key = (db_path, table_name) + cols = _TABLE_COLUMNS_CACHE.get(key) + if cols is None: + cols = session.loadTable(table_name, db_path).schema["name"].tolist() + _TABLE_COLUMNS_CACHE[key] = cols + return list(cols) + + +def clear_table_columns_cache() -> None: + """清空表列名缓存(DDL/写路径变更后调用)。""" + _TABLE_COLUMNS_CACHE.clear() diff --git a/qlib/data/data.py b/qlib/data/data.py index 2e4dd99a054..07f82adfbc3 100644 --- a/qlib/data/data.py +++ b/qlib/data/data.py @@ -9,6 +9,7 @@ import copy import queue import re +import threading from typing import Dict, List, Optional, Tuple, Union import numpy as np @@ -741,7 +742,9 @@ def split_list_by_length(lst, n): return [lst[i:i+n] for i in range(0, len(lst), n)] column_names:List = [column_names] if isinstance(column_names,str) else column_names - split_column_names:List[List[str]] = split_list_by_length(column_names, 30) + # 每批字段数可配置(C["ddb_field_chunk_size"],默认 30) + chunk_size: int = int(C.get("ddb_field_chunk_size", 30)) + split_column_names:List[List[str]] = split_list_by_length(column_names, chunk_size) dfs:List[pd.DataFrame] = [ExpressionD.expression(inst, fields, start_time, end_time, freq) for fields in split_column_names] dfs_filter_empty:List[pd.DataFrame] = [df for df in dfs if isinstance(df, (pd.Series, pd.DataFrame)) and not df.empty] @@ -776,11 +779,11 @@ def _process_group(group_data, inst, inst_processors): maxtasksperchild=C.maxtasksperchild, )(tasks) - # 合并结果 + # 合并结果并排序保证确定性,随后落入统一的 cache_to_origin_data 归一化路径 + # ⚠️ 此处不能提前 return:否则 inst_processors 的处理结果会被丢弃(历史 bug) if results: - data = pd.concat(results, axis=0) - return pd.DataFrame() - + data = pd.concat(results, axis=0).sort_index() + if not data.empty: # column_names,*_ = normalize_fields_to_ddb(column_names) # column_names:List[str] = list(column_names.keys()) @@ -1056,8 +1059,8 @@ def ddb_expression( self, instrument, field, start_time=None, end_time=None, freq="day" ): - # 使用ddb表达但也兼容qlib原有表达式,相比qlib表达,缺失对于前序计算期的支持 - # 比如计算MA10时wind为10所以应该在原有起始日期前10天开始计算.但ddb没有考虑这种情况. + # 使用ddb表达但也兼容qlib原有表达式;前序计算期(如 MA10 需向前多取 + # 10 日)由 fetch_features_from_ddb 按算子树解析后在 DDB 端外扩并截断 series:pd.Series = pd.Series(np.float32) try: # 这样便不会从dolphindb_storage中获取数据除非使用的feature调用 @@ -1623,13 +1626,25 @@ def __init__(self, uri="dolphindb://admin:123456@127.0.0.1:8848") -> None: # 初始化连接 config = DDBConnectionSpec(uri=self.uri) # 创建连接管理器 - connector = DDBClient(config) + self._connector = DDBClient(config) # 创建表操作对象 - self.session = connector.session - # 创建连接池 - self.pool = connector.pool + self.session = self._connector.session + # ⚠️ ddb.Session 非线程安全,且 feature 查询是「run(上传日期)→upload(变量)→run(主查询)」 + # 的多步会话对话,交错执行会互相覆盖服务器端变量(跨线程数据污染的根因)。 + # 约定:所有 DBClient.session 的触点必须包在 `with DBClient.session_lock:` 内。 + # 使用 RLock:storage 读取可能嵌套在 feature 查询持锁期间发生。 + self.session_lock = threading.RLock() register_ddb_functions_to_qlib(self.session) + @property + def pool(self): + """惰性获取连接池。 + + ⚠️ 不在 __init__ 中急切创建:连接池会向服务器开 4 条连接, + 每条都占服务器内存(社区版 2 核/8GB),且读路径并不使用连接池。 + """ + return self._connector.pool + import sys diff --git a/qlib/data/dataset/loader.py b/qlib/data/dataset/loader.py index 4e46b1d474f..7dfdcd25467 100644 --- a/qlib/data/dataset/loader.py +++ b/qlib/data/dataset/loader.py @@ -215,10 +215,16 @@ def __init__( ), f"freq(={self.freq}), inst_processors(={self.inst_processors}) cannot be None/empty" def load(self, instruments=None, start_time=None, end_time=None) -> pd.DataFrame: - # 串行化整个 load 体(含 is_group 两分支与 load_group_df 调用),阻断并发 load - # 共享 qlib 数据层时的跨线程数据污染。详见 _load_lock 注释。 - with QlibDataLoader._load_lock: - return super().load(instruments, start_time, end_time) + # 跨线程数据污染仅在 DolphinDB 后端出现(单一共享 session + 全局 H 缓存); + # 文件后端本就线程安全,加锁反而是吞吐回归,故锁收窄到 DDB 路径。 + # DDB 路径在会话级 session_lock(B1)之外保留本 load 级锁,作为全局 H + # MemCache 写模式的双保险。详见 _load_lock 注释。 + from ..data import is_using_dolphindb + + if is_using_dolphindb(): + with QlibDataLoader._load_lock: + return super().load(instruments, start_time, end_time) + return super().load(instruments, start_time, end_time) def load_group_df( self, @@ -722,7 +728,12 @@ def __init__( raise ValueError("db_name cannot be empty") # Initialize connections and state + # ⚠️ 这里持有的是进程级共享的 DBClient.session(storage/feature 路径同用), + # 本实例并不拥有它,因此 _owns_session=False,teardown 时不得关闭; + # 会话非线程安全,所有 self.session 调用必须持有 _session_lock self.session = DBClient.session + self._session_lock = DBClient.session_lock + self._owns_session = False self._table_validated = False self._mysql_bridge = None # DolphinDB mode only @@ -736,7 +747,9 @@ def _validate_table_exists(self) -> None: try: table_path = f"dfs://{self.db_name}" - if not self.session.existsTable(table_path, self.table_name): + with self._session_lock: + table_exists = self.session.existsTable(table_path, self.table_name) + if not table_exists: raise ValueError( f"Table '{self.table_name}' does not exist in database '{self.db_name}'. " f"Please check table name and database configuration." @@ -745,7 +758,7 @@ def _validate_table_exists(self) -> None: except Exception as e: if "does not exist" in str(e): raise ValueError(str(e)) - raise RuntimeError(f"Failed to validate table existence: {e}") + raise RuntimeError(f"Failed to validate table existence: {e}") from e def load(self, instruments=None, start_time=None, end_time=None) -> pd.DataFrame: """ @@ -817,44 +830,46 @@ def _load_via_dolphindb( pivot_values=None, ) -> pd.DataFrame: """Load data directly from DolphinDB.""" - # Load table and build query - tb = self.session.loadTable( - dbPath=f"dfs://{self.db_name}", tableName=self.table_name - ) + # ⚠️ loadTable 句柄创建与 toDF 执行是同一会话上的多步对话,整体持锁 + with self._session_lock: + # Load table and build query + tb = self.session.loadTable( + dbPath=f"dfs://{self.db_name}", tableName=self.table_name + ) - # Build select clause - select_fields = self._build_select_list( - fields, datetime_col, instruments_col, pivot, pivot_columns - ) + # Build select clause + select_fields = self._build_select_list( + fields, datetime_col, instruments_col, pivot, pivot_columns + ) - # 在透视模式下,确保包含值列 - if pivot and pivot_values and pivot_values not in select_fields: - select_fields.append(pivot_values) + # 在透视模式下,确保包含值列 + if pivot and pivot_values and pivot_values not in select_fields: + select_fields.append(pivot_values) - query = tb.select(select_fields) + query = tb.select(select_fields) - # Apply time filtering - if start_time and end_time: - time_filter = self._build_time_filter( - start_time, end_time, datetime_col, datetime_format, mysql_format=False - ) - query = query.where(time_filter) + # Apply time filtering + if start_time and end_time: + time_filter = self._build_time_filter( + start_time, end_time, datetime_col, datetime_format, mysql_format=False + ) + query = query.where(time_filter) - # Apply instrument filtering - if instruments: - instruments_filter = self._build_instruments_filter( - instruments, instruments_col, mysql_format=False - ) - query = query.where(instruments_filter) + # Apply instrument filtering + if instruments: + instruments_filter = self._build_instruments_filter( + instruments, instruments_col, mysql_format=False + ) + query = query.where(instruments_filter) - # 在透视模式下,添加字段筛选条件 - if pivot and fields and pivot_columns: - fields_filter = self._build_pivot_fields_filter( - fields, pivot_columns, mysql_format=False - ) - query = query.where(fields_filter) + # 在透视模式下,添加字段筛选条件 + if pivot and fields and pivot_columns: + fields_filter = self._build_pivot_fields_filter( + fields, pivot_columns, mysql_format=False + ) + query = query.where(fields_filter) - return query.toDF() + return query.toDF() def query_raw_data( self, @@ -948,7 +963,7 @@ def query_raw_data( return self._load_from_mysql_bridge(select, where, groupBy, having) except Exception as e: - raise RuntimeError(f"Failed to load data: {e}") + raise RuntimeError(f"Failed to load data: {e}") from e def _load_from_dolphindb( self, @@ -992,7 +1007,8 @@ def _load_from_dolphindb( ) # Execute query - result = self.session.run(expr) + with self._session_lock: + result = self.session.run(expr) # Ensure DataFrame return if not isinstance(result, pd.DataFrame): @@ -1023,7 +1039,9 @@ def get_table_info(self) -> Dict: "table_exists": False, } - if self.session.existsTable(f"dfs://{self.db_name}", self.table_name): + with self._session_lock: + table_exists = self.session.existsTable(f"dfs://{self.db_name}", self.table_name) + if table_exists: info["table_exists"] = True # Could add more table info like row count, column info, etc. @@ -1038,7 +1056,9 @@ def __enter__(self): def __exit__(self, exc_type, exc_val, exc_tb): """Context manager exit with proper resource cleanup.""" try: - if self.session: + # ⚠️ 历史 bug:曾无条件 close 进程级共享 session,破坏其他使用方; + # 仅在本实例独占会话(_owns_session=True)时才允许关闭 + if self._owns_session and self.session: self.session.close() except Exception as e: warnings.warn(f"Error closing DolphinDB session: {e}") @@ -1047,7 +1067,7 @@ def __exit__(self, exc_type, exc_val, exc_tb): def __del__(self): """Destructor for cleanup.""" try: - if hasattr(self, "session") and self.session: + if getattr(self, "_owns_session", False) and getattr(self, "session", None): self.session.close() except Exception: pass # Ignore cleanup errors in destructor diff --git a/qlib/data/storage/dolphindb_storage.py b/qlib/data/storage/dolphindb_storage.py index d73f11b68fc..a1683ef588e 100644 --- a/qlib/data/storage/dolphindb_storage.py +++ b/qlib/data/storage/dolphindb_storage.py @@ -55,9 +55,21 @@ def support_freq(self) -> List[str]: def uri(self) -> Tuple[str, str]: return self.db_path, self.table_name - def exists(self, db_path: str, table_name: str) -> bool: + # existsTable 的正向结果缓存:每个存储访问器都会 check(),未缓存时每次 + # 访问都多付一次 RPC。⚠️ 仅缓存 True——表不存在必须持续报错直到其被创建; + # 写路径经 ddb_qlib.invalidate_ddb_caches() 失效。 + _exists_cache: Dict[Tuple[str, str], bool] = {} - return DBClient.session.existsTable(db_path, table_name) + def exists(self, db_path: str, table_name: str) -> bool: + key = (db_path, table_name) + if self._exists_cache.get(key): + return True + # ⚠️ 约定:所有 DBClient.session 触点必须持有 session_lock(会话非线程安全) + with DBClient.session_lock: + result = bool(DBClient.session.existsTable(db_path, table_name)) + if result: + self._exists_cache[key] = True + return result def check(self): """check self.uri @@ -93,9 +105,10 @@ def _read_calendar(self) -> List[CalVT]: if not self.exists(self.db_path, self.table_name): self._write_calendar(values=[]) - df: pd.DataFrame = ( - DBClient.session.loadTable(self.table_name, self.db_path).select("*").toDF() - ) + with DBClient.session_lock: + df: pd.DataFrame = ( + DBClient.session.loadTable(self.table_name, self.db_path).select("*").toDF() + ) if df.empty: return [] @@ -120,17 +133,23 @@ def _freq_db(self) -> str: self._freq_file_cache = freq return self._freq_file_cache - @property - def data(self) -> List[CalVT]: - self.check() - # If cache is enabled, then return cache directly + def _cached_calendar(self) -> List[CalVT]: + """经 ``H["c"]`` 缓存读取原始日历(与 ``data`` 共用同一缓存键)。 + + ⚠️ 性能关键:``index()``/``__getitem__`` 曾绕过缓存直接 + ``_read_calendar()``,导致数据集对齐期间反复全量下载日历表。 + """ if self.enable_read_cache: key = "orig_file" + str(self.uri) if key not in H["c"]: H["c"][key] = self._read_calendar() - _calendar = H["c"][key] - else: - _calendar = self._read_calendar() + return H["c"][key] + return self._read_calendar() + + @property + def data(self) -> List[CalVT]: + self.check() + _calendar = self._cached_calendar() if Freq(self._freq_db) != Freq(self.freq): _calendar = resam_calendar( np.array(list(map(pd.Timestamp, _calendar))), @@ -142,12 +161,12 @@ def data(self) -> List[CalVT]: def index(self, value: CalVT) -> int: self.check() - calendar = self._read_calendar() + calendar = self._cached_calendar() return int(np.argwhere(calendar == value)[0]) def __getitem__(self, i: Union[int, slice]) -> Union[CalVT, List[CalVT]]: self.check() - return self._read_calendar()[i] + return self._cached_calendar()[i] def __len__(self) -> int: return len(self.data) @@ -169,20 +188,29 @@ def __init__(self, market: str, freq: str, provider_uri: dict = None, **kwargs): self.table_name: str = market.lower() def _read_instrument(self) -> Dict[InstKT, InstVT]: - + self.check() + # 经 H["i"] 缓存读取(此前每个访问器调用都全量下载 + 行循环重建 dict); + # 写路径经 ddb_qlib.invalidate_ddb_caches() 失效 + cache_key = "db_instrument_" + str(self.uri) + if cache_key in H["i"]: + # 浅拷贝防止调用方修改缓存本体(值 spans 列表按约定只读) + return dict(H["i"][cache_key]) + _instruments = dict() # sql = f"""SELECT * FROM {self.ddb_table}""" # df = DolphinDB.run(sql) - df: pd.DataFrame = ( - DBClient.session.loadTable(self.table_name, self.db_path).select("*").toDF() - ) + with DBClient.session_lock: + df: pd.DataFrame = ( + DBClient.session.loadTable(self.table_name, self.db_path).select("*").toDF() + ) for row in df.itertuples(index=False): _instruments.setdefault(row[0], []).append((row[1], row[2])) - return _instruments + H["i"][cache_key] = _instruments + return dict(_instruments) @property @@ -257,14 +285,17 @@ def __getitem__(self, i: Union[int, slice]) -> Union[Tuple[int, float], pd.Serie else: raise TypeError(f"不支持的索引类型: type(i) = {type(i)}") - df: pd.DataFrame = fetch_features_from_ddb( - DBClient.session, - self.instrument, - self.field, - start_time, - end_time, - self.freq - ) + # ⚠️ 整个 fetch 是「run→upload→run」的多步会话对话,必须整体持锁, + # 否则并发线程会互相覆盖服务器端的 instruments/expressions 变量 + with DBClient.session_lock: + df: pd.DataFrame = fetch_features_from_ddb( + DBClient.session, + self.instrument, + self.field, + start_time, + end_time, + self.freq + ) if df.empty: return pd.Series(dtype=float) diff --git a/scripts/benchmark_ddb_backend.py b/scripts/benchmark_ddb_backend.py new file mode 100644 index 00000000000..c8cc3117342 --- /dev/null +++ b/scripts/benchmark_ddb_backend.py @@ -0,0 +1,96 @@ +"""DDB 后端 live 基准脚本(可选,需要真实 DolphinDB 服务器)。 + +用法:: + + DDB_BENCH_URI="dolphindb://admin:123456@host:8848" python scripts/benchmark_ddb_backend.py + +对 D.calendar / D.instruments / D.features(纯字段与计算表达式)计时, +并通过包装 session.run/upload 统计 RPC 往返数,用于优化前后对比。 +不设 DDB_BENCH_URI 时直接退出(离线 CI 安全)。 +""" + +import functools +import os +import sys +import time +from pathlib import Path + +# 保证从仓库根目录可直接运行 +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +from loguru import logger + + +def main() -> None: + uri = os.environ.get("DDB_BENCH_URI") + if not uri: + logger.info("未设置 DDB_BENCH_URI,跳过 live 基准") + return + + import qlib + from qlib.config import REG_CN + from qlib.data import D + + qlib.init(database_uri=uri, region=REG_CN) + + # 包装共享 session 统计 RPC + from qlib.data.data import DBClient + + session = DBClient.session + counts = {"run": 0, "upload": 0, "loadTable": 0} + + def _count(name, func): + @functools.wraps(func) + def wrapper(*args, **kwargs): + counts[name] += 1 + return func(*args, **kwargs) + + return wrapper + + session.run = _count("run", session.run) + session.upload = _count("upload", session.upload) + session.loadTable = _count("loadTable", session.loadTable) + + def bench(label, fn): + for key in counts: + counts[key] = 0 + t0 = time.perf_counter() + result = fn() + elapsed = time.perf_counter() - t0 + size = len(result) if hasattr(result, "__len__") else "-" + logger.info( + f"{label}: {elapsed:.3f}s, rows={size}, " + f"RPC(run={counts['run']}, upload={counts['upload']}, loadTable={counts['loadTable']})" + ) + return result + + start, end = "2023-01-01", "2023-12-31" + + bench("D.calendar", lambda: D.calendar(start_time=start, end_time=end)) + instruments = bench( + "D.instruments+list", + lambda: D.list_instruments(D.instruments("csi300"), start_time=start, end_time=end, as_list=True), + ) + sample = instruments[:50] + + bench( + "D.features 纯字段×2", + lambda: D.features(sample, ["$close", "$open"], start_time=start, end_time=end), + ) + bench( + "D.features 计算×2", + lambda: D.features( + sample, ["Ref($close,1)/$close-1", "Mean($volume,5)"], start_time=start, end_time=end + ), + ) + # 二次调用:观测缓存生效后的 RPC 下降 + bench( + "D.features 计算×2 (二次)", + lambda: D.features( + sample, ["Ref($close,1)/$close-1", "Mean($volume,5)"], start_time=start, end_time=end + ), + ) + + +if __name__ == "__main__": + main() diff --git a/tests/ddb_mocks.py b/tests/ddb_mocks.py new file mode 100644 index 00000000000..a080ed68315 --- /dev/null +++ b/tests/ddb_mocks.py @@ -0,0 +1,146 @@ +"""DolphinDB 会话的可复用记录型 mock(离线测试基础设施)。 + +模拟 ``fetch_features_from_ddb`` / storage 所需的最小 SDK 表面: +``run`` / ``upload`` / ``runFile`` / ``existsTable`` / ``loadTable``(含 +``select/where/sort/exec/toDF/schema`` 链式调用),并对每类 RPC 计数, +供 RPC 往返数基线测试与分支回归测试共用。 +""" + +import re +from typing import Callable + +import numpy as np +import pandas as pd + + +class FakeQueryChain: + """模拟 loadTable 返回的表句柄及其链式查询。""" + + def __init__(self, session: "RecordingSession", table_name: str, db_path: str): + self._session = session + self.table_name = table_name + self.db_path = db_path + self.selects: list = [] + self.wheres: list[str] = [] + self.sorts: list = [] + self.exec_col: str | None = None + + # --- 链式查询 --- + def select(self, cols): + self.selects.append(cols) + return self + + def where(self, cond: str): + self.wheres.append(cond) + return self + + def sort(self, cols): + self.sorts.append(cols) + return self + + def exec(self, col: str): + self.exec_col = col + return self + + def toDF(self): + self._session.query_chains.append(self) + return self._session._resolve_table_result(self) + + @property + def schema(self) -> pd.DataFrame: + cols = self._session.table_columns.get( + (self.db_path, self.table_name), + self._session.table_columns.get(self.table_name, []), + ) + return pd.DataFrame({"name": cols}) + + +class RecordingSession: + """记录所有 RPC 调用的假 DolphinDB 会话。 + + :param calendar: 交易日历(np.ndarray[datetime64]),供 TradeDateUtils 使用 + :param table_columns: {(db_path, table_name) 或 table_name: [列名]},供 schema 查询 + :param table_results: {(db_path, table_name) 或 table_name: DataFrame 或 callable(chain)}, + 供 toDF 返回查询结果 + :param run_responses: [(正则, 结果或 callable(script))],按序匹配 run 脚本 + """ + + def __init__( + self, + calendar: np.ndarray | None = None, + table_columns: dict | None = None, + table_results: dict | None = None, + run_responses: list[tuple[str, object]] | None = None, + ): + self.calendar = calendar + self.table_columns = table_columns or {} + self.table_results = table_results or {} + self.run_responses = run_responses or [] + + # RPC 计数与调用记录 + self.counts: dict[str, int] = { + "run": 0, + "upload": 0, + "runFile": 0, + "existsTable": 0, + "loadTable": 0, + } + self.run_scripts: list[str] = [] + self.uploads: list[dict] = [] + self.run_files: list = [] + self.load_table_calls: list[tuple[str, str]] = [] + self.query_chains: list[FakeQueryChain] = [] + self.exists_result: bool = True # existsTable 的返回值(可按测试配置) + + # --- SDK 表面 --- + def run(self, script: str): + self.counts["run"] += 1 + self.run_scripts.append(script) + for pattern, result in self.run_responses: + if re.search(pattern, script): + return result(script) if isinstance(result, Callable) else result + return None + + def upload(self, variables: dict): + self.counts["upload"] += 1 + self.uploads.append(variables) + + def runFile(self, filepath): + self.counts["runFile"] += 1 + self.run_files.append(filepath) + + def existsTable(self, db_path: str, table_name: str) -> bool: + self.counts["existsTable"] += 1 + return self.exists_result + + def loadTable(self, tableName: str = None, dbPath: str = None, **kwargs): + # 兼容位置参数写法 loadTable(table, db) + table = tableName or kwargs.get("tableName") + db = dbPath or kwargs.get("dbPath") + self.counts["loadTable"] += 1 + self.load_table_calls.append((db, table)) + return FakeQueryChain(self, table, db) + + # --- 结果解析 --- + def _resolve_table_result(self, chain: FakeQueryChain): + # 日历查询:exec("TRADE_DAYS") + if chain.exec_col == "TRADE_DAYS": + if self.calendar is None: + raise AssertionError("测试未配置 calendar,但发生了日历查询") + return self.calendar + result = self.table_results.get( + (chain.db_path, chain.table_name), + self.table_results.get(chain.table_name), + ) + if callable(result): + return result(chain) + if result is None: + raise AssertionError( + f"测试未配置表结果: {chain.db_path}/{chain.table_name}" + ) + return result + + +def make_calendar(start: str, periods: int) -> np.ndarray: + """构造连续工作日的假交易日历(np.datetime64 数组,升序)。""" + return pd.date_range(start, periods=periods, freq="B").values diff --git a/tests/test_backend_sql_guards.py b/tests/test_backend_sql_guards.py new file mode 100644 index 00000000000..5d2d05b1d88 --- /dev/null +++ b/tests/test_backend_sql_guards.py @@ -0,0 +1,62 @@ +"""MySQL 同步 SQL 参数校验的离线单元测试。 + +背景:``QlibDDBMySQLInitializer`` 的同步方法把用户可控的日期/代码/交易所参数 +直接 f-string 拼入 SQL(唯一真实的注入面)。修复方式为拼接前做白名单校验: +- 日期:``^\\d{8}$``(Wind YYYYMMDD 约定) +- 标识符/代码:``^[A-Za-z0-9._-]+$``(覆盖合法 Wind 代码如 000300.SH) +合法输入原样通过,行为不变;仅病态输入提前抛 ValueError。 +""" + +import pytest + +from qlib.data.backend.ddb_qlib.utils import validate_date_str, validate_sql_identifier + + +class TestValidateDateStr: + def test_valid_dates_pass_through(self): + assert validate_date_str("20240101") == "20240101" + assert validate_date_str("19991231") == "19991231" + + @pytest.mark.parametrize( + "bad", + ["2024-01-01", "2024.01.01", "202401", "20240101; DROP TABLE x", "", "20240101'"], + ) + def test_invalid_dates_rejected(self, bad): + with pytest.raises(ValueError): + validate_date_str(bad) + + +class TestValidateSqlIdentifier: + def test_legitimate_values_pass_through(self): + # 合法 Wind 代码、交易所、库表名必须原样通过(行为不变) + for value in ["000300.SH", "SSE", "SZSE", "ASHAREEODPRICES", "csi_300", "H30021.CSI"]: + assert validate_sql_identifier(value) == value + + @pytest.mark.parametrize( + "bad", + ["'; DROP TABLE x--", "SSE'", "a b", 'x"y', "code;", "(select 1)", ""], + ) + def test_injection_payloads_rejected(self, bad): + with pytest.raises(ValueError): + validate_sql_identifier(bad) + + +class TestInitializerGuards: + """同步方法在构造 SQL 前即拒绝病态参数(不触达任何数据库连接)。""" + + def _make_initializer(self): + from qlib.data.backend.ddb_qlib.ddb_mysql_bridge import QlibDDBMySQLInitializer + + init = object.__new__(QlibDDBMySQLInitializer) + init._bridge = None # 校验应在 _get_bridge 之前发生 + return init + + def test_sync_index_daily_rejects_bad_code(self): + init = self._make_initializer() + with pytest.raises(ValueError): + init.sync_index_daily(["000300.SH'; DROP TABLE AINDEXEODPRICES--"]) + + def test_sync_index_daily_rejects_bad_date(self): + init = self._make_initializer() + with pytest.raises(ValueError): + init.sync_index_daily(["000300.SH"], start_date="2024-01-01") diff --git a/tests/test_ddb_calendar_cache.py b/tests/test_ddb_calendar_cache.py new file mode 100644 index 00000000000..a94c360f025 --- /dev/null +++ b/tests/test_ddb_calendar_cache.py @@ -0,0 +1,104 @@ +"""D1 日历缓存的离线回归测试。 + +历史问题: +1. ``TradeDateUtils.__init__`` 每次构造都全量下载交易日历——而 + ``fetch_features_from_ddb`` 每批字段(30 个/批)都会构造一次, + Alpha158 一次 ``D.features`` ≈ 6 次全量日历下载; +2. ``DBCalendarStorage.index()``/``__getitem__`` 绕过 ``H["c"]`` 缓存 + 直接 ``_read_calendar()``,数据集对齐期间反复全量下载。 + +修复:TradeDateUtils 模块级缓存 + storage 统一走 ``_cached_calendar()``; +写路径通过 ``invalidate_ddb_caches()`` 失效。 +""" + +import threading + +import numpy as np +import pandas as pd +import pytest + +from ddb_mocks import RecordingSession, make_calendar + +from qlib.data.backend.ddb_qlib import invalidate_ddb_caches +from qlib.data.backend.ddb_qlib.ddb_features import TradeDateUtils +from qlib.data.data import DBClient + +CAL = make_calendar("2024-01-01", 10) + + +@pytest.fixture(autouse=True) +def _clean_caches(): + invalidate_ddb_caches() + yield + invalidate_ddb_caches() + + +class TestTradeDateUtilsCache: + def test_calendar_loaded_once_across_instances(self): + session = RecordingSession(calendar=CAL) + utils_a = TradeDateUtils(session, "day") + utils_b = TradeDateUtils(session, "day") + assert session.counts["loadTable"] == 1, "第二次构造应命中缓存" + # 缓存共享同一份数据且行为一致 + assert np.array_equal(utils_a._calendar, utils_b._calendar) + start, end = utils_b.get_locate_date("2024-01-06", "2024-01-08") + assert pd.Timestamp(start) == pd.Timestamp("2024-01-08") # 周六周日 snap 到下一交易日 + + def test_clear_cache_forces_reload(self): + session = RecordingSession(calendar=CAL) + TradeDateUtils(session, "day") + TradeDateUtils.clear_cache() + TradeDateUtils(session, "day") + assert session.counts["loadTable"] == 2 + + def test_invalidate_ddb_caches_clears_calendar(self): + session = RecordingSession(calendar=CAL) + TradeDateUtils(session, "day") + invalidate_ddb_caches() + TradeDateUtils(session, "day") + assert session.counts["loadTable"] == 2 + + +class _FakeProvider: + def __init__(self, session): + self.session = session + self.session_lock = threading.RLock() + + +@pytest.fixture() +def calendar_storage(monkeypatch): + """构造带假 provider 的 DBCalendarStorage。""" + from qlib.config import C + from qlib.data.storage.dolphindb_storage import DBCalendarStorage + + if "region" not in C: + C["region"] = "cn" # DBCalendarStorage.__init__ 依赖,qlib.init 时才会设置 + + session = RecordingSession( + table_results={ + ("dfs://QlibCalendars", "day"): pd.DataFrame({"TRADE_DAYS": pd.DatetimeIndex(CAL)}) + } + ) + old = DBClient.__dict__.get("_provider") + DBClient.register(_FakeProvider(session)) + storage = DBCalendarStorage(freq="day", future=False) + yield storage, session + DBClient.register(old) + + +class TestDBCalendarStorageCache: + def test_index_and_getitem_share_cache(self, calendar_storage): + storage, session = calendar_storage + _ = storage.data + loads_after_data = session.counts["loadTable"] + _ = storage[0] + _ = storage[1:3] + loads_after_getitem = session.counts["loadTable"] + assert loads_after_getitem == loads_after_data, ( + "index/__getitem__ 绕过缓存重复下载日历(历史回归)" + ) + + def test_data_values_unchanged_by_cache(self, calendar_storage): + storage, session = calendar_storage + assert list(storage.data) == list(pd.DatetimeIndex(CAL)) + assert storage[0] == pd.DatetimeIndex(CAL)[0] diff --git a/tests/test_ddb_client.py b/tests/test_ddb_client.py new file mode 100644 index 00000000000..4318feec321 --- /dev/null +++ b/tests/test_ddb_client.py @@ -0,0 +1,172 @@ +"""DDBClient 连接池生命周期的离线单元测试(不依赖 DolphinDB 服务器)。 + +历史 bug 回归: +1. ``_pool_instance`` 曾是类变量——连接不同服务器的多个客户端会共享同一个池; +2. ``close_pool`` 曾引用不存在的 ``cls._pool_lock``(AttributeError 被裸 except 吞掉), + 且读取的类变量永远是 None,实际从未关闭过任何池; +3. ``tableAppender``/``tableUpsert`` 调用不存在的 ``self.get_session()``,属死代码,已删除。 +""" + +import pytest + +from qlib.data.backend.ddb_qlib import ddb_client as ddb_client_mod +from qlib.data.backend.ddb_qlib.ddb_client import DDBClient, DDBConnectionSpec + + +class _FakeSession: + """记录 connect/close 调用的假会话。""" + + def __init__(self, *args, **kwargs): + self.connect_calls = [] + self.closed = False + + def connect(self, **kwargs): + self.connect_calls.append(kwargs) + + def close(self): + self.closed = True + + +class _FakePool: + """记录构造参数与 shutDown 调用的假连接池。""" + + instances: list["_FakePool"] = [] + + def __init__(self, *args, **kwargs): + self.args = args + self.kwargs = kwargs + self.shutdown_called = 0 + _FakePool.instances.append(self) + + def shutDown(self): + self.shutdown_called += 1 + + +@pytest.fixture() +def fake_ddb(monkeypatch): + """把 ddb_client 模块内的 dolphindb SDK 替换为假实现。""" + _FakePool.instances = [] + monkeypatch.setattr(ddb_client_mod.ddb, "session", _FakeSession) + monkeypatch.setattr(ddb_client_mod.ddb, "DBConnectionPool", _FakePool) + + class _FakeSettings: + PROTOCOL_DDB = "ddb" + + monkeypatch.setattr(ddb_client_mod.ddb, "settings", _FakeSettings, raising=False) + return _FakePool + + +def _make_client(uri: str = "dolphindb://admin:pwd@127.0.0.1:8848") -> DDBClient: + return DDBClient(DDBConnectionSpec(uri=uri)) + + +def test_pool_created_lazily_and_once(fake_ddb): + client = _make_client() + assert fake_ddb.instances == [], "连接池不应在构造时创建" + pool_first = client.pool + pool_second = client.pool + assert pool_first is pool_second + assert len(fake_ddb.instances) == 1, "连接池应只创建一次" + + +def test_pool_not_shared_between_clients(fake_ddb): + """回归:池曾是类变量,导致不同客户端共享同一个池。""" + client_a = _make_client("dolphindb://admin:pwd@10.0.0.1:8848") + client_b = _make_client("dolphindb://admin:pwd@10.0.0.2:8848") + assert client_a.pool is not client_b.pool + assert len(fake_ddb.instances) == 2 + + +def test_close_pool_shuts_down_and_resets(fake_ddb): + client = _make_client() + pool = client.pool + client.close_pool() + assert pool.shutdown_called == 1 + # 关闭后再次访问会新建 + assert client.pool is not pool + # 未创建时关闭是空操作,不会新建池 + created_before = len(fake_ddb.instances) + fresh = _make_client() + fresh.close_pool() + assert len(fake_ddb.instances) == created_before + + +def test_close_pool_swallows_shutdown_error(fake_ddb): + client = _make_client() + pool = client.pool + pool.shutDown = lambda: (_ for _ in ()).throw(RuntimeError("boom")) + client.close_pool() # 不应抛出 + assert client._pool_instance is None + + +def test_close_closes_session_and_pool(fake_ddb): + client = _make_client() + pool = client.pool + client.close() + assert pool.shutdown_called == 1 + assert client.session.closed is True + + +def test_provider_pool_is_lazy(fake_ddb, monkeypatch): + """回归:DolphinDBClientProvider 曾在 init 时急切创建 4 连接的池(无人使用)。""" + import qlib.data.backend.ddb_qlib as ddb_pkg + from qlib.data.data import DolphinDBClientProvider + + monkeypatch.setattr(ddb_pkg, "register_ddb_functions_to_qlib", lambda session: None) + provider = DolphinDBClientProvider(uri="dolphindb://admin:pwd@127.0.0.1:8848") + assert fake_ddb.instances == [], "provider 初始化不应创建连接池" + _ = provider.pool + assert len(fake_ddb.instances) == 1, "首次访问 pool 才创建连接池" + + +def test_get_shared_client_reuses_per_uri(fake_ddb): + """D6:同 URI 复用共享客户端,不同 URI 各自独立。""" + from qlib.data.backend.ddb_qlib import ddb_operator as op + + op._shared_clients.clear() + a = op.get_shared_client("dolphindb://admin:pwd@10.0.0.1:8848") + b = op.get_shared_client("dolphindb://admin:pwd@10.0.0.1:8848") + c = op.get_shared_client("dolphindb://admin:pwd@10.0.0.2:8848") + assert a is b + assert a is not c + op._shared_clients.clear() + + +def test_write_df_to_ddb_uses_shared_client(fake_ddb, monkeypatch): + """D6 回归:write_df_to_ddb 曾每次调用新建 DDBClient(新 TCP 会话)。""" + import pandas as pd + + from qlib.data.backend.ddb_qlib import ddb_operator as op + + op._shared_clients.clear() + clients_used = [] + + class _FakeOperator: + def __init__(self, client): + clients_used.append(client) + + def table_appender(self, **kwargs): + pass + + def table_upsert(self, **kwargs): + pass + + monkeypatch.setattr(op, "DDBTableOperator", _FakeOperator) + uri = "dolphindb://admin:pwd@10.0.0.3:8848" + op.write_df_to_ddb("db", "tb", pd.DataFrame({"a": [1]}), uri=uri) + op.write_df_to_ddb("db", "tb", pd.DataFrame({"a": [1]}), uri=uri) + assert len(clients_used) == 2 + assert clients_used[0] is clients_used[1], "同 URI 的两次写入应复用同一客户端" + # 显式传 client 时优先使用 + explicit = _make_client("dolphindb://admin:pwd@10.0.0.4:8848") + op.write_df_to_ddb("db", "tb", pd.DataFrame({"a": [1]}), uri=None, client=explicit) + assert clients_used[-1] is explicit + op._shared_clients.clear() + + +def test_dead_write_methods_removed(fake_ddb): + """守卫:死代码 tableAppender/tableUpsert 不应被重新引入(正确实现在 DDBTableOperator)。""" + client = _make_client() + assert not hasattr(client, "tableAppender") + assert not hasattr(client, "tableUpsert") + assert not hasattr(client, "get_session") diff --git a/tests/test_ddb_concurrency.py b/tests/test_ddb_concurrency.py new file mode 100644 index 00000000000..0ecb0a40d5b --- /dev/null +++ b/tests/test_ddb_concurrency.py @@ -0,0 +1,110 @@ +"""DDB 会话级锁的并发回归测试(离线,不依赖 DolphinDB 服务器)。 + +背景:ddb.Session 非线程安全,且 feature 查询是「run(上传日期)→upload(变量)→run(主查询)」 +的多步会话对话;并发线程交错执行会互相覆盖服务器端的 instruments/expressions 变量, +造成跨线程数据污染(A 线程拿到 B 线程的股票池)。 + +约定(B1):所有 ``DBClient.session`` 触点必须持有 ``DBClient.session_lock``(RLock)。 +本测试验证 ``DBFeatureStorage.__getitem__`` 对整个 fetch 会话对话持锁串行化。 +""" + +import threading +import time + +import pandas as pd +import pytest + +import qlib.data.backend.ddb_qlib as ddb_pkg +from qlib.data.data import DBClient +from qlib.data.storage.dolphindb_storage import DBFeatureStorage + + +class _FakeProvider: + """带真实 RLock 的假 DolphinDBClientProvider。""" + + def __init__(self): + self.session = object() # 会话本体不被本测试触碰 + self.session_lock = threading.RLock() + + +@pytest.fixture() +def fake_provider(): + old = DBClient.__dict__.get("_provider") + provider = _FakeProvider() + DBClient.register(provider) + yield provider + DBClient.register(old) + + +def test_feature_fetch_serialized_across_threads(fake_provider, monkeypatch): + """并发调用 DBFeatureStorage 时,fetch 会话对话必须整体串行(不被交错)。""" + events: list[tuple[int, str]] = [] + events_guard = threading.Lock() + + def fake_fetch(session, instruments, fields, start_time, end_time, freq): + tid = threading.get_ident() + with events_guard: + events.append((tid, "enter")) + time.sleep(0.02) # 放大交错窗口 + with events_guard: + events.append((tid, "exit")) + return pd.DataFrame({"close": [1.0]}) + + monkeypatch.setattr(ddb_pkg, "fetch_features_from_ddb", fake_fetch) + + storage = DBFeatureStorage(instrument=["SH600000"], field="$close", freq="day") + + threads = [threading.Thread(target=lambda: storage[:]) for _ in range(4)] + for t in threads: + t.start() + for t in threads: + t.join(timeout=10) + + assert len(events) == 8 + # enter/exit 必须严格成对且属于同一线程——任何交错都意味着锁失效 + for i in range(0, len(events), 2): + assert events[i][1] == "enter" and events[i + 1][1] == "exit", f"事件交错: {events}" + assert events[i][0] == events[i + 1][0], f"跨线程交错: {events}" + + +def test_session_lock_is_reentrant(fake_provider): + """session_lock 必须可重入(RLock):storage 读取可嵌套在 feature 查询持锁期间。""" + lock = DBClient.session_lock + assert lock.acquire(blocking=False) + try: + assert lock.acquire(blocking=False), "session_lock 必须是 RLock(可重入)" + lock.release() + finally: + lock.release() + + +def test_load_lock_scoped_to_ddb_backend(monkeypatch): + """B2 回归:全局 _load_lock 仅在 DolphinDB 后端持有;文件后端 load 无锁。""" + import qlib.data.data as data_mod + from qlib.data.dataset.loader import DLWParser, QlibDataLoader + + lock_states: list[bool] = [] + + def fake_super_load(self, instruments=None, start_time=None, end_time=None): + lock_states.append(QlibDataLoader._load_lock.locked()) + return pd.DataFrame() + + monkeypatch.setattr(DLWParser, "load", fake_super_load) + loader = object.__new__(QlibDataLoader) # 绕过构造器,load 不依赖实例属性 + + monkeypatch.setattr(data_mod, "is_using_dolphindb", lambda: True) + loader.load() + monkeypatch.setattr(data_mod, "is_using_dolphindb", lambda: False) + loader.load() + + assert lock_states == [True, False], "DDB 路径应持锁,文件路径应无锁" + + +def test_provider_exposes_session_lock(): + """真实 DolphinDBClientProvider 必须暴露 session_lock(接口契约)。""" + import inspect + + from qlib.data.data import DolphinDBClientProvider + + source = inspect.getsource(DolphinDBClientProvider.__init__) + assert "session_lock" in source and "RLock" in source diff --git a/tests/test_ddb_dataset_processor.py b/tests/test_ddb_dataset_processor.py new file mode 100644 index 00000000000..7cfb64ff7ec --- /dev/null +++ b/tests/test_ddb_dataset_processor.py @@ -0,0 +1,149 @@ +"""ddb_dataset_processor 的离线回归测试(不依赖 DolphinDB 服务器)。 + +历史 bug 回归:当传入 inst_processors 时,`ddb_dataset_processor` 曾在并行处理后 +提前 `return pd.DataFrame()`,把处理结果整体丢弃并跳过 cache_to_origin_data 归一化。 +本文件通过注入假的 ExpressionD provider 与顺序执行的 ParallelExt 存根, +在无 DolphinDB 连接的情况下覆盖该路径。 +""" + +import numpy as np +import pandas as pd +import pytest + +from qlib.data import data as data_mod +from qlib.data.data import DatasetProvider +from qlib.data.inst_processor import InstProcessor + + +class _AddConst(InstProcessor): + """给所有列加常数的标记处理器,用于验证处理结果没有被丢弃。""" + + def __init__(self, const: float = 100.0): + self.const = const + + def __call__(self, df: pd.DataFrame, instrument, *args, **kwargs): + return df + self.const + + +class _FakeExpressionProvider: + """按固定值构造 (instrument, datetime) MultiIndex 结果的假 provider。 + + :param instruments: 股票代码列表 + :param dates: 交易日列表 + """ + + def __init__(self, instruments: list[str], dates: pd.DatetimeIndex): + self.instruments = instruments + self.dates = dates + self.calls: list[list[str]] = [] # 记录每次调用的字段批次 + + def expression(self, inst, fields, start_time, end_time, freq): + self.calls.append(list(fields)) + index = pd.MultiIndex.from_product( + [self.instruments, self.dates], names=["instrument", "datetime"] + ) + # 每个字段填充其在全部调用中的序号,便于断言列取值 + data = {f: float(i) for i, f in enumerate(fields)} + return pd.DataFrame(data, index=index, dtype=np.float32) + + +class _SequentialParallel: + """顺序执行 joblib delayed 任务的 ParallelExt 存根。""" + + def __init__(self, *args, **kwargs): + pass + + def __call__(self, tasks): + return [func(*args, **kwargs) for func, args, kwargs in tasks] + + +@pytest.fixture() +def fake_env(monkeypatch): + """注入假 provider / 并行存根 / 内核数配置,返回 provider 供断言。""" + instruments = ["SH600000", "SZ000001"] + dates = pd.date_range("2024-01-01", periods=3, freq="D") + provider = _FakeExpressionProvider(instruments, dates) + monkeypatch.setattr(data_mod.ExpressionD, "_provider", provider) + monkeypatch.setattr(data_mod, "ParallelExt", _SequentialParallel) + monkeypatch.setattr(type(data_mod.C), "get_kernels", lambda self, freq: 1, raising=False) + return provider + + +def test_inst_processors_result_not_discarded(fake_env): + """回归:inst_processors 非空时结果必须保留(历史上被丢弃返回空 DataFrame)。""" + result = DatasetProvider.ddb_dataset_processor( + inst=["SH600000", "SZ000001"], + column_names=["$close", "$open"], + start_time=pd.Timestamp("2024-01-01"), + end_time=pd.Timestamp("2024-01-03"), + freq="day", + inst_processors=[_AddConst(100.0)], + ) + assert not result.empty, "inst_processors 处理结果被丢弃(回归到历史 bug)" + # 基础值为字段序号(0/1),处理器加 100 + assert (result["$close"] == 100.0).all() + assert (result["$open"] == 101.0).all() + # 归一化路径必须走到:列名与索引结构保持约定 + assert list(result.columns) == ["$close", "$open"] + assert result.index.names == ["instrument", "datetime"] + # 结果按索引排序(确定性保证) + assert result.index.is_monotonic_increasing + + +def test_without_inst_processors_unchanged(fake_env): + """无 inst_processors 时行为与原实现一致。""" + result = DatasetProvider.ddb_dataset_processor( + inst=["SH600000", "SZ000001"], + column_names=["$close"], + start_time=pd.Timestamp("2024-01-01"), + end_time=pd.Timestamp("2024-01-03"), + freq="day", + ) + assert not result.empty + assert list(result.columns) == ["$close"] + assert (result["$close"] == 0.0).all() + + +def test_column_chunking_over_30(fake_env): + """超过 30 个字段按 30 一批分块调用 ExpressionD.expression。""" + fields = [f"$f{i}" for i in range(65)] + result = DatasetProvider.ddb_dataset_processor( + inst=["SH600000", "SZ000001"], + column_names=fields, + start_time=pd.Timestamp("2024-01-01"), + end_time=pd.Timestamp("2024-01-03"), + freq="day", + ) + assert [len(c) for c in fake_env.calls] == [30, 30, 5] + assert list(result.columns) == fields + + +def test_chunk_size_configurable(fake_env, monkeypatch): + """D7:每批字段数经 C[\"ddb_field_chunk_size\"] 可配置(默认 30)。""" + monkeypatch.setitem(data_mod.C, "ddb_field_chunk_size", 10) + fields = [f"$f{i}" for i in range(25)] + DatasetProvider.ddb_dataset_processor( + inst=["SH600000"], + column_names=fields, + start_time=pd.Timestamp("2024-01-01"), + end_time=pd.Timestamp("2024-01-03"), + freq="day", + ) + assert [len(c) for c in fake_env.calls] == [10, 10, 5] + + +def test_inst_processors_with_chunking(fake_env): + """分块 + inst_processors 组合路径:列齐全且处理生效。""" + fields = [f"$f{i}" for i in range(31)] + result = DatasetProvider.ddb_dataset_processor( + inst=["SH600000", "SZ000001"], + column_names=fields, + start_time=pd.Timestamp("2024-01-01"), + end_time=pd.Timestamp("2024-01-03"), + freq="day", + inst_processors=[_AddConst(1000.0)], + ) + assert not result.empty + assert list(result.columns) == fields + # 第二批只有 1 个字段,序号从 0 重新计:$f30 基础值为 0 + assert (result["$f30"] == 1000.0).all() diff --git a/tests/test_ddb_features.py b/tests/test_ddb_features.py index 61c6191b760..2769115ff19 100644 --- a/tests/test_ddb_features.py +++ b/tests/test_ddb_features.py @@ -110,11 +110,65 @@ def test_register_loads_ops_dos_first(self) -> None: register_ddb_functions_to_qlib(session) # type: ignore[arg-type] # 鸭子类型假会话 assert session.files[0] == "ops.dos" - def test_register_loads_every_script_exactly_once(self) -> None: - """所有 .dos 脚本都被加载且仅加载一次。""" + def test_register_loads_core_scripts_exactly_once(self) -> None: + """默认仅加载核心三件套(alpha 库 119KB 改为按需惰性加载)且无重复。""" session = _FakeSession() register_ddb_functions_to_qlib(session) # type: ignore[arg-type] # 鸭子类型假会话 + assert session.files == [ + "ops.dos", + "featureEngineering.dos", + "prepareInstruments.dos", + ] + + def test_register_preload_loads_every_script(self) -> None: + """preload_alpha_libs=True 恢复历史全量加载行为(ops.dos 仍置首)。""" + session = _FakeSession() + register_ddb_functions_to_qlib(session, preload_alpha_libs=True) # type: ignore[arg-type] script_dir = Path(register_ddb_functions_to_qlib.__code__.co_filename).parent / "ddb_scripts" expected_count = len(list(script_dir.glob("*.dos"))) + assert session.files[0] == "ops.dos" assert len(session.files) == expected_count assert len(set(session.files)) == expected_count # 无重复 + + +class TestLazyAlphaLibLoading: + """alpha 因子库按字段引用惰性加载(每会话每库仅一次)。""" + + def _registered_session(self) -> _FakeSession: + session = _FakeSession() + register_ddb_functions_to_qlib(session) # type: ignore[arg-type] + session.files.clear() # 只观察后续惰性加载 + return session + + def test_alpha_field_triggers_single_load(self) -> None: + from qlib.data.backend.ddb_qlib.ddb_features import ensure_alpha_libs_loaded + + session = self._registered_session() + ensure_alpha_libs_loaded(session, ["gtjaAlpha3($open,$close)", "$high"]) + assert session.files == ["gtja191Alpha.dos"] + # 第二次引用同库:不再加载 + ensure_alpha_libs_loaded(session, ["gtjaAlpha5($close)"]) + assert session.files == ["gtja191Alpha.dos"] + + def test_plain_fields_load_nothing(self) -> None: + from qlib.data.backend.ddb_qlib.ddb_features import ensure_alpha_libs_loaded + + session = self._registered_session() + ensure_alpha_libs_loaded(session, ["$close", "Ref($close,1)"]) + assert session.files == [] + + def test_case_insensitive_match(self) -> None: + from qlib.data.backend.ddb_qlib.ddb_features import ensure_alpha_libs_loaded + + session = self._registered_session() + ensure_alpha_libs_loaded(session, ["WQAlpha1($close)", "qlib158Alpha2($open)"]) + assert sorted(session.files) == ["qlib158Alpha.dos", "wq101alpha.dos"] + + def test_preloaded_session_skips_lazy_load(self) -> None: + from qlib.data.backend.ddb_qlib.ddb_features import ensure_alpha_libs_loaded + + session = _FakeSession() + register_ddb_functions_to_qlib(session, preload_alpha_libs=True) # type: ignore[arg-type] + session.files.clear() + ensure_alpha_libs_loaded(session, ["gtjaAlpha3($open,$close)"]) + assert session.files == [] diff --git a/tests/test_ddb_loader_teardown.py b/tests/test_ddb_loader_teardown.py new file mode 100644 index 00000000000..389cd5dc9b3 --- /dev/null +++ b/tests/test_ddb_loader_teardown.py @@ -0,0 +1,47 @@ +"""DolphinDBDataLoader teardown 的离线回归测试。 + +历史 bug:``__exit__``/``__del__`` 无条件调用 ``self.session.close()``, +而 ``self.session`` 是进程级共享的 ``DBClient.session``——一个 with 块或一次 GC +就会把全局会话关掉,破坏 storage 与 feature 路径的所有其他使用方。 +修复后引入 ``_owns_session`` 所有权标志:仅独占会话时才允许关闭。 +""" + +from qlib.data.dataset.loader import DolphinDBDataLoader + + +class _FakeSession: + def __init__(self): + self.close_calls = 0 + + def close(self): + self.close_calls += 1 + + +def _make_loader(owns_session: bool) -> tuple[DolphinDBDataLoader, _FakeSession]: + """绕过构造器(构造器依赖全局 DBClient)注入假会话。""" + loader = object.__new__(DolphinDBDataLoader) + session = _FakeSession() + loader.session = session + loader._owns_session = owns_session + return loader, session + + +def test_shared_session_never_closed(): + """共享会话(默认)在 __exit__ 与 __del__ 中都不得被关闭。""" + loader, session = _make_loader(owns_session=False) + loader.__exit__(None, None, None) + loader.__del__() + assert session.close_calls == 0, "回归:共享的全局 session 被关闭" + + +def test_owned_session_closed_on_exit(): + """独占会话在 __exit__ 时正常关闭。""" + loader, session = _make_loader(owns_session=True) + loader.__exit__(None, None, None) + assert session.close_calls == 1 + + +def test_del_without_attributes_is_safe(): + """构造中途失败(属性未设置)时 __del__ 不应抛出。""" + loader = object.__new__(DolphinDBDataLoader) + loader.__del__() # 不应抛出 diff --git a/tests/test_ddb_lookback.py b/tests/test_ddb_lookback.py new file mode 100644 index 00000000000..9c923a8a6cd --- /dev/null +++ b/tests/test_ddb_lookback.py @@ -0,0 +1,136 @@ +"""滚动算子取前序期(回看窗口外扩)的离线回归测试。 + +- 回看解析优先复用 qlib 算子树 ``get_extended_window_size``(嵌套/双臂/ + 未来引用均覆盖);qlib 无法实例化的表达式退回「正则扫窗口 + 配置兜底」。 +- 计算分支脚本以 ``daysStep,lookbackDays,rightDays`` 尾参把外扩量传给 + DDB 端 ``FeatureEngineeringByDate``,服务器端外扩查询并截断回请求区间。 + +使用 tests/ddb_mocks.py 的 RecordingSession,不依赖 DolphinDB 服务器。 +""" + +import numpy as np +import pandas as pd +import pytest + +from ddb_mocks import RecordingSession, make_calendar + +from qlib.data.backend.ddb_qlib import invalidate_ddb_caches +from qlib.data.backend.ddb_qlib.ddb_features import ( + batch_extended_window, + fetch_features_from_ddb, + get_expression_extended_window, +) + + +@pytest.fixture(autouse=True) +def _clear_caches(): + invalidate_ddb_caches() + yield + invalidate_ddb_caches() + + +class TestExpressionExtendedWindow: + """单表达式回看解析:qlib 算子树路径。""" + + @pytest.mark.parametrize( + "expr,expected", + [ + ("$close", (0, 0)), # 纯字段无外扩 + ("Mean($close,20)", (19, 0)), # 滚动窗口 N -> N-1 + ("Ref($close,5)", (5, 0)), # 引用 N 期前 -> N + ("Mean(Ref($close,5),20)", (24, 0)), # 嵌套:5 + 19 + ("Corr($close,$volume,10)", (9, 0)), # 双臂算子取 max + ("Ref($close,-2)", (0, 2)), # 未来引用 -> 向后外扩 + ("Ref($close,-2)/Ref($close,-1)-1", (0, 2)), # 标签表达式 + ("Mean($close,20)/Std($close,60)", (59, 0)), # 同表达式内取 max + ], + ) + def test_qlib_op_tree(self, expr, expected): + assert get_expression_extended_window(expr, 252) == expected + + def test_fallback_regex_window(self): + """qlib 不认识的函数:独立整数当窗口。""" + assert get_expression_extended_window("myCustomOp($close,60)", 252) == (60, 0) + + def test_fallback_ignores_identifier_digits(self): + """标识符里的数字(gtjaAlpha191_001)不当窗口,走配置兜底。""" + assert get_expression_extended_window("gtjaAlpha191_001($close)", 100) == (100, 0) + + def test_fallback_future_reference(self): + """兜底路径的未来引用:负数窗口仅按函数参数形式识别。""" + lft, rght = get_expression_extended_window("myCustomOp($close, -3)", 252) + assert rght == 3 + + def test_fallback_ignores_large_scaling_constant(self): + """大数值常量(缩放因子)不当窗口,避免误判为百万日回看,退回配置兜底。""" + assert get_expression_extended_window( + "gtjaAlpha191_005($volume/1000000, $close)", 252 + ) == (252, 0) + + def test_fallback_window_upper_bound(self): + """上界内整数仍视为窗口;超过上界视为常量并退回兜底。""" + assert get_expression_extended_window("myCustomOp($close, 2000)", 30)[0] == 2000 + assert get_expression_extended_window("myCustomOp($close, 2001)", 30)[0] == 30 + + def test_fallback_future_ignores_large_constant(self): + """未来引用兜底同样忽略超大常量,避免向后外扩到不合理天数。""" + _, rght = get_expression_extended_window("myCustomOp($close, -9999999)", 252) + assert rght == 0 + + +class TestBatchExtendedWindow: + def test_batch_takes_max(self): + exprs = ["Mean($close,20)", "Ref($close,-2)", "Std($close,60)"] + assert batch_extended_window(exprs, 252) == (59, 2) + + def test_empty_batch(self): + assert batch_extended_window([], 252) == (0, 0) + + +# 3 个交易日 × 2 只股票的固定小样本(与 test_fetch_features.py 一致) +CAL = make_calendar("2024-01-01", 5) +DATES = pd.DatetimeIndex(CAL[:3]) +CODES = ["SH600000", "SZ000001"] +START, END = "2024-01-01", "2024-01-03" + + +def _fe_response(script: str) -> dict: + values = np.arange(6, dtype=float).reshape(3, 2) # dates × codes + return {"ExprName0": [values, DATES, CODES]} + + +def _make_session() -> RecordingSession: + return RecordingSession( + calendar=CAL, + run_responses=[(r"FeatureEngineeringByDate", _fe_response)], + ) + + +class TestLookbackScriptWiring: + """计算分支脚本携带 lookbackDays/rightDays 尾参。""" + + def test_rolling_lookback_in_script(self): + session = _make_session() + fetch_features_from_ddb(session, CODES, ["Mean($close,20)"], START, END, "day") + fe_script = [s for s in session.run_scripts if "FeatureEngineeringByDate" in s][0] + assert ",252,19,0)" in fe_script + + def test_label_right_extension_in_script(self): + session = _make_session() + fetch_features_from_ddb( + session, CODES, ["Ref($close,-2)/Ref($close,-1)-1"], START, END, "day" + ) + fe_script = [s for s in session.run_scripts if "FeatureEngineeringByDate" in s][0] + assert ",252,0,2)" in fe_script + + def test_lookback_default_configurable(self, monkeypatch): + """qlib 解析不了且扫不到窗口的表达式,用 C["ddb_lookback_default"] 兜底。""" + from qlib.config import C + + monkeypatch.setitem(C, "ddb_lookback_default", 30) + session = _make_session() + fetch_features_from_ddb( + session, CODES, ["gtjaAlphaCustom($close)"], START, END, "day" + ) + fe_script = [s for s in session.run_scripts if "FeatureEngineeringByDate" in s][0] + assert ",252,30,0)" in fe_script diff --git a/tests/test_ddb_mysql_bridge.py b/tests/test_ddb_mysql_bridge.py new file mode 100644 index 00000000000..0277b6035de --- /dev/null +++ b/tests/test_ddb_mysql_bridge.py @@ -0,0 +1,46 @@ +"""DDBMySQLBridge 的离线单元测试(不依赖 DolphinDB / MySQL 服务器)。 + +历史 bug 回归:``load_mysql_plugin`` 曾误装 "lgbm" 插件而非 "mysql", +导致后续 ``mysql::connect`` 必然失败。 +""" + +import pandas as pd + +from qlib.data.backend.ddb_qlib.ddb_mysql_bridge import DDBMySQLBridge + + +class _FakeSession: + """记录 run 脚本的假会话。""" + + def __init__(self, raise_on_run: bool = False): + self.scripts: list[str] = [] + self.raise_on_run = raise_on_run + + def run(self, script: str): + if self.raise_on_run: + raise RuntimeError("boom") + self.scripts.append(script) + return pd.DataFrame() + + +def _make_bridge(session: _FakeSession) -> DDBMySQLBridge: + """绕过构造器(构造器会真实连接数据库)注入假会话。""" + bridge = object.__new__(DDBMySQLBridge) + bridge.ddb_session = session + return bridge + + +def test_load_mysql_plugin_installs_mysql_not_lgbm(): + session = _FakeSession() + bridge = _make_bridge(session) + bridge.load_mysql_plugin() + script = "\n".join(session.scripts) + assert 'installPlugin("mysql")' in script + assert 'loadPlugin("mysql")' in script + assert "lgbm" not in script, "回归:曾误装 lgbm 插件" + + +def test_close_swallows_error(): + """close 失败不应中断程序执行。""" + bridge = _make_bridge(_FakeSession(raise_on_run=True)) + bridge.close() # 不应抛出 diff --git a/tests/test_ddb_reshape_equivalence.py b/tests/test_ddb_reshape_equivalence.py new file mode 100644 index 00000000000..d0c2a3efa27 --- /dev/null +++ b/tests/test_ddb_reshape_equivalence.py @@ -0,0 +1,85 @@ +"""D4 结果重塑直构与旧路径的等价性测试。 + +``_computed_dict_to_panel`` 用一次分配直构 (instrument, datetime) 面板, +替代 ``_legacy_reshape`` 的 concat→unstack→stack→swaplevel 多次全景拷贝。 +本文件在各种边界形态下断言两条路径输出完全一致(排序后逐值比较), +并验证形状不一致时的运行时兜底。 +""" + +import numpy as np +import pandas as pd +import pytest + +from qlib.data.backend.ddb_qlib.ddb_features import ( + _computed_dict_to_panel, + _legacy_reshape, +) + + +def _assert_equivalent(data: dict) -> None: + new = _computed_dict_to_panel(data) + legacy = _legacy_reshape(data) + legacy.index.names = ["instrument", "datetime"] + legacy = legacy.sort_index().reindex(columns=list(new.columns)) + pd.testing.assert_frame_equal(new.sort_index(), legacy, check_names=True) + + +def _make(aliases: int, dates, codes, seed: int = 0, nan_ratio: float = 0.0) -> dict: + rng = np.random.default_rng(seed) + out = {} + for i in range(aliases): + values = rng.standard_normal((len(dates), len(codes))) + if nan_ratio: + mask = rng.random(values.shape) < nan_ratio + values = np.where(mask, np.nan, values) + out[f"ExprName{i}"] = [values, pd.DatetimeIndex(dates), list(codes)] + return out + + +DATES = pd.date_range("2024-01-01", periods=4, freq="B") +CODES = ["SZ000001", "SH600000", "SH600519"] # 故意乱序 + + +class TestEquivalence: + def test_basic_multi_alias(self): + _assert_equivalent(_make(3, DATES, CODES)) + + def test_unsorted_codes_and_dates(self): + dates = pd.DatetimeIndex(["2024-01-03", "2024-01-01", "2024-01-02"]) + _assert_equivalent(_make(2, dates, CODES, seed=1)) + + def test_with_nan(self): + _assert_equivalent(_make(2, DATES, CODES, seed=2, nan_ratio=0.3)) + + def test_single_alias(self): + _assert_equivalent(_make(1, DATES, CODES, seed=3)) + + def test_single_date(self): + _assert_equivalent(_make(2, DATES[:1], CODES, seed=4)) + + def test_single_code(self): + _assert_equivalent(_make(2, DATES, CODES[:1], seed=5)) + + +class TestFallback: + def test_shape_mismatch_falls_back_to_legacy(self): + """alias 间形状不一致时回退 legacy(不抛错、语义与旧版一致)。""" + data = _make(1, DATES, CODES) + # 第二个 alias 的日期轴不同(模拟 DDB 端异常占位返回) + other_dates = pd.date_range("2024-02-01", periods=4, freq="B") + data.update(_make(1, other_dates, CODES, seed=9)) + # legacy 能处理(unstack 对齐并集),直构应回退到相同结果 + result = _computed_dict_to_panel(data) + legacy = _legacy_reshape(data) + legacy.index.names = ["instrument", "datetime"] + pd.testing.assert_frame_equal( + result.sort_index(), legacy.sort_index().reindex(columns=list(result.columns)) + ) + + def test_malformed_entry_falls_back(self): + """条目缺轴(长度不足 3)时不抛错——与 legacy 行为对齐或回退。""" + data = _make(2, DATES, CODES) + data["ExprName0"] = [np.zeros((2, 2))] # 缺日期/代码轴 + with pytest.raises(Exception): + # legacy 与直构在此病态输入下都应报错(IndexError 语义保留) + _computed_dict_to_panel(data) diff --git a/tests/test_ddb_storage_caches.py b/tests/test_ddb_storage_caches.py new file mode 100644 index 00000000000..6581100661d --- /dev/null +++ b/tests/test_ddb_storage_caches.py @@ -0,0 +1,153 @@ +"""D3 存储层/schema/表达式翻译缓存的离线回归测试。 + +历史问题: +1. 每个存储访问器都 ``check()`` → ``existsTable`` RPC,无缓存; +2. ``DBInstrumentStorage`` 每次访问全量下载股票池表并行循环重建 dict; +3. ``build_field_expr``/``table_appender``/``table_upsert`` 每次调用重新 + 拉取表 schema; +4. ``adapt_qlib_expr_syntax_for_ddb``(~100 行递归解析)每字段每次重跑。 + +修复:existsTable 仅正向缓存、股票池走 ``H["i"]``、schema 走 +``get_table_columns`` 进程内缓存、表达式翻译走 lru_cache; +前三者由 ``invalidate_ddb_caches()`` 统一失效。 +""" + +import threading + +import pandas as pd +import pytest + +from ddb_mocks import RecordingSession + +from qlib.data.backend.ddb_qlib import invalidate_ddb_caches +from qlib.data.backend.ddb_qlib.ddb_features import ( + OPERATOR_MAPPING, + _adapt_cached, + adapt_qlib_expr_syntax_for_ddb, +) +from qlib.data.backend.ddb_qlib.utils import get_table_columns +from qlib.data.data import DBClient + + +@pytest.fixture(autouse=True) +def _clean_caches(): + invalidate_ddb_caches() + yield + invalidate_ddb_caches() + + +class _FakeProvider: + def __init__(self, session): + self.session = session + self.session_lock = threading.RLock() + + +@pytest.fixture() +def fake_provider(): + session = RecordingSession( + table_results={ + ("dfs://QlibInstruments", "csi300"): pd.DataFrame( + { + "instrument": ["SH600000", "SH600000", "SZ000001"], + "start_datetime": pd.to_datetime(["2020-01-01", "2022-01-01", "2020-01-01"]), + "end_datetime": pd.to_datetime(["2021-01-01", "2099-12-31", "2099-12-31"]), + } + ) + } + ) + old = DBClient.__dict__.get("_provider") + DBClient.register(_FakeProvider(session)) + yield session + DBClient.register(old) + + +class TestExistsCache: + def test_positive_result_cached(self, fake_provider): + from qlib.data.storage.dolphindb_storage import DBFeatureStorage + + storage = DBFeatureStorage(instrument="SH600000", field="close", freq="day") + storage.check() + storage.check() + assert fake_provider.counts["existsTable"] == 1, "正向 exists 结果应被缓存" + + def test_negative_result_not_cached(self, fake_provider): + """表不存在必须持续报错(负结果不缓存),表创建后立即可见。""" + from qlib.data.storage.dolphindb_storage import DBFeatureStorage + + fake_provider.exists_result = False + storage = DBFeatureStorage(instrument="SH600000", field="close", freq="day") + with pytest.raises(ValueError): + storage.check() + with pytest.raises(ValueError): + storage.check() + assert fake_provider.counts["existsTable"] == 2, "负结果不得缓存" + # 表被创建后(exists 变 True)无需失效缓存即可通过 + fake_provider.exists_result = True + storage.check() + + +class TestInstrumentCache: + def test_read_once_across_accessors(self, fake_provider): + from qlib.data.storage.dolphindb_storage import DBInstrumentStorage + + storage = DBInstrumentStorage(market="csi300", freq="day") + data_a = storage.data + data_b = storage.data + assert len(storage) == 2 + assert fake_provider.counts["loadTable"] == 1, "股票池应经 H['i'] 缓存" + # 内容正确且跨调用一致 + assert data_a == data_b + assert data_a["SH600000"] == [ + (pd.Timestamp("2020-01-01"), pd.Timestamp("2021-01-01")), + (pd.Timestamp("2022-01-01"), pd.Timestamp("2099-12-31")), + ] + # 返回浅拷贝:调用方修改不污染缓存 + data_a["FAKE"] = [] + assert "FAKE" not in storage.data + + def test_invalidate_forces_reload(self, fake_provider): + from qlib.data.storage.dolphindb_storage import DBInstrumentStorage + + storage = DBInstrumentStorage(market="csi300", freq="day") + _ = storage.data + invalidate_ddb_caches() + _ = storage.data + assert fake_provider.counts["loadTable"] == 2 + + +class TestSchemaCache: + def test_columns_loaded_once(self): + session = RecordingSession(table_columns={("dfs://Db", "Tb"): ["code", "date", "close"]}) + cols_a = get_table_columns(session, "dfs://Db", "Tb") + cols_b = get_table_columns(session, "dfs://Db", "Tb") + assert cols_a == cols_b == ["code", "date", "close"] + assert session.counts["loadTable"] == 1 + # 返回副本:调用方修改不污染缓存 + cols_a.append("hacked") + assert get_table_columns(session, "dfs://Db", "Tb") == ["code", "date", "close"] + + def test_invalidate_forces_reload(self): + session = RecordingSession(table_columns={("dfs://Db", "Tb"): ["code"]}) + get_table_columns(session, "dfs://Db", "Tb") + invalidate_ddb_caches() + get_table_columns(session, "dfs://Db", "Tb") + assert session.counts["loadTable"] == 2 + + +class TestExpressionMemoization: + def test_cache_hit_and_identical_output(self): + _adapt_cached.cache_clear() + exprs = ["Ref($close,1)/$close-1", "Mean($volume,5)", "Std($close,20)"] + first = [adapt_qlib_expr_syntax_for_ddb(e, OPERATOR_MAPPING, True) for e in exprs] + second = [adapt_qlib_expr_syntax_for_ddb(e, OPERATOR_MAPPING, True) for e in exprs] + assert first == second + assert _adapt_cached.cache_info().hits >= len(exprs) + + def test_custom_mapping_bypasses_cache(self): + """自定义映射(非默认 dict)不得命中默认映射的缓存。""" + custom = dict(OPERATOR_MAPPING) + custom["Ref"] = "customMove" + result = adapt_qlib_expr_syntax_for_ddb("Ref($close,1)", custom) + assert "customMove(" in result + # 默认映射结果不受影响 + assert "move(" in adapt_qlib_expr_syntax_for_ddb("Ref($close,1)") diff --git a/tests/test_ddb_storage_spans.py b/tests/test_ddb_storage_spans.py index 106bca46673..e01f0bc49d6 100644 --- a/tests/test_ddb_storage_spans.py +++ b/tests/test_ddb_storage_spans.py @@ -18,7 +18,7 @@ (出池不截断、调离不剔除)。 DDB 端本身对 dict 是完好支持的(纯字段分支 ``createDateStockMapping`` + - ``conditionalFilter``;非纯字段分支 ``FeatureEngine.buildSpanAwareWhereConditions``), + ``conditionalFilter``;非纯字段分支 ``FeatureEngine.buildWhereConditions``), 所以修复只需在存储层保持 dict 原型(仅键做 ``.upper()``)。 本测试为离线纯单元测试,锁定 ``DBFeatureStorage.__init__`` 三种入参 diff --git a/tests/test_fetch_features.py b/tests/test_fetch_features.py new file mode 100644 index 00000000000..985c410eff3 --- /dev/null +++ b/tests/test_fetch_features.py @@ -0,0 +1,258 @@ +"""fetch_features_from_ddb 四大分支的离线回归测试(D0 安全网)。 + +在任何性能重构(往返缩减/缓存/重塑重写)之前固化当前行为: +- 纯字段 × list / dict(spans) 两分支 +- 计算表达式 × list / dict(spans) 两分支 +- 空 instruments 与 DDB 空 dict 的早退路径 +- RPC 往返数基线(后续 D1/D2 优化后按目标值更新) + +使用 tests/ddb_mocks.py 的 RecordingSession,不依赖 DolphinDB 服务器。 +""" + +import numpy as np +import pandas as pd +import pytest + +from ddb_mocks import FakeQueryChain, RecordingSession, make_calendar + +from qlib.data.backend.ddb_qlib import invalidate_ddb_caches +from qlib.data.backend.ddb_qlib.ddb_features import fetch_features_from_ddb + + +@pytest.fixture(autouse=True) +def _clear_caches(): + """D1/D3 引入进程内缓存后,每个测试都从干净缓存开始以保证计数确定性。""" + invalidate_ddb_caches() + yield + invalidate_ddb_caches() + +# 3 个交易日 × 2 只股票的固定小样本 +CAL = make_calendar("2024-01-01", 5) +DATES = pd.DatetimeIndex(CAL[:3]) +CODES = ["SH600000", "SZ000001"] +START, END = "2024-01-01", "2024-01-03" + +FEATURE_KEY = ("dfs://QlibFeaturesDay", "Features") + + +def _pure_table_result(script: str) -> pd.DataFrame: + """按查询构造纯字段分支的长表结果(code/date/close/open)。""" + rows = [ + {"code": c, "date": d, "close": float(i + 1), "open": float((i + 1) * 10)} + for i, (c, d) in enumerate((c, d) for d in DATES for c in CODES) + ] + return pd.DataFrame(rows) + + +def _make_pure_session() -> RecordingSession: + # D2 后纯字段分支为单条 SQL 脚本(select ... from loadTable(...)) + return RecordingSession( + calendar=CAL, + table_columns={FEATURE_KEY: ["code", "date", "close", "open", "volume"]}, + run_responses=[(r"select .*from loadTable", _pure_table_result)], + ) + + +def _fe_response(script: str) -> dict: + """计算分支:模拟 FeatureEngineeringByDate 返回 {alias: [矩阵, 日期, 代码]}。""" + values0 = np.arange(6, dtype=float).reshape(3, 2) # dates × codes + values1 = values0 * 100 + return { + "ExprName0": [values0, DATES, CODES], + "ExprName1": [values1, DATES, CODES], + } + + +def _make_computed_session() -> RecordingSession: + return RecordingSession( + calendar=CAL, + run_responses=[(r"FeatureEngineeringByDate", _fe_response)], + ) + + +SPANS = { + "SH600000": [(pd.Timestamp("2024-01-01"), pd.Timestamp("2024-01-02"))], + "SZ000001": [(pd.Timestamp("2024-01-01"), pd.Timestamp("2024-01-03"))], +} + + +class TestPureFieldsBranch: + def test_list_instruments(self): + session = _make_pure_session() + df = fetch_features_from_ddb(session, CODES, ["$close", "$open"], START, END, "day") + + assert list(df.columns) == ["$close", "$open"] + assert df.index.names == ["instrument", "datetime"] + assert len(df) == 6 + # 单脚本查询:code in instruments 过滤 + 日期字面量 + 排序 + script = session.run_scripts[-1] + assert "code in instruments" in script + assert "date between pair(2024.01.01,2024.01.03)" in script + assert "order by date, code" in script + # 上传的变量含 instruments(且仅 instruments——纯字段分支不再上传表达式) + assert session.uploads[-1] == {"instruments": CODES} + + def test_dict_instruments_uses_conditional_filter(self): + session = _make_pure_session() + df = fetch_features_from_ddb(session, SPANS, ["$close"], START, END, "day") + + assert not df.empty + # spans 路径:映射创建与 conditionalFilter 主查询合并在同一脚本(单次往返) + script = session.run_scripts[-1] + assert "createDateStockMapping(2024.01.01,2024.01.03,instruments)" in script + assert "conditionalFilter(code, date, codeRangeFilter)" in script + assert session.counts["run"] == 1 + + def test_missing_base_field_padded_with_zero(self): + """表中不存在的基础字段用 '0 as 字段' 兜底。""" + session = RecordingSession( + calendar=CAL, + table_columns={FEATURE_KEY: ["code", "date", "close"]}, + run_responses=[ + ( + r"select .*from loadTable", + pd.DataFrame( + {"code": CODES, "date": [DATES[0]] * 2, "close": [1.0, 2.0], "not_exist": [0, 0]} + ), + ) + ], + ) + fetch_features_from_ddb(session, CODES, ["$close", "$not_exist"], START, END, "day") + assert "0 as not_exist" in session.run_scripts[-1] + + +class TestComputedBranch: + FIELDS = ["Ref($close,1)", "Mean($open,5)"] + + def test_list_instruments(self): + session = _make_computed_session() + df = fetch_features_from_ddb(session, CODES, self.FIELDS, START, END, "day") + + # 列名映射回原始表达式(空格移除后) + assert list(df.columns) == ["Ref($close,1)", "Mean($open,5)"] + assert df.index.names == ["instrument", "datetime"] + # 全组合:2 codes × 3 dates + assert len(df) == 6 + assert df.index.is_monotonic_increasing + # 数值正确性:ExprName0 矩阵为 dates×codes 的 0..5 + assert df.loc[("SH600000", DATES[0]), "Ref($close,1)"] == 0.0 + assert df.loc[("SZ000001", DATES[2]), "Ref($close,1)"] == 5.0 + assert df.loc[("SZ000001", DATES[2]), "Mean($open,5)"] == 500.0 + # 脚本用 instruments 变量(list 分支) + fe_script = [s for s in session.run_scripts if "FeatureEngineeringByDate" in s][0] + assert "FeatureEngineeringByDate(instruments," in fe_script + + def test_dict_instruments_applies_spans_mask(self): + session = _make_computed_session() + df = fetch_features_from_ddb(session, SPANS, self.FIELDS, START, END, "day") + + # dict 分支:脚本传键列表,spans 掩码在 Python 侧补齐 + fe_script = [s for s in session.run_scripts if "FeatureEngineeringByDate" in s][0] + assert "FeatureEngineeringByDate(keys(instruments)," in fe_script + # SH600000 的 2024-01-03 已出池,应被掩码剔除 + assert ("SH600000", DATES[2]) not in df.index + assert ("SZ000001", DATES[2]) in df.index + assert len(df) == 5 + + def test_days_step_default_and_configurable(self, monkeypatch): + """D7:日期分片经 C[\"ddb_days_step\"] 可配置(默认 252,行为不变)。 + + 脚本尾参为 daysStep,lookbackDays,rightDays;FIELDS 的批回看窗口为 + max(Ref(...,1)->1, Mean(...,5)->4) = 4 个交易日。 + """ + from qlib.config import C + + session = _make_computed_session() + fetch_features_from_ddb(session, CODES, self.FIELDS, START, END, "day") + assert ',252,4,0)' in [s for s in session.run_scripts if "FeatureEngineeringByDate" in s][0] + + monkeypatch.setitem(C, "ddb_days_step", 126) + session2 = _make_computed_session() + fetch_features_from_ddb(session2, CODES, self.FIELDS, START, END, "day") + assert ',126,4,0)' in [s for s in session2.run_scripts if "FeatureEngineeringByDate" in s][0] + + def test_unrecognized_token_triggers_alpha_retry(self): + """D5 兜底:未识别函数错误时全量加载 alpha 库并重试一次。""" + calls = {"n": 0} + + def _fe(script): + calls["n"] += 1 + if calls["n"] == 1: + raise RuntimeError("Syntax Error: Cannot recognize the token gtjaAlpha3") + return _fe_response(script) + + session = RecordingSession( + calendar=CAL, run_responses=[(r"FeatureEngineeringByDate", _fe)] + ) + df = fetch_features_from_ddb(session, CODES, self.FIELDS, START, END, "day") + assert not df.empty + loaded = {str(f).rsplit("/", 1)[-1] for f in session.run_files} + assert {"gtja191Alpha.dos", "qlib158Alpha.dos", "wq101alpha.dos"} <= loaded + assert calls["n"] == 2 + + def test_non_token_error_does_not_retry(self): + """非未识别函数错误:不重试、保留 RuntimeError 语义与异常链。""" + + def _fe(script): + raise RuntimeError("Server response: out of memory") + + session = RecordingSession( + calendar=CAL, run_responses=[(r"FeatureEngineeringByDate", _fe)] + ) + with pytest.raises(RuntimeError, match="DolphinDB 因子计算失败"): + fetch_features_from_ddb(session, CODES, self.FIELDS, START, END, "day") + assert session.run_files == [] # 未触发兜底加载 + + def test_empty_result_dict_returns_empty(self): + session = RecordingSession( + calendar=CAL, + run_responses=[(r"FeatureEngineeringByDate", {})], + ) + df = fetch_features_from_ddb(session, CODES, self.FIELDS, START, END, "day") + assert df.empty + + +class TestEarlyReturns: + def test_empty_instruments_returns_empty(self): + session = _make_computed_session() + df = fetch_features_from_ddb(session, [], ["Ref($close,1)"], START, END, "day") + assert df.empty + # 早退不应触发上传与主查询 + assert session.counts["upload"] == 0 + + +class TestRpcCountBaseline: + """RPC 往返数基线(优化后按新目标更新本测试)。 + + D1(日历缓存)+ D2(往返缩减)后的目标状态: + 计算分支每批 = 1 upload + 1 run(日历预热后); + 纯字段分支 = 1 upload + 1 run + schema 查询(D3 缓存后消失)。 + """ + + def test_computed_branch_rpc_counts(self): + session = _make_computed_session() + fetch_features_from_ddb(session, CODES, ["Ref($close,1)"], START, END, "day") + assert session.counts["loadTable"] == 1 # 仅日历首次下载(跨调用共享) + assert session.counts["run"] == 1 # 仅主查询(日期已内联为字面量) + assert session.counts["upload"] == 1 + + def test_pure_list_branch_rpc_counts(self): + session = _make_pure_session() + fetch_features_from_ddb(session, CODES, ["$close"], START, END, "day") + # 首次调用:日历下载 + build_field_expr schema 共 2 次 loadTable + assert session.counts["loadTable"] == 2 + assert session.counts["run"] == 1 # 单脚本主查询 + assert session.counts["upload"] == 1 + # 二次调用:日历与 schema 均命中缓存,仅 1 upload + 1 run + fetch_features_from_ddb(session, CODES, ["$close"], START, END, "day") + assert session.counts["loadTable"] == 2, "D3 回归:schema/日历缓存未生效" + assert session.counts["run"] == 2 + assert session.counts["upload"] == 2 + + def test_calendar_downloaded_once_across_calls(self): + """D1 回归:日历经模块级缓存共享,跨调用只下载一次(曾经每次都下载)。""" + session = _make_computed_session() + fetch_features_from_ddb(session, CODES, ["Ref($close,1)"], START, END, "day") + fetch_features_from_ddb(session, CODES, ["Ref($close,1)"], START, END, "day") + calendar_loads = [c for c in session.load_table_calls if c[0] == "dfs://QlibCalendars"] + assert len(calendar_loads) == 1, "日历缓存失效:跨调用重复全量下载" diff --git a/tests/test_operator_mapping.py b/tests/test_operator_mapping.py new file mode 100644 index 00000000000..967e1d39823 --- /dev/null +++ b/tests/test_operator_mapping.py @@ -0,0 +1,44 @@ +"""OPERATOR_MAPPING 的守卫测试。 + +历史问题:dict 字面量中 ``"Slope"``/``"Resi"`` 出现过两次(先映射到 DDB 内置的 +``mslr``/``mmse``,后映射到 ops.dos 自定义的 ``Slope``/``Resi``),后者静默覆盖 +前者——生效行为一直是 ops.dos 版本。清理后本测试锁定: +1. 生效映射不变(ops.dos 自定义算子); +2. dict 字面量中不再允许出现重复键(重复键会被 Python 静默覆盖,难以察觉)。 +""" + +import ast +import inspect + +import qlib.data.backend.ddb_qlib.ddb_features as ddb_features_mod +from qlib.data.backend.ddb_qlib.ddb_features import OPERATOR_MAPPING + + +def test_ops_dos_custom_operators_are_live(): + """ops.dos 自定义算子必须是生效映射(曾被前置重复键遮蔽混淆)。""" + assert OPERATOR_MAPPING["Slope"] == "Slope" + assert OPERATOR_MAPPING["Resi"] == "Resi" + assert OPERATOR_MAPPING["Rsquare"] == "Rsquare" + # 抽查若干常规映射未被误删 + assert OPERATOR_MAPPING["Ref"] == "move" + assert OPERATOR_MAPPING["Mean"] == "mavg" + assert OPERATOR_MAPPING["Std"] == "mstd" + + +def test_no_duplicate_keys_in_mapping_literal(): + """dict 字面量禁止重复键(后者静默覆盖前者,属隐蔽 bug 温床)。""" + tree = ast.parse(inspect.getsource(ddb_features_mod)) + for node in ast.walk(tree): + # OPERATOR_MAPPING: Dict = {...} 是 AnnAssign;兼容无注解的 Assign 写法 + if isinstance(node, ast.AnnAssign): + target_id = getattr(node.target, "id", "") + elif isinstance(node, ast.Assign): + target_id = getattr(node.targets[0], "id", "") + else: + continue + if target_id == "OPERATOR_MAPPING" and isinstance(node.value, ast.Dict): + keys = [k.value for k in node.value.keys if isinstance(k, ast.Constant)] + dupes = {k for k in keys if keys.count(k) > 1} + assert not dupes, f"OPERATOR_MAPPING 存在重复键: {dupes}" + return + raise AssertionError("未找到 OPERATOR_MAPPING 字面量定义")