harden budget reservation recovery

This commit is contained in:
user 2026-04-29 21:06:30 -07:00
parent 926de696a1
commit 09503ebb8f
7 changed files with 117 additions and 7 deletions

View file

@ -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

View file

@ -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)

View file

@ -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,
)

View file

@ -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],

View file

@ -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"])

View file

@ -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,
)

View file

@ -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