diff --git a/sanguo_factor/alpha_lab.py b/sanguo_factor/alpha_lab.py index 127050a..4e7abdb 100644 --- a/sanguo_factor/alpha_lab.py +++ b/sanguo_factor/alpha_lab.py @@ -8,23 +8,25 @@ if _VNPY_SRC not in sys.path: class AlphaLabSession: """Session manager for vnpy.alpha AlphaLab operations.""" - + def __init__(self, lab_path: str): """ Initialize AlphaLab session. - + Args: lab_path: Path to AlphaLab directory """ from vnpy.alpha.lab import AlphaLab - + self.lab_path = lab_path self.lab = AlphaLab(lab_path) + self._loaded_symbols: list[str] = [] + self._loaded_bars: dict[str, list] = {} def load_symbols(self, symbols: list[str], start: str, end: str, cfg) -> None: """ Load symbol data from database and save to AlphaLab. - + Args: symbols: List of vt_symbols to load start: Start date (YYYY-MM-DD) @@ -33,8 +35,55 @@ class AlphaLabSession: """ from sanguo_data.datareader import read_db_daily from .data_adapter import save_alpha_lab_data - + for symbol in symbols: bars = read_db_daily(symbol, start, end, cfg) if bars: save_alpha_lab_data(bars, self.lab_path) + # Cache bars for compute_factors + if symbol not in self._loaded_symbols: + self._loaded_symbols.append(symbol) + self._loaded_bars[symbol] = bars + + def compute_factors(self, factor_names: list[str], train_period: tuple, valid_period: tuple, test_period: tuple): + """ + Compute factors using cached bars and vnpy.alpha AlphaDataset. + + Args: + factor_names: List of factor names to compute + train_period: Training period tuple (start, end) + valid_period: Validation period tuple (start, end) + test_period: Test period tuple (start, end) + + Returns: + polars DataFrame with computed factors for test period + """ + # Lazy imports to avoid ImportError on local Python 3.14 without polars/vnpy.alpha + import polars as pl + from vnpy.alpha.dataset import AlphaDataset, Segment + from .registry import get_factor + from .data_adapter import convert_bars_to_alpha_df + + # Gather all cached bars across loaded symbols + all_bars = [] + for symbol in self._loaded_symbols: + all_bars.extend(self._loaded_bars.get(symbol, [])) + + # Convert bars to AlphaLab DataFrame format + df = convert_bars_to_alpha_df(all_bars) + + # Create AlphaDataset with the specified periods + ds = AlphaDataset(df, train_period, valid_period, test_period) + + # Add each factor to the dataset + for name in factor_names: + factor = get_factor(name) + if factor is None: + continue # Skip unknown factors + ds.add_feature(name, factor["expression"]) + + # Prepare data (compute features) + ds.prepare_data(max_workers=1) + + # Return test period data + return ds.fetch_raw(Segment.TEST) diff --git a/tests/factor/test_alpha_lab.py b/tests/factor/test_alpha_lab.py index 3b1e3a8..257a15a 100644 --- a/tests/factor/test_alpha_lab.py +++ b/tests/factor/test_alpha_lab.py @@ -35,3 +35,92 @@ def test_load_symbols_calls_read_db_daily(): 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)