fa7237b996
- analyzer.py: 提取 IC 值到 ic_summary (mean/std/icir/t_stat),periods 提参 (默认 1,5,10) - alpha_lab.py: _loaded_bars 缓存 LRU 上限 (_MAX_CACHED_SYMBOLS=50) - runner.py: 统一阶段文案 (参数优化中/因子分析中),worker 类型标注,_wait_future 文档 - pool.py: submit_work 添加 task_id debug 日志 Co-Authored-By: Claude <noreply@anthropic.com>
217 lines
8.0 KiB
Python
217 lines
8.0 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.
|
|
|
|
Requires polars - runs in container, skips locally.
|
|
"""
|
|
import pytest
|
|
pytest.importorskip("polars")
|
|
|
|
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)
|
|
|
|
|
|
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
|