diff --git a/tests/backtest/test_cta_engine.py b/tests/backtest/test_cta_engine.py index 7ca012c..555be72 100644 --- a/tests/backtest/test_cta_engine.py +++ b/tests/backtest/test_cta_engine.py @@ -1,26 +1,30 @@ """Tests for sanguo_backtest.cta_engine module.""" # Mock vnpy and tzlocal modules before importing anything that depends on them import sys +import importlib.util from unittest.mock import MagicMock -# Mock vnpy/tzlocal/empyrical BEFORE importing cta_engine (which imports them at top). -# Save originals so we can RESTORE after this module — otherwise the global sys.modules -# pollution breaks every later test that imports real vnpy/empyrical (datareader/spike/factor/metrics). -_MOCKED_KEYS = ( - "tzlocal", "vnpy.trader.setting", "vnpy.trader.constant", "vnpy.trader.object", - "vnpy.trader.database", "vnpy_ctastrategy.backtesting", "empyrical", -) -_SAVED_MODULES = {k: sys.modules.get(k) for k in _MOCKED_KEYS} - +# Only mock these when the real module is NOT importable (e.g. local env missing empyrical). +# In the container (authoritative test env) everything imports fine, so NO mock is installed +# → zero sys.modules pollution leaking into other test modules at collection time. +# (Previous unconditional sys.modules[...]=MagicMock() here broke datareader/spike/factor/metrics +# tests because it ran at collection time and never restored.) mock_tzlocal = MagicMock() mock_tzlocal.get_localzone_name = MagicMock(return_value="UTC") -sys.modules["tzlocal"] = mock_tzlocal -sys.modules["vnpy.trader.setting"] = MagicMock() -sys.modules["vnpy.trader.constant"] = MagicMock() -sys.modules["vnpy.trader.object"] = MagicMock() -sys.modules["vnpy.trader.database"] = MagicMock() -sys.modules["vnpy_ctastrategy.backtesting"] = MagicMock() -sys.modules["empyrical"] = MagicMock() +_MOCK_FACTORIES = { + "tzlocal": lambda: mock_tzlocal, + "vnpy.trader.setting": MagicMock, + "vnpy.trader.constant": MagicMock, + "vnpy.trader.object": MagicMock, + "vnpy.trader.database": MagicMock, + "vnpy_ctastrategy.backtesting": MagicMock, + "empyrical": MagicMock, +} +_SAVED_MODULES = {} +for _name, _factory in _MOCK_FACTORIES.items(): + if importlib.util.find_spec(_name) is None: + _SAVED_MODULES[_name] = sys.modules.get(_name) + sys.modules[_name] = _factory() import pytest import json