Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand Down
23 changes: 23 additions & 0 deletions tests/test_offline_network_policy.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Offline test-lane network policy tests."""

import asyncio
import socket

import pytest
Expand Down Expand Up @@ -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())
19 changes: 9 additions & 10 deletions tests/test_yahoo_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
12 changes: 10 additions & 2 deletions tests/test_yahoo_provider_edge_cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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")
Expand Down
Loading