mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
harden budget reservation recovery
This commit is contained in:
parent
926de696a1
commit
09503ebb8f
7 changed files with 117 additions and 7 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue