diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index bf9ed6c376c..00b14f63755 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -22,6 +22,7 @@ Works across multiple proxy instances via DualCache (in-memory + Redis). """ import os +import secrets from typing import TYPE_CHECKING, Any, Optional, Union from litellm import DualCache @@ -62,9 +63,11 @@ return new_val # Default TTL for session budget counters (1 hour) DEFAULT_MAX_BUDGET_PER_SESSION_TTL = 3600 -_RESERVED_COST_KEY = "_litellm_session_budget_reserved_cost" -_RESERVATION_RELEASED_KEY = "_litellm_session_budget_reservation_released" -_RESERVATION_SESSION_KEY = "_litellm_session_budget_session_id" +_RESERVATION_KEY = "_litellm_session_budget_reservation" +_RESERVATION_TOKEN_FIELD = "token" +_RESERVATION_SESSION_FIELD = "session_id" +_RESERVATION_COST_FIELD = "reserved_cost" +_RESERVATION_RELEASED_FIELD = "released" _EPSILON = 1e-12 @@ -82,6 +85,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): def __init__(self, internal_usage_cache: InternalUsageCache): self.internal_usage_cache = internal_usage_cache + self._reservation_token = secrets.token_hex(16) self.ttl = int( os.getenv( "LITELLM_MAX_BUDGET_PER_SESSION_TTL", @@ -261,28 +265,27 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): """ await self._reconcile_reservation(container=request_data, actual_cost=0.0) - async def _reconcile_reservation(self, container: dict, actual_cost: float) -> bool: + async def _reconcile_reservation(self, container: Any, actual_cost: float) -> bool: """ Adjust the session counter from the reserved amount to `actual_cost`. - Returns True when a reservation existed (whether or not it still needed - adjusting), so the success path knows to skip the legacy full-cost - increment. Idempotent across the failure/post-call-failure callbacks via - the released marker. + Only reservations this handler created are honored: the stored token + must match this instance's secret, so caller-supplied metadata can't + forge a reservation and drive a negative adjustment on another session. + Returns True when such a reservation existed, so the success path knows + to skip the legacy full-cost increment. Idempotent via the released + marker on the reservation record. """ - reserved_cost = self._lookup_reserved_cost(container) - if reserved_cost is None: + reservation = self._load_active_reservation(container) + if reservation is None: return False - if self._reservation_released(container): - return True - session_id = self._lookup_reservation_session_id(container) - if session_id is not None: - adjustment = actual_cost - reserved_cost - if adjustment != 0: - await self._increment_spend(self._make_cache_key(session_id), adjustment) - - self._mark_reservation_released(container) + reservation[_RESERVATION_RELEASED_FIELD] = True + session_id = str(reservation[_RESERVATION_SESSION_FIELD]) + reserved_cost = float(reservation[_RESERVATION_COST_FIELD]) + adjustment = actual_cost - reserved_cost + if adjustment != 0: + await self._increment_spend(self._make_cache_key(session_id), adjustment) return True def _estimate_reservation_cost(self, data: dict, call_type: str) -> float | None: @@ -320,54 +323,51 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): ) def _stash_reservation(self, data: dict, session_id: str, reserved_cost: float) -> None: - self._stash_in_metadata_channels(data=data, key=_RESERVED_COST_KEY, value=reserved_cost) - self._stash_in_metadata_channels(data=data, key=_RESERVATION_SESSION_KEY, value=session_id) + reservation = { + _RESERVATION_TOKEN_FIELD: self._reservation_token, + _RESERVATION_SESSION_FIELD: session_id, + _RESERVATION_COST_FIELD: reserved_cost, + _RESERVATION_RELEASED_FIELD: False, + } + channels = self._metadata_channels(data) + if not channels: + data["metadata"] = {} + channels = (data["metadata"],) + for channel in channels: + channel[_RESERVATION_KEY] = reservation - @staticmethod - def _stash_in_metadata_channels(data: dict, key: str, value: Any) -> None: - for channel in ("metadata", "litellm_metadata"): - existing = data.get(channel) - if isinstance(existing, dict): - existing[key] = value - elif channel == "metadata": - data[channel] = {key: value} - - def _lookup_reserved_cost(self, container: Any) -> float | None: - value = self._lookup_stashed_value(container, _RESERVED_COST_KEY) - if value is None: - return None - try: - return float(value) - except (TypeError, ValueError): - return None - - def _lookup_reservation_session_id(self, container: Any) -> str | None: - value = self._lookup_stashed_value(container, _RESERVATION_SESSION_KEY) - return str(value) if value is not None else None - - def _reservation_released(self, container: Any) -> bool: - return self._lookup_stashed_value(container, _RESERVATION_RELEASED_KEY) is True - - def _mark_reservation_released(self, container: Any) -> None: - if isinstance(container, dict): - self._stash_in_metadata_channels(data=container, key=_RESERVATION_RELEASED_KEY, value=True) - - @staticmethod - def _lookup_stashed_value(container: Any, key: str) -> Any: - """Resolve a stashed value from any metadata channel a callback sees.""" - if not isinstance(container, dict): - return None - for channel in ("metadata", "litellm_metadata"): - channel_dict = container.get(channel) - if isinstance(channel_dict, dict) and key in channel_dict: - return channel_dict.get(key) - litellm_params = container.get("litellm_params") - if isinstance(litellm_params, dict): - lp_metadata = litellm_params.get("metadata") - if isinstance(lp_metadata, dict) and key in lp_metadata: - return lp_metadata.get(key) + def _load_active_reservation(self, container: Any) -> dict | None: + """ + Return the reservation record this handler stamped on the request, or + None. A record is honored only when its token matches this instance's + secret (proving provenance) and it has not already been released. + """ + for channel in self._metadata_channels(container): + reservation = channel.get(_RESERVATION_KEY) + if not isinstance(reservation, dict): + continue + if reservation.get(_RESERVATION_TOKEN_FIELD) != self._reservation_token: + continue + if reservation.get(_RESERVATION_RELEASED_FIELD) is True: + return None + session_id = reservation.get(_RESERVATION_SESSION_FIELD) + reserved_cost = reservation.get(_RESERVATION_COST_FIELD) + if isinstance(session_id, str) and isinstance(reserved_cost, (int, float)): + return reservation return None + @staticmethod + def _metadata_channels(container: Any) -> "tuple[dict, ...]": + if not isinstance(container, dict): + return () + litellm_params = container.get("litellm_params") + candidates = ( + container.get("metadata"), + container.get("litellm_metadata"), + litellm_params.get("metadata") if isinstance(litellm_params, dict) else None, + ) + return tuple(channel for channel in candidates if isinstance(channel, dict)) + def _get_session_id(self, data: dict) -> Optional[str]: """Extract session_id from request metadata.""" metadata = data.get("metadata") or {} diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py index 4f06a870a42..ee64e43458c 100644 --- a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py @@ -292,6 +292,46 @@ async def test_reservation_refunded_on_failure(): assert await handler._get_current_spend(handler._make_cache_key(session_id)) == pytest.approx(0.0) +@pytest.mark.asyncio +async def test_forged_reservation_metadata_is_ignored(): + """ + Reconciliation must only honor reservations this handler created. A caller + that injects reservation keys into request metadata must not be able to + apply a negative adjustment to (and drain the budget of) an arbitrary + session. + """ + handler = _make_handler() + victim_session = "victim-session" + victim_key = handler._make_cache_key(victim_session) + + await handler._increment_spend(victim_key, 5.0) + assert await handler._get_current_spend(victim_key) == pytest.approx(5.0) + + forged_metadata = { + "session_id": victim_session, + "_litellm_session_budget_reservation": { + "token": "attacker-guessed-token", + "session_id": victim_session, + "reserved_cost": 1000.0, + "released": False, + }, + } + + await handler.async_log_failure_event( + kwargs={"litellm_params": {"metadata": forged_metadata}}, + response_obj=None, + start_time=None, + end_time=None, + ) + await handler.async_post_call_failure_hook( + request_data={"metadata": forged_metadata}, + original_exception=Exception("boom"), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-attacker"), + ) + + assert await handler._get_current_spend(victim_key) == pytest.approx(5.0) + + @pytest.mark.asyncio async def test_reservation_refund_is_idempotent_across_failure_hooks(): """