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
4 changes: 2 additions & 2 deletions alloc/lib/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,8 @@
# Maps method names to the cache-type string used by settings.get_cache_ttl().
# Add entries here to enable caching for additional Polygon methods.
CACHE_MAP: dict[str, str] = {
"get_aggs": "historical_data",
"get_ticker_details": "ticker_details",
"get_aggregate_bars": "historical_data",
"get_snapshot": "ticker_details",
}


Expand Down
24 changes: 12 additions & 12 deletions alloc/models/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,36 +199,36 @@ def get_multi_asset_data(
formatted = ticker.upper()
try:
# --- hourly ---
bars = client.get_aggs(
ticker=formatted,
bars = client.get_aggregate_bars(
symbol=formatted,
multiplier=1,
timespan="hour",
from_=hourly_start.strftime("%Y-%m-%d"),
to=end_date.strftime("%Y-%m-%d"),
from_date=hourly_start.strftime("%Y-%m-%d"),
to_date=end_date.strftime("%Y-%m-%d"),
limit=5000,
)
if bars:
result[ticker]["hourly"] = [float(b.close) for b in bars]

# --- daily ---
bars = client.get_aggs(
ticker=formatted,
bars = client.get_aggregate_bars(
symbol=formatted,
multiplier=1,
timespan="day",
from_=daily_start.strftime("%Y-%m-%d"),
to=end_date.strftime("%Y-%m-%d"),
from_date=daily_start.strftime("%Y-%m-%d"),
to_date=end_date.strftime("%Y-%m-%d"),
limit=5000,
)
if bars:
result[ticker]["daily"] = [float(b.close) for b in bars]

# --- weekly ---
bars = client.get_aggs(
ticker=formatted,
bars = client.get_aggregate_bars(
symbol=formatted,
multiplier=1,
timespan="week",
from_=weekly_start.strftime("%Y-%m-%d"),
to=end_date.strftime("%Y-%m-%d"),
from_date=weekly_start.strftime("%Y-%m-%d"),
to_date=end_date.strftime("%Y-%m-%d"),
limit=5000,
)
if bars:
Expand Down
64 changes: 32 additions & 32 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,8 +29,8 @@ def _mock_method(name: str, return_value) -> MagicMock:
def mock_rest_client() -> MagicMock:
"""Return a MagicMock standing in for polygon.StocksClient."""
client = MagicMock()
client.get_aggs = _mock_method("get_aggs", {"results": [1, 2, 3]})
client.get_ticker_details = _mock_method("get_ticker_details", {"ticker": "AAPL"})
client.get_aggregate_bars = _mock_method("get_aggregate_bars", {"results": [1, 2, 3]})
client.get_snapshot = _mock_method("get_snapshot", {"ticker": "AAPL"})
client.get_news = _mock_method("get_news", [{"title": "Breaking"}])
return client

Expand Down Expand Up @@ -84,7 +84,7 @@ def test_uses_default_cache_map(
def test_accepts_custom_cache_map(
self, mock_rest_client: MagicMock, cache: DiskCache
) -> None:
custom = {"get_aggs": "latest_prices"}
custom = {"get_aggregate_bars": "latest_prices"}
with patch("alloc.lib.client.StocksClient", return_value=mock_rest_client):
c = PolygonClient(api_key="k", cache=cache, cache_map=custom)
assert c._cache_map == custom
Expand All @@ -97,31 +97,31 @@ def test_accepts_custom_cache_map(
class TestCaching:
"""Tests that cached methods actually cache."""

def test_get_aggs_is_cached(self, client: PolygonClient) -> None:
"""Second call to get_aggs should hit cache, not upstream."""
client.get_aggs("AAPL", 1, "day", "2024-01-01", "2024-01-31")
client.get_aggs("AAPL", 1, "day", "2024-01-01", "2024-01-31")
# Underlying client's get_aggs called only once
assert client._client.get_aggs.call_count == 1
def test_get_aggregate_bars_is_cached(self, client: PolygonClient) -> None:
"""Second call to get_aggregate_bars should hit cache, not upstream."""
client.get_aggregate_bars("AAPL", 1, "day", "2024-01-01", "2024-01-31")
client.get_aggregate_bars("AAPL", 1, "day", "2024-01-01", "2024-01-31")
# Underlying client's get_aggregate_bars called only once
assert client._client.get_aggregate_bars.call_count == 1

def test_get_ticker_details_is_cached(self, client: PolygonClient) -> None:
client.get_ticker_details("AAPL")
client.get_ticker_details("AAPL")
assert client._client.get_ticker_details.call_count == 1
def test_get_snapshot_is_cached(self, client: PolygonClient) -> None:
client.get_snapshot("AAPL")
client.get_snapshot("AAPL")
assert client._client.get_snapshot.call_count == 1

def test_different_args_miss_cache(self, client: PolygonClient) -> None:
client.get_aggs("AAPL", 1, "day", "2024-01-01", "2024-01-31")
client.get_aggs("MSFT", 1, "day", "2024-01-01", "2024-01-31")
assert client._client.get_aggs.call_count == 2
client.get_aggregate_bars("AAPL", 1, "day", "2024-01-01", "2024-01-31")
client.get_aggregate_bars("MSFT", 1, "day", "2024-01-01", "2024-01-31")
assert client._client.get_aggregate_bars.call_count == 2

def test_disabled_cache_skips_caching(
self, mock_rest_client: MagicMock, disabled_cache: DiskCache
) -> None:
with patch("alloc.lib.client.StocksClient", return_value=mock_rest_client):
c = PolygonClient(api_key="k", cache=disabled_cache)
c.get_aggs("AAPL", 1, "day", "2024-01-01", "2024-01-31")
c.get_aggs("AAPL", 1, "day", "2024-01-01", "2024-01-31")
assert mock_rest_client.get_aggs.call_count == 2
c.get_aggregate_bars("AAPL", 1, "day", "2024-01-01", "2024-01-31")
c.get_aggregate_bars("AAPL", 1, "day", "2024-01-01", "2024-01-31")
assert mock_rest_client.get_aggregate_bars.call_count == 2


# =====================================================================
Expand Down Expand Up @@ -152,11 +152,11 @@ def _noop(*a, **k):
return None

bare_client = SimpleNamespace()
_noop.__name__ = "get_aggs"
bare_client.get_aggs = _noop
_noop.__name__ = "get_aggregate_bars"
bare_client.get_aggregate_bars = _noop
_noop2 = _noop
_noop2.__name__ = "get_ticker_details"
bare_client.get_ticker_details = _noop2
_noop2.__name__ = "get_snapshot"
bare_client.get_snapshot = _noop2

with patch("alloc.lib.client.StocksClient", return_value=bare_client):
c = PolygonClient(api_key="k", cache=cache)
Expand All @@ -175,16 +175,16 @@ def test_cache_valid_false_not_cached(
self, mock_rest_client: MagicMock, cache: DiskCache
) -> None:
"""A result with __cache_valid__: False should not be cached."""
mock_rest_client.get_ticker_details.return_value = {
mock_rest_client.get_snapshot.return_value = {
"ticker": "AAPL",
"__cache_valid__": False,
}
with patch("alloc.lib.client.StocksClient", return_value=mock_rest_client):
c = PolygonClient(api_key="k", cache=cache)
r1 = c.get_ticker_details("AAPL")
r2 = c.get_ticker_details("AAPL")
r1 = c.get_snapshot("AAPL")
r2 = c.get_snapshot("AAPL")
# Called twice — never cached
assert mock_rest_client.get_ticker_details.call_count == 2
assert mock_rest_client.get_snapshot.call_count == 2
# Sentinel stripped from return value
assert "__cache_valid__" not in r1
assert r1 == {"ticker": "AAPL"}
Expand All @@ -202,13 +202,13 @@ def test_missing_method_on_restclient_is_skipped(
) -> None:
"""If a method in cache_map doesn't exist on StocksClient, skip it."""
custom_map = {
"get_aggs": "historical_data",
"get_aggregate_bars": "historical_data",
"nonexistent_method_xyz": "latest_prices",
}
with patch("alloc.lib.client.StocksClient", return_value=mock_rest_client):
c = PolygonClient(api_key="k", cache=cache, cache_map=custom_map)
# get_aggs should be wrapped
assert hasattr(c, "get_aggs")
# get_aggregate_bars should be wrapped
assert hasattr(c, "get_aggregate_bars")
# nonexistent_method_xyz should not cause an error
# and should fall through to __getattr__ (which will raise)

Expand All @@ -221,8 +221,8 @@ class TestCacheMap:
"""Tests for the CACHE_MAP module constant."""

def test_cache_map_has_expected_keys(self) -> None:
assert "get_aggs" in CACHE_MAP
assert "get_ticker_details" in CACHE_MAP
assert "get_aggregate_bars" in CACHE_MAP
assert "get_snapshot" in CACHE_MAP

def test_cache_map_values_are_known_types(self) -> None:
for v in CACHE_MAP.values():
Expand Down
50 changes: 25 additions & 25 deletions tests/test_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,15 +76,15 @@ class TestGetMultiAssetData:
"""Tests for get_multi_asset_data."""

def test_returns_dict_with_ticker_keys(self, mock_client: MagicMock, end_date: datetime) -> None:
mock_client.get_aggs.return_value = [_make_bar(100.0)]
mock_client.get_aggregate_bars.return_value = [_make_bar(100.0)]
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
)
assert "AAPL" in result

def test_returns_dict_with_frequency_keys(self, mock_client: MagicMock, end_date: datetime) -> None:
mock_client.get_aggs.return_value = [_make_bar(100.0)]
mock_client.get_aggregate_bars.return_value = [_make_bar(100.0)]
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
Expand All @@ -94,7 +94,7 @@ def test_returns_dict_with_frequency_keys(self, mock_client: MagicMock, end_date

def test_hourly_prices_extracted_correctly(self, mock_client: MagicMock, end_date: datetime) -> None:
bars = [_make_bar(100.0), _make_bar(101.0), _make_bar(102.0)]
mock_client.get_aggs.return_value = bars
mock_client.get_aggregate_bars.return_value = bars
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
Expand All @@ -103,7 +103,7 @@ def test_hourly_prices_extracted_correctly(self, mock_client: MagicMock, end_dat

def test_daily_prices_extracted_correctly(self, mock_client: MagicMock, end_date: datetime) -> None:
bars = [_make_bar(200.0), _make_bar(205.0)]
mock_client.get_aggs.return_value = bars
mock_client.get_aggregate_bars.return_value = bars
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
Expand All @@ -112,51 +112,51 @@ def test_daily_prices_extracted_correctly(self, mock_client: MagicMock, end_date

def test_weekly_prices_extracted_correctly(self, mock_client: MagicMock, end_date: datetime) -> None:
bars = [_make_bar(300.0)]
mock_client.get_aggs.return_value = bars
mock_client.get_aggregate_bars.return_value = bars
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
)
assert result["AAPL"]["weekly"] == [300.0]

def test_calls_get_aggs_with_correct_timespan_hourly(
def test_calls_get_aggregate_bars_with_correct_timespan_hourly(
self, mock_client: MagicMock, end_date: datetime
) -> None:
mock_client.get_aggs.return_value = []
mock_client.get_aggregate_bars.return_value = []
get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=7, daily_days=30, weekly_weeks=12,
)
# First call should be hourly
call_args = mock_client.get_aggs.call_args_list[0]
call_args = mock_client.get_aggregate_bars.call_args_list[0]
assert call_args[1]["timespan"] == "hour"

def test_calls_get_aggs_with_correct_timespan_daily(
def test_calls_get_aggregate_bars_with_correct_timespan_daily(
self, mock_client: MagicMock, end_date: datetime
) -> None:
mock_client.get_aggs.return_value = []
mock_client.get_aggregate_bars.return_value = []
get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=7, daily_days=30, weekly_weeks=12,
)
# Second call should be daily
call_args = mock_client.get_aggs.call_args_list[1]
call_args = mock_client.get_aggregate_bars.call_args_list[1]
assert call_args[1]["timespan"] == "day"

def test_calls_get_aggs_with_correct_timespan_weekly(
def test_calls_get_aggregate_bars_with_correct_timespan_weekly(
self, mock_client: MagicMock, end_date: datetime
) -> None:
mock_client.get_aggs.return_value = []
mock_client.get_aggregate_bars.return_value = []
get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=7, daily_days=30, weekly_weeks=12,
)
# Third call should be week
call_args = mock_client.get_aggs.call_args_list[2]
call_args = mock_client.get_aggregate_bars.call_args_list[2]
assert call_args[1]["timespan"] == "week"

def test_handles_empty_bars(self, mock_client: MagicMock, end_date: datetime) -> None:
mock_client.get_aggs.return_value = []
mock_client.get_aggregate_bars.return_value = []
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
Expand All @@ -166,7 +166,7 @@ def test_handles_empty_bars(self, mock_client: MagicMock, end_date: datetime) ->
assert result["AAPL"]["weekly"] == []

def test_handles_multiple_tickers(self, mock_client: MagicMock, end_date: datetime) -> None:
mock_client.get_aggs.return_value = [_make_bar(100.0)]
mock_client.get_aggregate_bars.return_value = [_make_bar(100.0)]
result = get_multi_asset_data(
tickers=["AAPL", "MSFT"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
Expand All @@ -175,17 +175,17 @@ def test_handles_multiple_tickers(self, mock_client: MagicMock, end_date: dateti
assert "MSFT" in result

def test_ticker_uppercased(self, mock_client: MagicMock, end_date: datetime) -> None:
mock_client.get_aggs.return_value = []
mock_client.get_aggregate_bars.return_value = []
get_multi_asset_data(
tickers=["aapl"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
)
# First call's ticker arg should be uppercased
call_args = mock_client.get_aggs.call_args_list[0]
assert call_args[1]["ticker"] == "AAPL"
call_args = mock_client.get_aggregate_bars.call_args_list[0]
assert call_args[1]["symbol"] == "AAPL"

def test_defaults_end_date_to_today(self, mock_client: MagicMock) -> None:
mock_client.get_aggs.return_value = []
mock_client.get_aggregate_bars.return_value = []
with patch("alloc.models.data.datetime") as mock_dt:
mock_dt.today.return_value = datetime(2024, 1, 1)
mock_dt.side_effect = lambda *a, **k: datetime(*a, **k)
Expand All @@ -194,18 +194,18 @@ def test_defaults_end_date_to_today(self, mock_client: MagicMock) -> None:
hourly_days=1, daily_days=1, weekly_weeks=1,
)
# Should have been called
assert mock_client.get_aggs.called
assert mock_client.get_aggregate_bars.called

def test_error_on_ticker_does_not_crash_others(
self, mock_client: MagicMock, end_date: datetime
) -> None:
"""If one ticker raises, others should still be processed."""
def side_effect(*args, **kwargs):
if kwargs.get("ticker") == "BAD":
if kwargs.get("symbol") == "BAD":
raise RuntimeError("API error")
return [_make_bar(100.0)]

mock_client.get_aggs.side_effect = side_effect
mock_client.get_aggregate_bars.side_effect = side_effect
result = get_multi_asset_data(
tickers=["GOOD", "BAD"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
Expand All @@ -217,13 +217,13 @@ def side_effect(*args, **kwargs):

def test_uses_injected_client_not_singleton(self, mock_client: MagicMock, end_date: datetime) -> None:
"""Verify the function uses the passed client, not a module-level singleton."""
mock_client.get_aggs.return_value = [_make_bar(42.0)]
mock_client.get_aggregate_bars.return_value = [_make_bar(42.0)]
result = get_multi_asset_data(
tickers=["AAPL"], client=mock_client, end_date=end_date,
hourly_days=1, daily_days=1, weekly_weeks=1,
)
assert result["AAPL"]["hourly"] == [42.0]
assert mock_client.get_aggs.called
assert mock_client.get_aggregate_bars.called


# =====================================================================
Expand Down
Loading