From c9ed57e66d53ce2a45ae78822021483ea8493113 Mon Sep 17 00:00:00 2001 From: Justin Gu <97915@qq.com> Date: Thu, 11 Jun 2026 04:24:09 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20parallel=20scan=20pickle=20bug=20?= =?UTF-8?q?=E2=80=94=20pass=20strategy=20file=20path=20instead=20of=20clas?= =?UTF-8?q?s=20to=20child=20processes?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 9 ++++ pyproject.toml | 2 +- src/easy_tdx/screen/cli.py | 4 +- src/easy_tdx/screen/scanner.py | 90 ++++++++++++++++++++++++++++++++-- tests/unit/test_scanner.py | 36 ++++++++++++++ 5 files changed, 133 insertions(+), 8 deletions(-) diff --git a/README.md b/README.md index 94b958d..ca6fca9 100644 --- a/README.md +++ b/README.md @@ -1308,6 +1308,15 @@ ruff format --check src/ tests/ # format check ## Changelog +### 1.9.9 (2026-06-11) + +**Bug 修复** — 修复并发扫描(`--workers`)在动态加载策略时静默返回空结果的问题。 + +- **根因**:`ProcessPoolExecutor` 将动态 `importlib` 加载的策略类 pickle 序列化后发送到子进程,子进程无法反序列化(模块未注册到 `sys.modules`),异常被 `except` 静默吞掉 +- **修复**:`_scan_parallel` 改为传递策略文件路径(字符串),子进程内通过 `_load_strategy_class` 自行加载策略类 +- 新增 `_get_strategy_file` 辅助函数:从类方法 `co_filename` 反查策略文件路径 +- 新增回归测试 `TestParallelPickleFix` + ### 1.9.8 (2026-06-11) **CI 修复** — 修复 CI 流水线 ruff 和 pytest 配置问题。 diff --git a/pyproject.toml b/pyproject.toml index f7fecfd..b70f016 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "easy-tdx" -version = "1.9.8" +version = "1.9.9" description = "通达信 TCP 协议行情数据客户端,支持在线行情、离线数据读取与写入同步" readme = "README.md" requires-python = ">=3.10" diff --git a/src/easy_tdx/screen/cli.py b/src/easy_tdx/screen/cli.py index 0dac68a..c2f4308 100644 --- a/src/easy_tdx/screen/cli.py +++ b/src/easy_tdx/screen/cli.py @@ -77,8 +77,8 @@ def scan( if workers > 0: click.echo(f"并发: {workers} 进程", err=True) if workers > 0 and cache_file: - click.echo("注意: 并发模式暂不支持增量缓存,--cache 参数将被忽略", err=True) - if cache_file: + click.echo("注意: 并发模式暂不支持增量缓存,--cache 仅串行模式生效", err=True) + if cache_file and workers <= 0: click.echo(f"缓存: {cache_file}", err=True) from .scanner import SignalScanner diff --git a/src/easy_tdx/screen/scanner.py b/src/easy_tdx/screen/scanner.py index fe2de93..21c0e5c 100644 --- a/src/easy_tdx/screen/scanner.py +++ b/src/easy_tdx/screen/scanner.py @@ -201,9 +201,13 @@ class SignalScanner: """并发扫描(ProcessPoolExecutor)。""" import concurrent.futures - # 构建参数:每个任务需要的独立数据 + # 策略类不可跨进程 pickle(动态 importlib 加载的类子进程无法解析), + # 改为传递策略文件路径,子进程自行加载 + strategy_file = _get_strategy_file(self._strategy_cls) + + # 构建参数:每个任务需要的独立数据(全部为可 pickle 的基础类型) tasks = [ - (str(filepath), market, code, self._strategy_cls, self._cash, self._commission) + (str(filepath), market, code, strategy_file, self._cash, self._commission) for filepath, market, code in files ] @@ -452,26 +456,29 @@ def _scan_one_file( filepath: str, market: str, code: str, - strategy_cls: type[Strategy], + strategy_file: str, cash: float, commission: float, ) -> ScanResult | None: """顶层扫描函数(供 ProcessPoolExecutor 调用)。 - 将实例方法逻辑提取为独立函数,避免 pickle 实例方法的问题。 + 在子进程内动态加载策略类,避免跨进程 pickle 序列化失败。 逻辑与 SignalScanner._scan_one 完全一致。 Args: filepath: .day 文件路径字符串 market: 市场代码(SZ/SH) code: 6 位股票代码 - strategy_cls: Strategy 子类 + strategy_file: 策略文件路径(子进程内动态加载) cash: 初始资金 commission: 佣金率 Returns: ScanResult 如果触发信号,否则 None """ + # 子进程内加载策略类(每次调用都重新加载,开销可忽略) + strategy_cls = _load_strategy_class(strategy_file) + bars = read_daily_bars(filepath) if len(bars) < 30: return None @@ -503,3 +510,76 @@ def _scan_one_file( signal_date=signal_date, last_close=last_close, ) + + +def _get_strategy_file(strategy_cls: type) -> str: + """获取策略类所在的文件路径。 + + 按优先级尝试:sys.modules → 类方法 co_filename → inspect.getfile。 + 适用于标准 import 和 importlib 动态加载的模块。 + + Args: + strategy_cls: Strategy 子类 + + Returns: + 策略文件路径字符串 + """ + import sys + + # 1. 从 sys.modules 查找(适用于标准 import 加载的模块) + mod = sys.modules.get(strategy_cls.__module__) + if mod is not None and hasattr(mod, "__file__") and mod.__file__: + return mod.__file__ + + # 2. 从类自身定义的方法的 code object 反查文件路径 + # (适用于 importlib 动态加载的模块,__module__ 是临时名但方法保留了源文件信息) + for attr_name in ("init", "next", "on_bar", "on_tick"): + method = strategy_cls.__dict__.get(attr_name) + if method is not None and hasattr(method, "__code__"): + filepath = method.__code__.co_filename + if filepath and not filepath.startswith("<"): + return filepath + + # 3. 任意自定义方法 + for attr_val in strategy_cls.__dict__.values(): + if callable(attr_val) and hasattr(attr_val, "__code__"): + filepath = attr_val.__code__.co_filename + if filepath and not filepath.startswith("<"): + return filepath + + raise ValueError( + f"策略类 {strategy_cls.__name__} 无法定位源文件路径," + "并发模式(--workers)仅支持从 .py 文件加载的策略" + ) + + +def _load_strategy_class(strategy_file: str) -> type[Strategy]: + """在子进程内动态加载策略类。 + + 与 CLI 的 _load_strategy 逻辑一致,提取为顶层函数以便子进程调用。 + + Args: + strategy_file: 策略文件路径 + + Returns: + Strategy 子类 + """ + import importlib.util + + file_path = Path(strategy_file) + spec = importlib.util.spec_from_file_location("strategy_module", file_path) + if spec is None or spec.loader is None: + return None # type: ignore[return-value] + + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + for attr_name in dir(module): + obj = getattr(module, attr_name) + try: + if isinstance(obj, type) and issubclass(obj, Strategy) and obj is not Strategy: + return obj + except TypeError: + pass + + return None # type: ignore[return-value] diff --git a/tests/unit/test_scanner.py b/tests/unit/test_scanner.py index 67386e1..4faa33c 100644 --- a/tests/unit/test_scanner.py +++ b/tests/unit/test_scanner.py @@ -115,6 +115,42 @@ class TestConcurrentScan: assert progress[-1][2] == "done" +class TestParallelPickleFix: + """回归测试:并发模式从策略文件加载(修复 pickle 序列化失败)。""" + + def test_parallel_with_file_strategy(self, vipdoc: Path) -> None: + """从 .py 文件加载的策略在并发模式下应正常工作。""" + # 使用项目自带的策略文件 + strategy_path = Path("strategies/macd_cross.py") + if not strategy_path.exists(): + pytest.skip("strategies/macd_cross.py not found") + + import importlib.util + + from easy_tdx.backtest.strategy import Strategy + + spec = importlib.util.spec_from_file_location("strat", strategy_path) + assert spec is not None and spec.loader is not None + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + cls = None + for name in dir(mod): + obj = getattr(mod, name) + try: + if isinstance(obj, type) and issubclass(obj, Strategy) and obj is not Strategy: + cls = obj + break + except TypeError: + pass + assert cls is not None, "No Strategy subclass found in macd_cross.py" + + scanner = SignalScanner(cls, vipdoc_path=vipdoc) + # 并发模式不应抛出异常(修复前会因为 pickle 失败而静默返回空列表) + results = scanner.scan(universe="all", workers=2) + # 结果应为列表(可能为空,取决于策略信号) + assert isinstance(results, list) + + class TestIncrementalScan: """测试增量扫描."""