diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index adbba821bf3e..cb49490ad74b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -470,6 +470,7 @@ def generate_feedback_box(): from litellm.proxy.route_llm_request import route_request from litellm.proxy.search_endpoints.endpoints import router as search_router from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager +from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.spend_tracking.spend_management_endpoints import ( router as spend_management_router, ) @@ -2264,111 +2265,133 @@ async def increment_spend_counters( budget_reservation["finalized"] = True return - if token is not None: - # token arrives pre-hashed from metadata["user_api_key"] (auth flow + cost: float = response_cost + + async def _key_scope(key_token: str) -> None: + # key_token arrives pre-hashed from metadata["user_api_key"] (auth flow # hashes raw "sk-..." keys before they reach the callback). The # startswith("sk-") check is a safety net matching update_cache — # if a raw key somehow arrives, hash it; otherwise use as-is to # avoid double-hashing (budget checks read valid_token.token which # is single-hashed). - hashed_token = hash_token(token=token) if isinstance(token, str) and token.startswith("sk-") else token + hashed_token = ( + hash_token(token=key_token) if isinstance(key_token, str) and key_token.startswith("sk-") else key_token + ) key_counter_key = f"spend:key:{hashed_token}" if key_counter_key not in reserved_counter_keys: await _init_and_increment_spend_counter( counter_key=key_counter_key, source_cache_key=hashed_token, - increment=response_cost, + increment=cost, ) - # Increment per-window budget counters for multi-budget keys key_obj = await user_api_key_cache.async_get_cache(key=hashed_token) - if key_obj is not None: - key_budget_limits = getattr(key_obj, "budget_limits", None) or ( - key_obj.get("budget_limits") if isinstance(key_obj, dict) else None - ) - if isinstance(key_budget_limits, str): - key_budget_limits = json.loads(key_budget_limits) - if isinstance(key_budget_limits, list): - for window in key_budget_limits: - duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration - key_window_counter = f"spend:key:{hashed_token}:window:{duration}" - if key_window_counter not in reserved_counter_keys: - from litellm.proxy.spend_tracking.budget_reservation import ( - get_budget_window_start, - ) - - await _init_and_increment_window_spend_counter( - counter_key=key_window_counter, - entity_type="Key", - entity_id=hashed_token, - window_start=get_budget_window_start(window), - increment=response_cost, - ) + if key_obj is None: + return + key_budget_limits = getattr(key_obj, "budget_limits", None) or ( + key_obj.get("budget_limits") if isinstance(key_obj, dict) else None + ) + if isinstance(key_budget_limits, str): + key_budget_limits = json.loads(key_budget_limits) + if not isinstance(key_budget_limits, list): + return + for window in key_budget_limits: + duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration + key_window_counter = f"spend:key:{hashed_token}:window:{duration}" + if key_window_counter not in reserved_counter_keys: + await _init_and_increment_window_spend_counter( + counter_key=key_window_counter, + entity_type="Key", + entity_id=hashed_token, + window_start=get_budget_window_start(window), + increment=cost, + ) - if team_id is not None: - team_counter_key = f"spend:team:{team_id}" + async def _team_scope(scope_team_id: str) -> None: + team_counter_key = f"spend:team:{scope_team_id}" if team_counter_key not in reserved_counter_keys: await _init_and_increment_spend_counter( counter_key=team_counter_key, - source_cache_key=f"team_id:{team_id}", - increment=response_cost, - ) - - # Increment per-window budget counters for multi-budget teams - team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{team_id}") - if team_obj is not None: - team_budget_limits = getattr(team_obj, "budget_limits", None) or ( - team_obj.get("budget_limits") if isinstance(team_obj, dict) else None - ) - if isinstance(team_budget_limits, str): - team_budget_limits = json.loads(team_budget_limits) - if isinstance(team_budget_limits, list): - for window in team_budget_limits: - duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration - team_window_counter = f"spend:team:{team_id}:window:{duration}" - if team_window_counter not in reserved_counter_keys: - from litellm.proxy.spend_tracking.budget_reservation import ( - get_budget_window_start, - ) - - await _init_and_increment_window_spend_counter( - counter_key=team_window_counter, - entity_type="Team", - entity_id=team_id, - window_start=get_budget_window_start(window), - increment=response_cost, - ) - - if user_id is not None and team_id is not None: - team_member_counter_key = f"spend:team_member:{user_id}:{team_id}" - if team_member_counter_key not in reserved_counter_keys: - await _init_and_increment_spend_counter( - counter_key=team_member_counter_key, - source_cache_key=f"team_membership:{user_id}:{team_id}", - increment=response_cost, + source_cache_key=f"team_id:{scope_team_id}", + increment=cost, ) - if user_id is not None: - user_counter_key = f"spend:user:{user_id}" - if user_counter_key not in reserved_counter_keys: - await _init_and_increment_spend_counter( - counter_key=user_counter_key, - source_cache_key=user_id, - increment=response_cost, - ) + team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}") + if team_obj is None: + return + team_budget_limits = getattr(team_obj, "budget_limits", None) or ( + team_obj.get("budget_limits") if isinstance(team_obj, dict) else None + ) + if isinstance(team_budget_limits, str): + team_budget_limits = json.loads(team_budget_limits) + if not isinstance(team_budget_limits, list): + return + for window in team_budget_limits: + duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration + team_window_counter = f"spend:team:{scope_team_id}:window:{duration}" + if team_window_counter not in reserved_counter_keys: + await _init_and_increment_window_spend_counter( + counter_key=team_window_counter, + entity_type="Team", + entity_id=scope_team_id, + window_start=get_budget_window_start(window), + increment=cost, + ) - await _increment_end_user_and_tag_spend_counters( - end_user_id=end_user_id, - tags=tags, - response_cost=response_cost, - reserved_counter_keys=reserved_counter_keys, - ) + async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None: + team_member_counter_key = f"spend:team_member:{scope_user_id}:{scope_team_id}" + if team_member_counter_key in reserved_counter_keys: + return + await _init_and_increment_spend_counter( + counter_key=team_member_counter_key, + source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}", + increment=cost, + ) - await _increment_org_spend_counter( - org_id=org_id, - response_cost=response_cost, - reserved_counter_keys=reserved_counter_keys, + async def _user_scope(scope_user_id: str) -> None: + user_counter_key = f"spend:user:{scope_user_id}" + if user_counter_key in reserved_counter_keys: + return + await _init_and_increment_spend_counter( + counter_key=user_counter_key, + source_cache_key=scope_user_id, + increment=cost, + ) + + scope_coros = tuple( + coro + for coro in ( + _key_scope(token) if token is not None else None, + _team_scope(team_id) if team_id is not None else None, + _team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None, + _user_scope(user_id) if user_id is not None else None, + _increment_end_user_and_tag_spend_counters( + end_user_id=end_user_id, + tags=tags, + response_cost=cost, + reserved_counter_keys=reserved_counter_keys, + ) + if end_user_id is not None or tags is not None + else None, + _increment_org_spend_counter( + org_id=org_id, + response_cost=cost, + reserved_counter_keys=reserved_counter_keys, + ) + if org_id is not None + else None, + ) + if coro is not None ) + + # return_exceptions so a failing scope does not leave its siblings running + # as orphaned tasks that race the caller's reservation-counter invalidation; + # all scopes settle, then the first error propagates as before. + scope_results = await asyncio.gather(*scope_coros, return_exceptions=True) + scope_errors = [r for r in scope_results if isinstance(r, BaseException)] + if scope_errors: + raise scope_errors[0] + if budget_reservation is not None: budget_reservation["finalized"] = True diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index a839d82984cf..51980342a1de 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -20,6 +20,7 @@ from __future__ import annotations +import asyncio from datetime import datetime from unittest.mock import AsyncMock, MagicMock @@ -420,6 +421,191 @@ async def _fake_coalesced(**kwargs): } +class _ConcurrencyProbe: + """Stand-in for redis_cache.async_increment that pins concurrency. + + Each call registers itself as in-flight and blocks on ``release`` until the + test lets it proceed. ``all_arrived`` fires once ``expected`` distinct scope + increments are simultaneously suspended here, which can only happen if the + per-scope increments are gathered rather than awaited one after another. + """ + + def __init__(self, expected_concurrency: int): + self.expected = expected_concurrency + self.in_flight = 0 + self.max_in_flight = 0 + self.all_arrived = asyncio.Event() + self.release = asyncio.Event() + self.values: dict[str, float] = {} + + async def async_increment(self, *, key, value, refresh_ttl=True): + self.in_flight += 1 + self.max_in_flight = max(self.max_in_flight, self.in_flight) + if self.in_flight >= self.expected: + self.all_arrived.set() + if not self.release.is_set(): + await self.release.wait() + self.in_flight -= 1 + self.values[key] = self.values.get(key, 0.0) + value + return self.values[key] + + +@pytest.mark.asyncio +async def test_increment_spend_counters_runs_scopes_concurrently(monkeypatch): + """The six independent scopes (key, team, team_member, user, end_user+tags, + org) must be incremented concurrently. The probe only fires once all six are + suspended in async_increment at the same time, which is impossible if the + awaits are chained sequentially.""" + probe = _ConcurrencyProbe(expected_concurrency=6) + fake_cache = _make_spend_counter_cache(redis_get_value=None) + fake_cache.redis_cache.async_increment = probe.async_increment + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + task = asyncio.create_task( + ps.increment_spend_counters( + token="hashed-tok", + team_id="t1", + user_id="u1", + org_id="org1", + end_user_id="eu1", + tags=["a", "b"], + response_cost=5.0, + ) + ) + + try: + await asyncio.wait_for(probe.all_arrived.wait(), timeout=2.0) + except asyncio.TimeoutError: + probe.release.set() + await task + pytest.fail( + "scope increments did not run concurrently; sequential awaits " + f"detected (peak in-flight was {probe.max_in_flight}, expected 6)" + ) + + assert probe.in_flight == 6 + assert probe.max_in_flight == 6 + probe.release.set() + await task + + assert probe.values == { + "spend:key:hashed-tok": 5.0, + "spend:team:t1": 5.0, + "spend:team_member:u1:t1": 5.0, + "spend:user:u1": 5.0, + "spend:end_user:eu1": 5.0, + "spend:tag:a": 5.0, + "spend:tag:b": 5.0, + "spend:org:org1": 5.0, + } + + +@pytest.mark.asyncio +async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch): + """Counters already reserved by a budget reservation are skipped, every + other scope is still incremented exactly once, and the reservation is + finalized after the gathered work completes.""" + import litellm.proxy.spend_tracking.budget_reservation as br + + reserved = {"spend:key:hashed-tok", "spend:org:org1"} + monkeypatch.setattr( + br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved)) + ) + monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock()) + + recorded: dict[str, float] = {} + + async def _record_increment(*, key, value, refresh_ttl=True): + recorded[key] = recorded.get(key, 0.0) + value + return recorded[key] + + fake_cache = _make_spend_counter_cache(redis_get_value=None) + fake_cache.redis_cache.async_increment = _record_increment + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + reservation = {"finalized": False} + await ps.increment_spend_counters( + token="hashed-tok", + team_id="t1", + user_id="u1", + org_id="org1", + end_user_id="eu1", + tags=["a"], + response_cost=5.0, + budget_reservation=reservation, + ) + + assert reservation["finalized"] is True + assert recorded == { + "spend:team:t1": 5.0, + "spend:team_member:u1:t1": 5.0, + "spend:user:u1": 5.0, + "spend:end_user:eu1": 5.0, + "spend:tag:a": 5.0, + } + + +@pytest.mark.asyncio +async def test_increment_spend_counters_failing_scope_propagates_after_siblings_settle( + monkeypatch, +): + """A failure in one scope must propagate to the caller (so it can invalidate + reserved counters) while every other scope still settles rather than being + left as an orphaned background task, and the reservation is not finalized.""" + recorded: dict[str, float] = {} + + async def _increment(*, key, value, refresh_ttl=True): + if key == "spend:team:t1": + raise RuntimeError("redis increment failed") + recorded[key] = recorded.get(key, 0.0) + value + return recorded[key] + + fake_cache = _make_spend_counter_cache(redis_get_value=None) + fake_cache.redis_cache.async_increment = _increment + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + reservation = {"finalized": False} + with pytest.raises(RuntimeError, match="redis increment failed"): + await ps.increment_spend_counters( + token="hashed-tok", + team_id="t1", + user_id="u1", + org_id="org1", + end_user_id="eu1", + tags=["a"], + response_cost=5.0, + budget_reservation=reservation, + ) + + assert reservation["finalized"] is False + assert recorded == { + "spend:key:hashed-tok": 5.0, + "spend:team_member:u1:t1": 5.0, + "spend:user:u1": 5.0, + "spend:end_user:eu1": 5.0, + "spend:tag:a": 5.0, + "spend:org:org1": 5.0, + } + + @pytest.mark.asyncio async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation( monkeypatch,