Files
sanguo_vnpy_v2/tests/factor/test_alpha_lab.py
T
claude_dev a1048690c1 test: 前后端对齐—容器复跑验证+清理僵尸测试
- test_main: FastAPI 0.139 _IncludedRouter 不再 flatten,改用 TestClient 探测路由
- datareader: 文件名(sh600000_daily)/patch target(vnpy.trader.database) 对齐 lazy import 实现
- alpha_lab/analyzer/data_adapter: vnpy.alpha/alphalens 容器专用本地 skip
- 删 4 个测废弃 sanguo_web 的僵尸测试(-1267 行死代码)
- pytest.ini: asyncio_mode=auto
- frontend: package.json 加 test script(npm test 可跑)
- NAS 容器 309 passed 全绿验证(Python 3.10,本机 303+6skip)
2026-07-11 11:45:26 +08:00

350 lines
14 KiB
Python

"""Test alpha_lab module - AlphaLab session management."""
import sys
import os
_VNPY_SRC = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "vnpy_v4.4.0"))
sys.path.insert(0, _VNPY_SRC)
from unittest.mock import Mock, patch
import tempfile
def test_builtin_factors_registered_on_import():
"""Test that importing sanguo_factor registers built-in factors (Root cause A fix)."""
# Reimport to ensure registration runs
import importlib
import sanguo_factor
importlib.reload(sanguo_factor)
from sanguo_factor.registry import get_factor
# Verify built-in factors are registered
ma5_factor = get_factor("ma5")
assert ma5_factor is not None, "ma5 factor not registered after import"
assert ma5_factor["expression"] == "ts_mean(close, 5)"
assert ma5_factor["category"] == "builtin"
# Verify other built-in factors
assert get_factor("ma10") is not None, "ma10 factor not registered"
assert get_factor("ma20") is not None, "ma20 factor not registered"
assert get_factor("vol_ma5") is not None, "vol_ma5 factor not registered"
def test_alpha_lab_session_init():
"""Test AlphaLabSession initialization without calling real __init__."""
from pathlib import Path
from sanguo_factor.alpha_lab import AlphaLabSession
with tempfile.TemporaryDirectory() as tmpdir:
lab_path = str(Path(tmpdir) / "alpha_lab")
# Patch __init__ to skip AlphaLab import
with patch.object(AlphaLabSession, "__init__", lambda self, lab_path: None):
session = AlphaLabSession(lab_path=lab_path)
session.lab_path = lab_path
assert session.lab_path == lab_path
def test_load_symbols_calls_read_db_daily():
"""Test that load_symbols method exists on AlphaLabSession."""
from pathlib import Path
from sanguo_factor.alpha_lab import AlphaLabSession
with tempfile.TemporaryDirectory() as tmpdir:
lab_path = str(Path(tmpdir) / "alpha_lab")
# Patch __init__ to skip AlphaLab import
with patch.object(AlphaLabSession, "__init__", lambda self, lab_path: setattr(self, "lab_path", lab_path)), \
patch.object(AlphaLabSession, "load_symbols"):
session = AlphaLabSession(lab_path=lab_path)
# Verify load_symbols method exists
assert hasattr(session, "load_symbols")
def test_compute_factors_calls_prepare_and_fetch(tmp_path):
"""Test compute_factors calls AlphaDataset methods correctly.
Requires polars - runs in container, skips locally.
"""
import pytest
pytest.importorskip("polars")
# vnpy.alpha 仅 vnpy_v4.4.0 源码提供(容器内 sys.path 引用);本地 pip 装的 vnpy 无此子模块
try:
import vnpy.alpha # noqa: F401
except ImportError:
pytest.skip("vnpy.alpha 仅容器内可用(vnpy_v4.4.0 源码)")
from pathlib import Path
from datetime import datetime
from unittest.mock import MagicMock
from sanguo_factor.alpha_lab import AlphaLabSession
import polars as pl
# Create mock bar data
mock_bar = MagicMock()
mock_bar.vt_symbol = "600000.SSE"
mock_bar.datetime = datetime(2024, 1, 1)
mock_bar.open_price = 1.0
mock_bar.high_price = 1.0
mock_bar.low_price = 1.0
mock_bar.close_price = 1.0
mock_bar.volume = 1
mock_bar.turnover = 0
mock_bar.open_interest = 0
# Mock AlphaLab session initialization
with patch("vnpy.alpha.lab.AlphaLab"), \
patch("sanguo_data.datareader.read_db_daily") as mock_read, \
patch("sanguo_factor.data_adapter.save_alpha_lab_data"), \
patch("sanguo_factor.data_adapter.convert_bars_to_alpha_df") as mock_convert, \
patch("sanguo_factor.registry.get_factor") as mock_get_factor, \
patch("vnpy.alpha.dataset.AlphaDataset") as mock_dataset_class:
# Setup mock returns
mock_read.return_value = [mock_bar]
mock_get_factor.return_value = {"expression": "ts_mean(close,5)"}
# Create mock AlphaDataset instance
mock_dataset = MagicMock()
mock_dataset.fetch_raw.return_value = pl.DataFrame({
"datetime": [],
"vt_symbol": [],
"ma5": []
})
mock_dataset_class.return_value = mock_dataset
# Create DataFrame for convert_bars_to_alpha_df
# The code will localize the datetime column to Asia/Shanghai, so the mock needs to return a DataFrame with naive datetime
initial_df = pl.DataFrame({
"vt_symbol": ["600000.SSE"],
"datetime": [datetime(2024, 1, 1)],
"open": [1.0],
"high": [1.0],
"low": [1.0],
"close": [1.0],
"volume": [1.0],
"turnover": [0.0],
"open_interest": [0.0]
})
mock_convert.return_value = initial_df
# Create session and load symbols
lab_path = str(tmp_path / "alpha_lab")
session = AlphaLabSession(lab_path=lab_path)
session.load_symbols(["600000"], "2024-01-01", "2024-06-30", cfg=MagicMock())
# Compute factors
df = session.compute_factors(
["ma5"],
("2024-01-01", "2024-04-30"),
("2024-05-01", "2024-05-15"),
("2024-05-16", "2024-06-30")
)
# Verify AlphaDataset was created with Asia/Shanghai-aware periods
from zoneinfo import ZoneInfo
_SH = ZoneInfo("Asia/Shanghai")
# The fix: compute_factors should convert period strings to Asia/Shanghai-aware datetimes
expected_train_period = (datetime(2024, 1, 1, tzinfo=_SH), datetime(2024, 4, 30, tzinfo=_SH))
expected_valid_period = (datetime(2024, 5, 1, tzinfo=_SH), datetime(2024, 5, 15, tzinfo=_SH))
expected_test_period = (datetime(2024, 5, 16, tzinfo=_SH), datetime(2024, 6, 30, tzinfo=_SH))
# Check that AlphaDataset was called once (don't compare DataFrames to avoid polars comparison issues)
assert mock_dataset_class.call_count == 1
call_args = mock_dataset_class.call_args
# Verify the periods are correct
assert call_args[0][1] == expected_train_period
assert call_args[0][2] == expected_valid_period
assert call_args[0][3] == expected_test_period
# Verify add_feature was called for the factor
mock_dataset.add_feature.assert_called_once_with("ma5", "ts_mean(close,5)")
# Verify prepare_data was called
mock_dataset.prepare_data.assert_called_once_with(max_workers=1)
# Verify fetch_raw was called
assert mock_dataset.fetch_raw.called
# Verify return value is a DataFrame
assert isinstance(df, pl.DataFrame)
def test_compute_factors_passes_aware_periods_to_alpha_dataset(tmp_path):
"""Test that compute_factors converts period strings to Asia/Shanghai-aware datetimes.
This ensures vnpy.alpha's to_datetime() preserves timezone awareness,
preventing SchemaError when comparing aware datetime column with naive literals.
Requires polars - runs in container, skips locally.
"""
import pytest
pytest.importorskip("polars")
# vnpy.alpha 仅 vnpy_v4.4.0 源码提供(容器内 sys.path 引用);本地 pip 装的 vnpy 无此子模块
try:
import vnpy.alpha # noqa: F401
except ImportError:
pytest.skip("vnpy.alpha 仅容器内可用(vnpy_v4.4.0 源码)")
from pathlib import Path
from datetime import datetime
from unittest.mock import MagicMock, call
from zoneinfo import ZoneInfo
from sanguo_factor.alpha_lab import AlphaLabSession
import polars as pl
# Create mock bar data
mock_bar = MagicMock()
mock_bar.vt_symbol = "600000.SSE"
mock_bar.datetime = datetime(2024, 1, 1)
mock_bar.open_price = 1.0
mock_bar.high_price = 1.0
mock_bar.low_price = 1.0
mock_bar.close_price = 1.0
mock_bar.volume = 1
mock_bar.turnover = 0
mock_bar.open_interest = 0
# Mock AlphaLab session initialization
with patch("vnpy.alpha.lab.AlphaLab"), \
patch("sanguo_data.datareader.read_db_daily") as mock_read, \
patch("sanguo_factor.data_adapter.save_alpha_lab_data"), \
patch("sanguo_factor.data_adapter.convert_bars_to_alpha_df") as mock_convert, \
patch("sanguo_factor.registry.get_factor") as mock_get_factor, \
patch("vnpy.alpha.dataset.AlphaDataset") as mock_dataset_class:
# Setup mock returns
mock_read.return_value = [mock_bar]
mock_get_factor.return_value = {"expression": "ts_mean(close,5)"}
# Create mock AlphaDataset instance
mock_dataset = MagicMock()
mock_dataset.fetch_raw.return_value = pl.DataFrame({
"datetime": [],
"vt_symbol": [],
"ma5": []
})
mock_dataset_class.return_value = mock_dataset
# Create DataFrame for convert_bars_to_alpha_df
# The code will localize the datetime column to Asia/Shanghai, so the mock needs to return a DataFrame with naive datetime
initial_df = pl.DataFrame({
"vt_symbol": ["600000.SSE"],
"datetime": [datetime(2024, 1, 1)],
"open": [1.0],
"high": [1.0],
"low": [1.0],
"close": [1.0],
"volume": [1.0],
"turnover": [0.0],
"open_interest": [0.0]
})
mock_convert.return_value = initial_df
# Create session and load symbols
lab_path = str(tmp_path / "alpha_lab")
session = AlphaLabSession(lab_path=lab_path)
session.load_symbols(["600000"], "2024-01-01", "2024-06-30", cfg=MagicMock())
# Test with string periods (should be converted to aware datetimes)
session.compute_factors(
["ma5"],
("2024-01-01", "2024-04-30"), # String periods
("2024-05-01", "2024-05-15"),
("2024-05-16", "2024-06-30")
)
# Verify AlphaDataset was called with Asia/Shanghai-aware datetimes
call_args = mock_dataset_class.call_args
_, train_period, valid_period, test_period = call_args[0]
# Check that all period boundaries are datetime objects with Asia/Shanghai timezone
_SH = ZoneInfo("Asia/Shanghai")
for period_name, period in [("train", train_period), ("valid", valid_period), ("test", test_period)]:
start, end = period
assert isinstance(start, datetime), f"{period_name} period start should be datetime object, got {type(start)}"
assert isinstance(end, datetime), f"{period_name} period end should be datetime object, got {type(end)}"
assert start.tzinfo == _SH, f"{period_name} period start should be Asia/Shanghai-aware, got {start.tzinfo}"
assert end.tzinfo == _SH, f"{period_name} period end should be Asia/Shanghai-aware, got {end.tzinfo}"
def test_loaded_bars_cache_eviction(tmp_path):
"""Test that _loaded_bars cache evicts oldest entries when exceeding cap."""
from datetime import datetime
from unittest.mock import MagicMock
from sanguo_factor.alpha_lab import AlphaLabSession, _MAX_CACHED_SYMBOLS
from collections import OrderedDict
# Create mock bar data
def create_mock_bar(symbol: str):
mock_bar = MagicMock()
mock_bar.vt_symbol = symbol
mock_bar.datetime = datetime(2024, 1, 1)
mock_bar.open_price = 1.0
mock_bar.high_price = 1.0
mock_bar.low_price = 1.0
mock_bar.close_price = 1.0
mock_bar.volume = 1
mock_bar.turnover = 0
mock_bar.open_interest = 0
return [mock_bar]
# Create a mock session and manually test cache eviction logic
lab_path = str(tmp_path / "alpha_lab")
# Manually initialize the session to avoid vnpy.alpha import
import sanguo_factor.alpha_lab as alpha_lab_module
original_init = AlphaLabSession.__init__
def mock_init(self, lab_path):
self.lab_path = lab_path
self._loaded_symbols = []
self._loaded_bars = OrderedDict()
# Temporarily replace __init__ and test the cache logic
AlphaLabSession.__init__ = mock_init
try:
session = AlphaLabSession(lab_path=lab_path)
# Directly simulate the cache eviction logic from load_symbols
symbols_to_load = [f"60000{i}.SSE" for i in range(_MAX_CACHED_SYMBOLS + 10)]
for i, symbol in enumerate(symbols_to_load):
# Simulate adding symbol to cache (from load_symbols logic)
bars = create_mock_bar(symbol)
# Add symbol to _loaded_symbols
if symbol not in session._loaded_symbols:
session._loaded_symbols.append(symbol)
# Add or update symbol in cache (LRU logic)
if symbol in session._loaded_bars:
del session._loaded_bars[symbol]
session._loaded_bars[symbol] = bars
# Enforce cache cap - evict oldest symbol if exceeded
while len(session._loaded_bars) > _MAX_CACHED_SYMBOLS:
oldest_symbol = next(iter(session._loaded_bars))
del session._loaded_bars[oldest_symbol]
if oldest_symbol in session._loaded_symbols:
session._loaded_symbols.remove(oldest_symbol)
# Check that cache size never exceeds cap
assert len(session._loaded_bars) <= _MAX_CACHED_SYMBOLS, \
f"Cache exceeded cap at iteration {i}: {len(session._loaded_bars)} > {_MAX_CACHED_SYMBOLS}"
# Final check: cache should be exactly at cap
assert len(session._loaded_bars) == _MAX_CACHED_SYMBOLS
# Verify that the oldest symbols were evicted (first loaded symbols should be gone)
oldest_symbols = symbols_to_load[:10] # First 10 symbols should be evicted
for symbol in oldest_symbols:
assert symbol not in session._loaded_bars, f"Oldest symbol {symbol} should have been evicted"
# Verify that the newest symbols are still in cache
newest_symbols = symbols_to_load[-10:] # Last 10 symbols should be present
for symbol in newest_symbols:
assert symbol in session._loaded_bars, f"Newest symbol {symbol} should be in cache"
finally:
# Restore original __init__
AlphaLabSession.__init__ = original_init