Files
sanguo_vnpy_v2/tests/factor/test_alpha_lab.py
T
claude_dev 8972d1058f feat(factor): alpha_lab compute_factors 完整化
添加 compute_factors 方法到 AlphaLabSession:
- 在 __init__ 添加 _loaded_symbols 和 _loaded_bars 缓存
- load_symbols 现在缓存 bar 数据供 compute_factors 使用
- compute_factors 使用缓存的 bars 调用 AlphaDataset
- 使用懒导入避免本地 Python 3.14 缺少 polars/vnpy.alpha 的 ImportError

测试 (container only):
- test_compute_factors_calls_prepare_and_fetch 验证 AlphaDataset 调用

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-06 19:14:12 +08:00

127 lines
4.6 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_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."""
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
mock_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 = mock_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 correct periods
mock_dataset_class.assert_called_once_with(
mock_df,
("2024-01-01", "2024-04-30"),
("2024-05-01", "2024-05-15"),
("2024-05-16", "2024-06-30")
)
# 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)