diff --git a/tests/conftest.py b/tests/conftest.py index 8d5f191..c8a76fb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -110,6 +110,23 @@ def guarded_getaddrinfo(host, *args, **kwargs): monkeypatch.setattr(socket.socket, "connect", guarded_connect) monkeypatch.setattr(socket.socket, "connect_ex", guarded_connect_ex) monkeypatch.setattr(socket, "getaddrinfo", guarded_getaddrinfo) + + # yfinance sends through curl_cffi, whose libcurl opens sockets in C and never + # reaches the patches above. + try: + from curl_cffi import requests as curl_requests + except ImportError: + curl_requests = None + if curl_requests is not None: + + def guarded_curl_request(session, method, url, *args, **kwargs): + reject_external(url) + + async def guarded_async_curl_request(session, method, url, *args, **kwargs): + reject_external(url) + + monkeypatch.setattr(curl_requests.Session, "request", guarded_curl_request) + monkeypatch.setattr(curl_requests.AsyncSession, "request", guarded_async_curl_request) yield if violations and not request.node.get_closest_marker("network_guard_probe"): pytest.fail(f"Offline test attempted external network access: {violations!r}") diff --git a/tests/test_offline_network_policy.py b/tests/test_offline_network_policy.py index 5374007..0d108b8 100644 --- a/tests/test_offline_network_policy.py +++ b/tests/test_offline_network_policy.py @@ -1,5 +1,6 @@ """Offline test-lane network policy tests.""" +import asyncio import socket import pytest @@ -28,3 +29,25 @@ def test_default_lane_allows_loopback_connections() -> None: with socket.create_connection(server.getsockname(), timeout=1): connection, _ = server.accept() connection.close() + + +@pytest.mark.network_guard_probe +def test_default_lane_rejects_curl_cffi_requests() -> None: + """yfinance's transport is libcurl, which never touches Python's socket module.""" + curl_requests = pytest.importorskip("curl_cffi.requests") + + with curl_requests.Session() as session: + with pytest.raises(RuntimeError, match="Offline test attempted"): + session.get("https://query1.finance.yahoo.com/v8/finance/chart/AAPL") + + +@pytest.mark.network_guard_probe +def test_default_lane_rejects_async_curl_cffi_requests() -> None: + curl_requests = pytest.importorskip("curl_cffi.requests") + + async def fetch() -> None: + async with curl_requests.AsyncSession() as session: + await session.get("https://query1.finance.yahoo.com/v8/finance/chart/AAPL") + + with pytest.raises(RuntimeError, match="Offline test attempted"): + asyncio.run(fetch()) diff --git a/tests/test_yahoo_provider.py b/tests/test_yahoo_provider.py index ed77aac..2e43d3f 100644 --- a/tests/test_yahoo_provider.py +++ b/tests/test_yahoo_provider.py @@ -105,20 +105,19 @@ def test_fetch_minute_data(self, mock_download: MagicMock) -> None: @patch("ml4t.data.providers.yahoo.yf.download") def test_empty_response_handling(self, mock_download: MagicMock) -> None: """Test handling of empty responses.""" - # Mock empty response + # Mock empty response, and an empty retry through Ticker.history mock_download.return_value = pd.DataFrame() provider = YahooFinanceProvider() - # Empty response should raise SymbolNotFoundError - from ml4t.data.core.exceptions import SymbolNotFoundError - - with pytest.raises(SymbolNotFoundError): - provider.fetch_ohlcv( - symbol="INVALID", - start="2024-01-01", - end="2024-01-03", - ) + with patch("ml4t.data.providers.yahoo.yf.Ticker") as mock_ticker: + mock_ticker.return_value.history.return_value = pd.DataFrame() + with pytest.raises(SymbolNotFoundError): + provider.fetch_ohlcv( + symbol="INVALID", + start="2024-01-01", + end="2024-01-03", + ) @patch("ml4t.data.providers.yahoo.yf.download") def test_rate_limiting(self, mock_download: MagicMock) -> None: diff --git a/tests/test_yahoo_provider_edge_cases.py b/tests/test_yahoo_provider_edge_cases.py index 8b9c784..1c698c9 100644 --- a/tests/test_yahoo_provider_edge_cases.py +++ b/tests/test_yahoo_provider_edge_cases.py @@ -27,8 +27,12 @@ def provider(self): def test_empty_data_raises_symbol_not_found(self, provider): """Test empty download raises SymbolNotFoundError.""" - with patch("yfinance.download") as mock_download: + with ( + patch("yfinance.download") as mock_download, + patch("ml4t.data.providers.yahoo.yf.Ticker") as mock_ticker, + ): mock_download.return_value = pd.DataFrame() + mock_ticker.return_value.history.return_value = pd.DataFrame() with pytest.raises(SymbolNotFoundError) as exc_info: provider._fetch_and_transform_data( @@ -74,8 +78,12 @@ def test_unknown_frequency_uses_daily_default(self, provider): def test_propagates_symbol_not_found_error(self, provider): """Test SymbolNotFoundError is propagated without wrapping.""" - with patch("yfinance.download") as mock_download: + with ( + patch("yfinance.download") as mock_download, + patch("ml4t.data.providers.yahoo.yf.Ticker") as mock_ticker, + ): mock_download.return_value = pd.DataFrame() # Empty = not found + mock_ticker.return_value.history.return_value = pd.DataFrame() with pytest.raises(SymbolNotFoundError): provider._fetch_and_transform_data("MISSING", "2024-01-01", "2024-01-31", "daily")