From db2cc8c53156013662b2321f51dcff64124f91fb Mon Sep 17 00:00:00 2001 From: claude_dev Date: Mon, 6 Jul 2026 22:18:51 +0800 Subject: [PATCH] =?UTF-8?q?fix(factor):=20compute=5Ffactors=20=E6=97=B6?= =?UTF-8?q?=E5=8C=BA=E5=AF=B9=E9=BD=90=EF=BC=88=E8=BE=B9=E7=95=8C=20Asia/S?= =?UTF-8?q?hanghai=20aware=EF=BC=8C=E7=9C=9F=E6=95=B0=E6=8D=AE=E8=B7=91?= =?UTF-8?q?=E9=80=9A=E5=9B=A0=E5=AD=90=E7=AE=A1=E7=BA=BF=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Root cause: vnpy.alpha's to_datetime() creates naive datetimes from strings, causing SchemaError when comparing with timezone-aware DataFrame columns - Fix: Convert period boundaries to Asia/Shanghai-aware datetimes + localize DataFrame datetime column before passing to AlphaDataset - Restore data_adapter.py to fa7237b (removed ineffective tz stripping) - Add test_compute_factors_passes_aware_periods_to_alpha_dataset - Real data verification: 600000.SSE ma5 factor analysis successful - Container tests: 67 passed Co-Authored-By: Claude --- sanguo_factor/alpha_lab.py | 26 +++++++ sanguo_factor/data_adapter.py | 2 +- tests/factor/test_alpha_lab.py | 123 ++++++++++++++++++++++++++++++--- 3 files changed, 140 insertions(+), 11 deletions(-) diff --git a/sanguo_factor/alpha_lab.py b/sanguo_factor/alpha_lab.py index b7ca3a6..e1838d4 100644 --- a/sanguo_factor/alpha_lab.py +++ b/sanguo_factor/alpha_lab.py @@ -82,6 +82,26 @@ class AlphaLabSession: from .registry import get_factor from .data_adapter import convert_bars_to_alpha_df + # Convert period boundaries to Asia/Shanghai-aware datetime objects + # This ensures vnpy.alpha's to_datetime() preserves timezone awareness, + # preventing SchemaError when comparing aware datetime column with naive literals + from datetime import datetime + from zoneinfo import ZoneInfo + _SH = ZoneInfo("Asia/Shanghai") + + def _to_aware_period(period: tuple) -> tuple: + """Convert period boundary strings or naive datetimes to Asia/Shanghai-aware datetimes.""" + start, end = period + def conv(x): + if isinstance(x, datetime): + return x if x.tzinfo else x.replace(tzinfo=_SH) + return datetime.strptime(x, "%Y-%m-%d").replace(tzinfo=_SH) + return (conv(start), conv(end)) + + train_period = _to_aware_period(train_period) + valid_period = _to_aware_period(valid_period) + test_period = _to_aware_period(test_period) + # Gather all cached bars across loaded symbols all_bars = [] for symbol in self._loaded_symbols: @@ -90,6 +110,12 @@ class AlphaLabSession: # Convert bars to AlphaLab DataFrame format df = convert_bars_to_alpha_df(all_bars) + # Localize the datetime column to Asia/Shanghai-aware to match vnpy.alpha's expectations + # This ensures the DataFrame's datetime column has the same timezone as the period boundaries + df = df.with_columns( + pl.col("datetime").dt.replace_time_zone("Asia/Shanghai") + ) + # Create AlphaDataset with the specified periods ds = AlphaDataset(df, train_period, valid_period, test_period) diff --git a/sanguo_factor/data_adapter.py b/sanguo_factor/data_adapter.py index 12d980c..4f721bc 100644 --- a/sanguo_factor/data_adapter.py +++ b/sanguo_factor/data_adapter.py @@ -66,4 +66,4 @@ def save_alpha_lab_data(bars: list[BarData], lab_path: str) -> Path: lab = AlphaLab(lab_path) lab.save_bar_data(bars) - return lab.daily_path + return lab.daily_path \ No newline at end of file diff --git a/tests/factor/test_alpha_lab.py b/tests/factor/test_alpha_lab.py index cd6eb26..f173e33 100644 --- a/tests/factor/test_alpha_lab.py +++ b/tests/factor/test_alpha_lab.py @@ -85,7 +85,8 @@ def test_compute_factors_calls_prepare_and_fetch(tmp_path): mock_dataset_class.return_value = mock_dataset # Create DataFrame for convert_bars_to_alpha_df - mock_df = pl.DataFrame({ + # 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], @@ -96,7 +97,7 @@ def test_compute_factors_calls_prepare_and_fetch(tmp_path): "turnover": [0.0], "open_interest": [0.0] }) - mock_convert.return_value = mock_df + mock_convert.return_value = initial_df # Create session and load symbols lab_path = str(tmp_path / "alpha_lab") @@ -111,13 +112,22 @@ def test_compute_factors_calls_prepare_and_fetch(tmp_path): ("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 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)") @@ -132,6 +142,99 @@ def test_compute_factors_calls_prepare_and_fetch(tmp_path): 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") + + 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 @@ -213,4 +316,4 @@ def test_loaded_bars_cache_eviction(tmp_path): finally: # Restore original __init__ - AlphaLabSession.__init__ = original_init + AlphaLabSession.__init__ = original_init \ No newline at end of file