From 09503ebb8fb41db6e0f36784d7b5686993808bc2 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Wed, 29 Apr 2026 21:06:30 -0700 Subject: [PATCH] harden budget reservation recovery --- litellm/proxy/_types.py | 2 +- litellm/proxy/auth/user_api_key_auth.py | 2 + .../proxy/hooks/proxy_track_cost_callback.py | 19 ++++++- .../spend_tracking/budget_reservation.py | 12 +++++ .../proxy/auth/test_user_api_key_auth.py | 27 ++++++++++ .../hooks/test_proxy_track_cost_callback.py | 13 +++-- .../proxy/test_budget_reservation.py | 49 +++++++++++++++++++ 7 files changed, 117 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8be73f3427e..05e6135c16d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2567,7 +2567,7 @@ class UserAPIKeyAuth( user_spend: Optional[float] = None user_max_budget: Optional[float] = None request_route: Optional[str] = None - budget_reservation: Optional[Dict[str, Any]] = None + budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) user: Optional[Any] = None # Expanded user object when expand=user is used created_by_user: Optional[Any] = ( None # Expanded created_by user when expand=user is used diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 7e70e7bb3fa..0d60b36d35e 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1894,6 +1894,7 @@ async def _reserve_budget_after_common_checks( proxy_logging_obj: ProxyLogging, skip_budget_checks: bool, ) -> None: + user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: return @@ -1962,6 +1963,7 @@ async def user_api_key_auth( request_data=request_data, custom_litellm_key_header=custom_litellm_key_header, ) + user_api_key_auth_obj.budget_reservation = None ## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ## RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 71ad962d584..b5e659757ad 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -446,7 +446,9 @@ async def _update_database_and_spend_counters( ) except Exception: if budget_reservation is not None: - await _release_budget_reservation(budget_reservation=budget_reservation) + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) raise @@ -461,3 +463,18 @@ async def _release_budget_reservation(budget_reservation: Optional[dict]) -> Non await release_budget_reservation( budget_reservation=budget_reservation, ) + + +async def _invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + ) + + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation, + ) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 130a3e20275..18ef6dff0e6 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -153,6 +153,18 @@ async def release_budget_reservation(budget_reservation: Optional[dict]) -> None ) +async def invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.proxy_server import _invalidate_spend_counter + + for counter_key in get_reserved_counter_keys(budget_reservation=budget_reservation): + await _invalidate_spend_counter(counter_key=counter_key) + + async def _get_budget_counters( valid_token: UserAPIKeyAuth, team_object: Optional[LiteLLM_TeamTable], diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 23b1de881f7..bcd56cf872b 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -23,6 +23,7 @@ from litellm.proxy._types import ( from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( + _reserve_budget_after_common_checks, _run_centralized_common_checks, _run_post_custom_auth_checks, get_api_key, @@ -47,6 +48,32 @@ def test_get_api_key(): ) == (api_key, passed_in_key) +@pytest.mark.asyncio +async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): + user_api_key_auth_obj = UserAPIKeyAuth( + token="test_token", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_token"}], + }, + ) + + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "free-model"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=True, + ) + + assert user_api_key_auth_obj.budget_reservation is None + + @pytest.mark.asyncio async def test_custom_auth_does_not_enforce_key_model_access_by_default(): valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"]) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index dc052d8f012..ec7f06ac099 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -312,16 +312,19 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda @pytest.mark.asyncio -async def test_update_database_and_spend_counters_releases_reservation_when_counter_update_fails(): +async def test_update_database_and_spend_counters_invalidates_reservation_when_counter_update_fails(): proxy_logging_obj = MagicMock() proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() increment_spend_counters = AsyncMock(side_effect=Exception("counter unavailable")) - budget_reservation = {"reserved_cost": 0.5, "entries": []} + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } with patch( - "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", new_callable=AsyncMock, - ) as mock_release_budget_reservation: + ) as mock_invalidate_budget_reservation_counters: with pytest.raises(Exception, match="counter unavailable"): await _update_database_and_spend_counters( proxy_logging_obj=proxy_logging_obj, @@ -339,7 +342,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_coun budget_reservation=budget_reservation, ) - mock_release_budget_reservation.assert_awaited_once_with( + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( budget_reservation=budget_reservation, ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 9e154dcb842..c14582638c4 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -14,6 +14,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.spend_tracking.budget_reservation import ( estimate_request_max_cost, + invalidate_budget_reservation_counters, release_budget_reservation, reserve_budget_for_request, ) @@ -50,6 +51,20 @@ def _request_body() -> dict: } +def test_should_not_serialize_budget_reservation_on_user_api_key_auth(): + auth = UserAPIKeyAuth( + token="key-budget-runtime-state", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:key-budget-runtime-state"}], + }, + ) + + assert "budget_reservation" not in auth.model_dump() + assert "budget_reservation" not in auth.model_dump(exclude_none=True) + assert "budget_reservation" not in auth.model_dump_json() + + @pytest.mark.asyncio async def test_should_prevent_second_key_reservation_over_budget( spend_counter_state, @@ -467,6 +482,40 @@ async def test_should_retry_partial_release_without_double_decrement( ) == pytest.approx(0.0) +@pytest.mark.asyncio +async def test_should_invalidate_reserved_counters_after_persisted_spend_failure( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + await counter_cache.async_increment_cache( + key="spend:key:key-budget-invalidate", + value=0.4, + ) + await counter_cache.async_increment_cache( + key="spend:team:team-budget-invalidate", + value=0.4, + ) + + await invalidate_budget_reservation_counters( + { + "reserved_cost": 0.4, + "entries": [ + {"counter_key": "spend:key:key-budget-invalidate"}, + {"counter_key": "spend:team:team-budget-invalidate"}, + ], + } + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-invalidate") + is None + ) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:team-budget-invalidate") + is None + ) + + @pytest.mark.asyncio async def test_should_reserve_all_budgeted_counters(spend_counter_state): counter_cache, key_cache = spend_counter_state