diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 3e13848db02..1d7afcbee8f 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -313,7 +313,9 @@ class DualCache(BaseCache): else: self.last_redis_batch_access_time[key] = previous_time - async def _prepare_batch_get(self, keys: list[str], local_only: bool, **kwargs: object) -> PendingBatchRead: + async def _prepare_batch_get( + self, keys: list[str], local_only: bool, throttle_redis: bool = True, **kwargs: object + ) -> PendingBatchRead: result: list[object | None] = [None] * len(keys) if self.in_memory_cache is not None: in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) @@ -324,7 +326,10 @@ class DualCache(BaseCache): redis_keys: list[str] = [] previous_access_times: dict[str, float | None] = {} if None in result and self.redis_cache is not None and local_only is False: - redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result) + if throttle_redis: + redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result) + else: + redis_keys = [key for key, value in zip(keys, result) if value is None] return PendingBatchRead( keys=keys, result=result, redis_keys=redis_keys, previous_access_times=previous_access_times ) @@ -349,10 +354,13 @@ class DualCache(BaseCache): keys: list, parent_otel_span: Span | None = None, local_only: bool = False, + throttle_redis: bool = True, **kwargs, ): + """With ``throttle_redis`` False every key memory cannot serve is read from Redis, exactly as a per-key + ``async_get_cache`` would read it, instead of skipping keys that missed within ``redis_batch_cache_expiry``.""" try: - pending: Final = await self._prepare_batch_get(keys, local_only, **kwargs) + pending: Final = await self._prepare_batch_get(keys, local_only, throttle_redis, **kwargs) # Only hit Redis for keys memory could not serve and enough time has passed since last access. if not pending.redis_keys or self.redis_cache is None: return pending.result diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e3ce9bcd850..ed4d63fb9fc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3006,41 +3006,40 @@ async def _run_centralized_common_checks( skip_budget_checks=skip_budget_checks, project_object=project_object, ) + if not skip_budget_checks: + await _check_team_model_budget( + valid_token=user_api_key_auth_obj, + model_max_budget_limiter=model_max_budget_limiter, + models=_get_model_names_for_budget_checks( + model=_get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + team_id=user_api_key_auth_obj.team_id, + ) + ), + ) + + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request=request, + request_data=request_data, + route=route, + llm_router=llm_router, + team_object=team_object, + user_object=user_object, + end_user_id=end_user_id, + end_user_object=end_user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + skip_budget_checks=skip_budget_checks, + general_settings=general_settings, + ) finally: release_spend_counter_batch() - if not skip_budget_checks: - await _check_team_model_budget( - valid_token=user_api_key_auth_obj, - model_max_budget_limiter=model_max_budget_limiter, - models=_get_model_names_for_budget_checks( - model=_get_model_from_request_context( - request_data=request_data, - route=route, - request=request, - llm_router=llm_router, - team_id=user_api_key_auth_obj.team_id, - ) - ), - ) - - await _reserve_budget_after_common_checks( - user_api_key_auth_obj=user_api_key_auth_obj, - request=request, - request_data=request_data, - route=route, - llm_router=llm_router, - team_object=team_object, - user_object=user_object, - end_user_id=end_user_id, - end_user_object=end_user_object, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - skip_budget_checks=skip_budget_checks, - general_settings=general_settings, - ) - async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 0178465739b..05995d22293 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -30,6 +30,7 @@ from litellm.proxy.db.db_spend_update_writer import ( get_llm_router, ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.spend_tracking.spend_counter_batch import post_call_counter_keys, spend_counter_batch_scope from litellm.proxy.spend_tracking.spend_event import ( ObjectMapping, SpendEventBuildError, @@ -695,10 +696,67 @@ async def _update_database_and_spend_counters( model_access_groups: Sequence[str] | None = None, project_id: str | None = None, ) -> bool: + """The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then + spans the database write and the counter update, so the post-call counters are read with a single MGET after the + write and their increments leave in a single pipeline.""" + from litellm.proxy.proxy_server import spend_counter_cache + from litellm.proxy.spend_tracking.budget_reservation import get_reserved_counter_keys + if budget_reservation is not None: await _reconcile_budget_reservation_before_db_update( budget_reservation=budget_reservation, response_cost=response_cost ) + counter_keys: Final = frozenset( + get_reserved_counter_keys(budget_reservation=budget_reservation) + ) | post_call_counter_keys( + token=user_api_key, + team_id=team_id, + user_id=user_id, + org_id=org_id, + end_user_id=end_user_id, + tags=request_tags, + model_access_groups=model_access_groups, + project_id=project_id, + ) + with spend_counter_batch_scope(spend_counter_cache.redis_cache, counter_keys=counter_keys): + return await _update_database_and_spend_counters_in_batch( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + org_id=org_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + budget_reservation=budget_reservation, + request_tags=request_tags, + model_access_groups=model_access_groups, + project_id=project_id, + ) + + +async def _update_database_and_spend_counters_in_batch( + proxy_logging_obj: "ProxyLogging", + increment_spend_counters: _IncrementSpendCounters, + user_api_key: str | None, + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + org_id: str | None, + kwargs: dict, + completion_response: object, + start_time: datetime | None, + end_time: datetime | None, + response_cost: float, + budget_reservation: dict | None, + request_tags: list[str] | None, + model_access_groups: Sequence[str] | None, + project_id: str | None, +) -> bool: try: charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database( token=user_api_key, @@ -762,11 +820,13 @@ async def _reconcile_budget_reservation_before_db_update( budget_reservation: dict, # mutable-ok: reconcile_budget_reservation stamps applied_adjustment on the caller's shared reservation dict response_cost: float, ) -> None: + """Reseeds the reserved counters that were flushed since reservation; the adjustments themselves are written by ``increment_spend_counters`` in the same pipeline as its increments, or by + the release / invalidation that runs when the spend write fails.""" from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation try: - await reconcile_budget_reservation( - budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False + _ = await reconcile_budget_reservation( + budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False, apply_consistent=False ) except Exception: # noqa: BLE001 # a failed reconcile must not block the spend write; the counters are dropped instead verbose_proxy_logger.warning( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6a221f2bfed..b4a497ea1e4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -759,6 +759,7 @@ from litellm.proxy.shutdown.scheduled_jobs import ( from litellm.proxy.spend_tracking.budget_reservation import ( get_budget_window_start, release_unbound_budget_reservation, + stamp_budget_reservation_actual_cost, ) from litellm.proxy.spend_tracking.spend_capture_rate import ( run_scheduled_spend_capture_rate_check, @@ -2846,13 +2847,16 @@ async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None if spend_counter_cache.redis_cache is not None: forget_spend_counter(counter_key) try: - await spend_counter_cache.redis_cache.async_set_max(key=counter_key, value=db_spend) + repaired: Final = await spend_counter_cache.redis_cache.async_set_max(key=counter_key, value=db_spend) except Exception: verbose_proxy_logger.debug( "Unable to repair stale spend counter %s in Redis", counter_key, exc_info=True, ) + return + if repaired is not None: + record_spend_counter_value(counter_key, repaired) async def reseed_spend_counter_from_db(counter_key: str) -> bool: @@ -3049,13 +3053,17 @@ async def _increment_spend_counters_batched( model_access_groups: Sequence[str] | None, project_id: str | None = None, ): - """Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET.""" - reserved_counter_keys: Final = await _reconcile_budget_reservation_for_counter_update( + """Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET, and + the reconcile adjustments go out in the same INCRBYFLOAT pipeline as the counter increments.""" + reservation_update: Final = await _reconcile_budget_reservation_for_counter_update( budget_reservation=budget_reservation, response_cost=response_cost, ) + reserved_counter_keys: Final = reservation_update.reserved_counter_keys if response_cost is None or response_cost == 0: + await _apply_spend_counter_increments(pending=reservation_update.pending) + stamp_budget_reservation_actual_cost(budget_reservation=budget_reservation, actual_cost=response_cost) if budget_reservation is not None: budget_reservation["finalized"] = True return @@ -3276,7 +3284,8 @@ async def _increment_spend_counters_batched( for item in scope if not isinstance(item, BaseException) ) - await _apply_spend_counter_increments(pending=pending) + await _apply_spend_counter_increments(pending=reservation_update.pending + pending) + stamp_budget_reservation_actual_cost(budget_reservation=budget_reservation, actual_cost=response_cost) if scope_errors: raise scope_errors[0] @@ -3284,12 +3293,21 @@ async def _increment_spend_counters_batched( budget_reservation["finalized"] = True +@dataclass(frozen=True, slots=True) +class _ReservationCounterUpdate: + """The reserved counters the direct increment must skip, and the adjustments that settle them on the actual + cost, still to be written; both empty when the reservation could not be reconciled and was dropped.""" + + reserved_counter_keys: frozenset[str] = frozenset() + pending: tuple[PendingSpendIncrement, ...] = () + + async def _reconcile_budget_reservation_for_counter_update( budget_reservation: dict | None, response_cost: float | None, -) -> set[str]: +) -> _ReservationCounterUpdate: if budget_reservation is None or budget_reservation.get("finalized") is True: - return set() + return _ReservationCounterUpdate() from litellm.proxy.spend_tracking.budget_reservation import ( get_reserved_counter_keys, @@ -3299,10 +3317,11 @@ async def _reconcile_budget_reservation_for_counter_update( reserved_counter_keys: Final = get_reserved_counter_keys(budget_reservation=budget_reservation) try: - await reconcile_budget_reservation( + pending: Final = await reconcile_budget_reservation( budget_reservation=budget_reservation, actual_cost=response_cost or 0.0, finalize=False, + apply_consistent=False, ) except Exception: verbose_proxy_logger.warning( @@ -3315,8 +3334,8 @@ async def _reconcile_budget_reservation_for_counter_update( verbose_proxy_logger.exception( "Failed to invalidate reserved counters after reservation reconciliation failed" ) - return set() - return reserved_counter_keys + return _ReservationCounterUpdate() + return _ReservationCounterUpdate(reserved_counter_keys=frozenset(reserved_counter_keys), pending=pending) async def _prepare_end_user_and_tag_spend_increments( @@ -3702,31 +3721,81 @@ async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncremen raise -async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> None: - """One INCRBYFLOAT+EXPIRE pipeline for every pending counter; on failure every counter is invalidated - before the error propagates, so no caller can read a half-applied batch.""" +async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]: + """One INCRBYFLOAT+EXPIRE pipeline for every pending counter, returning each counter's new value in order; on + failure every counter is invalidated before the error propagates, so no caller can read a half-applied batch.""" + if spend_counter_cache.redis_cache is None: + return await run_spend_counter_pipeline(pending=pending) + try: + return await run_spend_counter_pipeline(pending=pending) + except Exception: + await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) + raise + + +async def run_spend_counter_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]: + """The pipeline behind ``increment_spend_counters_pipeline`` without its invalidation: the caller decides what + happens to counters whose increment may or may not have landed when the pipeline fails.""" if not pending: - return + return () redis_cache: Final = spend_counter_cache.redis_cache if redis_cache is None: - for item in pending: - await SpendCounterReseed.increment_in_memory( - spend_counter_cache=spend_counter_cache, counter_key=item.counter_key, increment=item.increment - ) - return + return tuple( + [ + await SpendCounterReseed.increment_in_memory( + spend_counter_cache=spend_counter_cache, counter_key=item.counter_key, increment=item.increment + ) + for item in pending + ] + ) ttl: Final = redis_cache.get_ttl() increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation] RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl) for item in pending ] - try: - results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list) - except Exception: - await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) - raise + results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list) for item, current_value in zip(pending, results or ()): spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value) record_spend_counter_value(item.counter_key, float(current_value)) + return tuple(float(current_value) for current_value in results or ()) + + +def _update_cache_read_keys( + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + tags: Sequence[object] | None, + response_cost: float | None, +) -> tuple[str, ...]: + if response_cost is None: + return () + user_keys: tuple[str, ...] = (user_id, GLOBAL_PROXY_SPEND_CACHE_KEY) if user_id is not None else () + end_user_keys: tuple[str, ...] = (end_user_cache_key(end_user_id),) if end_user_id is not None else () + team_keys: tuple[str, ...] = (f"team_id:{team_id}",) if team_id is not None else () + tag_keys: tuple[str, ...] = tuple(tag_cache_key(tag) for tag in tags or () if isinstance(tag, str) and tag) + return user_keys + end_user_keys + team_keys + tag_keys + + +async def _read_update_cache_values(keys: Sequence[str], parent_otel_span: Span | None) -> Mapping[str, object]: + """One batched read for every object ``update_cache`` refreshes; a failed read leaves them all untouched, + exactly as a failed per-object GET left that object untouched.""" + if not keys: + return MappingProxyType({}) + try: + values: Final = await user_api_key_cache.async_batch_get_cache( + keys=list(keys), parent_otel_span=parent_otel_span, throttle_redis=False + ) + except Exception as e: + verbose_proxy_logger.warning( + "Spend tracking - failed to read cached spend objects. Budget enforcement may use stale spend values. " + "keys=%s - %s", + keys, + str(e), + ) + return MappingProxyType({}) + if values is None: + return MappingProxyType({}) + return MappingProxyType({key: value for key, value in zip(keys, values) if value is not None}) async def update_cache( @@ -3745,6 +3814,12 @@ async def update_cache( """ values_to_update_in_cache: Final[list[tuple[str, object]]] = [] + cached_values: Final = await _read_update_cache_values( + keys=_update_cache_read_keys( + user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost + ), + parent_otel_span=parent_otel_span, + ) ### UPDATE KEY SPEND ### async def _update_key_cache(token: str, response_cost: float): @@ -3810,7 +3885,7 @@ async def update_cache( # Fetch the existing cost for the given user if _id is None: continue - cached_user = await user_api_key_cache.async_get_cache(key=_id) + cached_user = cached_values.get(_id) if cached_user is None: # do nothing if there is no cache value return @@ -3833,11 +3908,11 @@ async def update_cache( ) ) ## UPDATE GLOBAL PROXY ## - global_proxy_spend: Final = await user_api_key_cache.async_get_cache(key=GLOBAL_PROXY_SPEND_CACHE_KEY) - if global_proxy_spend is None: + global_proxy_spend: Final = cached_values.get(GLOBAL_PROXY_SPEND_CACHE_KEY) + if not isinstance(global_proxy_spend, (int, float)): # do nothing if not in cache return - elif response_cost is not None and global_proxy_spend is not None: + elif response_cost is not None: increment: Final = global_proxy_spend + response_cost values_to_update_in_cache.append((GLOBAL_PROXY_SPEND_CACHE_KEY, increment)) except Exception as e: @@ -3859,7 +3934,7 @@ async def update_cache( _id: Final = end_user_cache_key(end_user_id) try: # Fetch the existing cost for the given user - cached_end_user: Final = await user_api_key_cache.async_get_cache(key=_id) + cached_end_user: Final = cached_values.get(_id) if cached_end_user is None: # if user does not exist in LiteLLM_UserTable, create a new user # do nothing if end-user not in api key cache @@ -3900,7 +3975,7 @@ async def update_cache( _id: Final = f"team_id:{team_id}" try: - cached_team: Final = await user_api_key_cache.async_get_cache(key=_id) + cached_team: Final = cached_values.get(_id) if cached_team is None: # do nothing if team not in api key cache return @@ -3950,7 +4025,7 @@ async def update_cache( cache_key = tag_cache_key(tag_name) # Fetch the existing tag object from cache - cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) + cached_tag = cached_values.get(cache_key) if cached_tag is None: # do nothing if tag not in api key cache continue diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index e28fa2c06a4..c094e91c6c0 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -4,7 +4,7 @@ import asyncio import json import math import time -from collections.abc import Mapping, Sequence +from collections.abc import AsyncIterator, Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta, timezone from types import MappingProxyType @@ -290,7 +290,6 @@ async def reserve_budget_for_request( raw_body=raw_body, ) - current_spend_by_counter_key: Final[dict[str, float]] = {} reservation_cost = estimate_request_max_cost( request_body=request_body, route=route, @@ -306,46 +305,17 @@ async def reserve_budget_for_request( applied_entries: Final[list[dict[str, float | str]]] = [] try: with _counters_batch_scope(frozenset(counter.counter_key for counter in counters)): - for counter in counters: - entry = _counter_to_reservation_entry( - counter=counter, - reserved_cost=reservation_cost, - ) - applied_entries.append(entry) - try: - reserved_value = await _reserve_counter( - counter=counter, - reservation_cost=reservation_cost, - ) - except _CounterReservationUnavailable as exc: - if exc.touched_counter and not exc.counter_invalidated: - await _release_applied_entries_best_effort( - entries=[entry], - default_reserved_cost=reservation_cost, - ) - applied_entries.remove(entry) - if fail_closed_budget_enforcement: - _raise_reservation_unavailable(counter_key=counter.counter_key) - continue - - if reserved_value is not None: - current_spend = reserved_value - else: - cached_spend = current_spend_by_counter_key.get(counter.counter_key) - if cached_spend is None: - cached_spend = await _get_current_counter_value(counter=counter) - current_spend = cached_spend + reservation_cost - if current_spend > counter.max_budget: - reservation_cost = await _apply_over_budget_reservation_policy( - counter=counter, - valid_token=valid_token, - entry=entry, - applied_entries=applied_entries, - reservation_cost=reservation_cost, - current_spend=current_spend, - fail_closed_budget_enforcement=fail_closed_budget_enforcement, - ) - continue + reservable: Final = await _initialize_reservation_counters( + counters=counters, + fail_closed_budget_enforcement=fail_closed_budget_enforcement, + ) + reservation_cost = await _reserve_reservable_counters( + reservable=reservable, + valid_token=valid_token, + applied_entries=applied_entries, + reservation_cost=reservation_cost, + fail_closed_budget_enforcement=fail_closed_budget_enforcement, + ) except Exception: await _release_applied_entries_best_effort( entries=applied_entries, @@ -381,19 +351,39 @@ async def reconcile_budget_reservation( budget_reservation: dict | None, actual_cost: float | None, finalize: bool = True, -) -> None: + apply_consistent: bool = True, +) -> tuple[PendingSpendIncrement, ...]: + """Settle every reserved counter on ``actual_cost``. With ``apply_consistent`` False the adjustments for + counters that still hold the reservation are returned instead of written, so the caller can pipeline them with + its own increments and then call ``stamp_budget_reservation_actual_cost``.""" if not budget_reservation or budget_reservation.get("finalized") is True: - return + return () reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0) actual: Final = float(actual_cost or 0.0) - await _set_reserved_entries_actual_cost( + pending: Final = await _set_reserved_entries_actual_cost( entries=budget_reservation.get("entries") or [], actual_cost=actual, default_reserved_cost=reserved_cost, + apply_consistent=apply_consistent, ) if finalize: budget_reservation["finalized"] = True + return pending + + +def stamp_budget_reservation_actual_cost(budget_reservation: dict | None, actual_cost: float | None) -> None: + """Record that every reserved counter now holds ``actual_cost``, once the adjustments handed back by + ``reconcile_budget_reservation(apply_consistent=False)`` have been written.""" + if not budget_reservation: + return + reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0) + actual: Final = float(actual_cost or 0.0) + for entry in budget_reservation.get("entries") or []: + if "counter_key" in entry: + entry["applied_adjustment"] = actual - _get_entry_reserved_cost( + entry=entry, default_reserved_cost=reserved_cost + ) async def release_budget_reservation(budget_reservation: dict | None) -> None: @@ -917,18 +907,40 @@ def _coerce_window(window: object) -> Mapping[str, object]: return dumped if isinstance(dumped, Mapping) else {} -async def _reserve_counter( - counter: _BudgetCounter, - reservation_cost: float, -) -> float | None: +async def _initialize_reservation_counters( + counters: Sequence[_BudgetCounter], + fail_closed_budget_enforcement: bool, +) -> tuple[_BudgetCounter, ...]: + """The counters whose current value is loaded, in order; one that cannot be loaded is skipped (or rejects the + request under fail-closed enforcement) exactly as it was when each counter was reserved on its own.""" + return tuple([counter async for counter in _loaded_reservation_counters(counters, fail_closed_budget_enforcement)]) + + +async def _loaded_reservation_counters( + counters: Sequence[_BudgetCounter], fail_closed_budget_enforcement: bool +) -> AsyncIterator[_BudgetCounter]: + for counter in counters: + if await _reservation_counter_loaded(counter, fail_closed_budget_enforcement): + yield counter + + +async def _reservation_counter_loaded(counter: _BudgetCounter, fail_closed_budget_enforcement: bool) -> bool: + try: + await _initialize_reservation_counter(counter=counter) + except _CounterReservationUnavailable: + if fail_closed_budget_enforcement: + _raise_reservation_unavailable(counter_key=counter.counter_key) + return False + return True + + +async def _initialize_reservation_counter(counter: _BudgetCounter) -> None: from litellm.proxy.proxy_server import ( _ensure_spend_counter_initialized, _ensure_window_spend_counter_initialized, - _increment_spend_counter_cache, _invalidate_spend_counter, ) - attempted_increment = False try: if counter.source_cache_key is not None: await _ensure_spend_counter_initialized( @@ -949,13 +961,6 @@ async def _reserve_counter( counter.counter_key, ) raise _CounterReservationUnavailable - - attempted_increment = True - reserved_value: Final = await _increment_spend_counter_cache( - counter_key=counter.counter_key, - increment=reservation_cost, - ) - return float(reserved_value) if reserved_value is not None else None except _CounterReservationUnavailable: raise except Exception: @@ -964,20 +969,121 @@ async def _reserve_counter( counter.counter_key, exc_info=True, ) - counter_invalidated = False try: await _invalidate_spend_counter(counter_key=counter.counter_key) - counter_invalidated = True except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", counter.counter_key, exc_info=True, ) - raise _CounterReservationUnavailable( - touched_counter=attempted_increment, - counter_invalidated=counter_invalidated, + raise _CounterReservationUnavailable + + +async def _reserve_reservable_counters( + reservable: Sequence[_BudgetCounter], + valid_token: UserAPIKeyAuth | None, + applied_entries: list[dict[str, float | str]], + reservation_cost: float, + fail_closed_budget_enforcement: bool, +) -> float: + """Charge the counters group by group (see ``_reservation_groups``), settling the over-budget policy on each + group before the next is charged, and hand back the reservation cost the policy left standing.""" + current_spend_by_counter_key: Final = { + counter.counter_key: await _get_current_counter_value(counter=counter) for counter in reservable + } + for group in _reservation_groups( + counters=reservable, + current_spend_by_counter_key=current_spend_by_counter_key, + reservation_cost=reservation_cost, + ): + charged_cost = reservation_cost + entries = tuple(_counter_to_reservation_entry(counter=counter, reserved_cost=charged_cost) for counter in group) + applied_entries.extend(entries) + reserved_values = await _reserve_counters(counters=group, entries=entries, reservation_cost=charged_cost) + if reserved_values is None: + for entry in entries: + applied_entries.remove(entry) + if fail_closed_budget_enforcement: + _raise_reservation_unavailable(counter_key=group[0].counter_key) + continue + for counter, entry, reserved_value in zip(group, entries, reserved_values): + if entry not in applied_entries: + continue + if reserved_value is not None: + current_spend = reserved_value - (charged_cost - reservation_cost) + else: + current_spend = current_spend_by_counter_key[counter.counter_key] + reservation_cost + if current_spend > counter.max_budget: + reservation_cost = await _apply_over_budget_reservation_policy( + counter=counter, + valid_token=valid_token, + entry=entry, + applied_entries=applied_entries, + reservation_cost=reservation_cost, + current_spend=current_spend, + fail_closed_budget_enforcement=fail_closed_budget_enforcement, + ) + return reservation_cost + + +def _reservation_groups( + counters: Sequence[_BudgetCounter], + current_spend_by_counter_key: Mapping[str, float], + reservation_cost: float, +) -> tuple[tuple[_BudgetCounter, ...], ...]: + """Every counter the batch read says still has room for the estimate is charged in one pipeline. As soon as one + does not, the counters are charged one at a time so the over-budget policy settles each before the next is + touched, and a rejection charges nothing after it.""" + if not counters: + return () + if all( + current_spend_by_counter_key[counter.counter_key] + reservation_cost <= counter.max_budget + for counter in counters + ): + return (tuple(counters),) + return tuple((counter,) for counter in counters) + + +async def _reserve_counters( + counters: Sequence[_BudgetCounter], + entries: Sequence[dict[str, float | str]], + reservation_cost: float, +) -> tuple[float | None, ...] | None: + """One INCRBYFLOAT pipeline reserves every counter. When it fails each counter is dropped, and one that cannot + be dropped is released instead in case its increment landed, so nothing is left to release by the caller.""" + from litellm.proxy.proxy_server import _invalidate_spend_counter, run_spend_counter_pipeline + + if not counters: + return () + try: + reserved: Final = await run_spend_counter_pipeline( + pending=tuple( + PendingSpendIncrement(counter_key=counter.counter_key, increment=reservation_cost) + for counter in counters + ) ) + except Exception: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because spend counter reservation failed", + tuple(counter.counter_key for counter in counters), + exc_info=True, + ) + for counter, entry in zip(counters, entries): + try: + await _invalidate_spend_counter(counter_key=counter.counter_key) + except Exception: + verbose_proxy_logger.warning( + "Failed to invalidate spend counter after budget reservation failure for %s", + counter.counter_key, + exc_info=True, + ) + await _release_applied_entries_best_effort( + entries=[entry], # mutable-ok: the release takes the reservation's list of entries + default_reserved_cost=reservation_cost, + ) + return None + return tuple(reserved) + (None,) * (len(counters) - len(reserved)) async def _get_current_counter_value(counter: _BudgetCounter) -> float: @@ -1026,9 +1132,11 @@ async def _set_reserved_entries_actual_cost( actual_cost: float, default_reserved_cost: float, reseed_on_inconsistent: bool = True, -) -> None: - """Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline. - A counter that was flushed or reseeded since reservation is settled on its own after the pipeline.""" + apply_consistent: bool = True, +) -> tuple[PendingSpendIncrement, ...]: + """Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline, or are + returned unwritten when ``apply_consistent`` is False. A counter that was flushed or reseeded since reservation + is settled on its own after the pipeline.""" from litellm.proxy.proxy_server import increment_spend_counters_pipeline with _counters_batch_scope(frozenset(str(entry["counter_key"]) for entry in entries if "counter_key" in entry)): @@ -1055,15 +1163,16 @@ async def _set_reserved_entries_actual_cost( f"Cannot resize budget reservation against inconsistent counter {inconsistent[0].counter_key}" ) applicable: Final = tuple(item for item, ok in zip(adjustments, consistent) if ok) - await increment_spend_counters_pipeline( - pending=tuple( - PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable - ) + applicable_pending: Final = tuple( + PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable ) + if apply_consistent: + await increment_spend_counters_pipeline(pending=applicable_pending) for item in inconsistent: await _reseed_reserved_entry(item=item, actual_cost=actual_cost) - for item in adjustments: + for item in adjustments if apply_consistent else inconsistent: item.entry["applied_adjustment"] = item.target_adjustment + return () if apply_consistent else applicable_pending async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> None: diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index ddb074ae023..a6694895a27 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -144,25 +144,43 @@ def release_spend_counter_batch() -> None: batch.close() -def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> Iterator[str]: - if token.token is not None: - yield f"spend:key:{token.token}" - if token.team_id is not None: - yield f"spend:team:{token.team_id}" - if token.user_id is not None: - yield f"spend:team_member:{token.user_id}:{token.team_id}" - if token.user_id is not None: - yield f"spend:user:{token.user_id}" - if end_user_id is not None: +def _iter_entity_counter_keys( + token: object, + team_id: object, + user_id: object, + org_id: object, + project_id: object, + end_user_id: object, +) -> Iterator[str]: + """Only string ids name a counter; anything else (None, or an unresolved placeholder in synthetic + logging payloads) simply has no counter to bind.""" + if isinstance(token, str): + yield f"spend:key:{token}" + if isinstance(team_id, str): + yield f"spend:team:{team_id}" + if isinstance(user_id, str): + yield f"spend:team_member:{user_id}:{team_id}" + if isinstance(user_id, str): + yield f"spend:user:{user_id}" + if isinstance(end_user_id, str): yield f"spend:end_user:{end_user_id}" - if token.org_id is not None: - yield f"spend:org:{token.org_id}" - if token.project_id is not None: - yield project_spend_counter_key(token.project_id) + if isinstance(org_id, str): + yield f"spend:org:{org_id}" + if isinstance(project_id, str): + yield project_spend_counter_key(project_id) def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]: - return frozenset(_iter_admission_counter_keys(token, end_user_id)) + return frozenset( + _iter_entity_counter_keys( + token=token.token, + team_id=token.team_id, + user_id=token.user_id, + org_id=token.org_id, + project_id=token.project_id, + end_user_id=end_user_id, + ) + ) def post_call_counter_keys( @@ -176,9 +194,15 @@ def post_call_counter_keys( project_id: str | None = None, ) -> frozenset[str]: """Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read.""" - entity_keys: Final = admission_counter_keys( - UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id), - end_user_id, + entity_keys: Final = frozenset( + _iter_entity_counter_keys( + token=token, + team_id=team_id, + user_id=user_id, + org_id=org_id, + project_id=project_id, + end_user_id=end_user_id, + ) ) tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str)) group_keys: Final = frozenset( 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 470db99108a..78b281d5c78 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 @@ -9290,3 +9290,67 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state(): assert result.budget_reservation == reservation assert websocket.state.budget_reservation is reservation assert websocket.scope["state"]["budget_reservation"] is reservation + + +@pytest.mark.asyncio +async def test_admission_and_budget_reservation_read_the_key_spend_counter_with_one_redis_mget(): + from fastapi import Request + from starlette.datastructures import URL + + import litellm.proxy.proxy_server as _proxy_server_mod + from litellm.proxy.spend_tracking.spend_counter_batch import ( + read_batched_spend_counter, + spend_counter_batch_scope, + ) + + token = UserAPIKeyAuth(api_key="sk-test", token="hashed", max_budget=10.0) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + reads: list[tuple[str, tuple[float | None, bool] | None]] = [] + + async def _admission_reads_spend(**kwargs): + reads.append(("admission", await read_batched_spend_counter("spend:key:hashed"))) + + async def _reservation_reads_spend(**kwargs): + reads.append(("reservation", await read_batched_spend_counter("spend:key:hashed"))) + + redis = MagicMock() + redis.async_batch_get_cache = AsyncMock(return_value={"spend:key:hashed": 4.0}) + attrs = { + **_proxy_attrs_for_centralized_checks(user_custom_auth=None), + "prisma_client": MagicMock(), + "spend_counter_cache": MagicMock(redis_cache=redis), + } + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( # test-quality-ok: authorization has its own tests above; this one checks the shared counter read + "litellm.proxy.auth.user_api_key_auth.common_checks", + new=AsyncMock(side_effect=_admission_reads_spend), + ), + patch( # test-quality-ok: the reservation helper imports reserve_budget_for_request in its body + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + side_effect=_reservation_reads_spend, + ), + spend_counter_batch_scope(redis), + ): + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + ) + reads.append(("after admission", await read_batched_spend_counter("spend:key:hashed"))) + finally: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, originals[k]) + + assert reads == [ + ("admission", (4.0, True)), + ("reservation", (4.0, True)), + ("after admission", None), + ], "admission and reservation share one snapshot, and read-then-write callers go to Redis once it closes" + assert redis.async_batch_get_cache.await_count == 1 + assert "spend:key:hashed" in redis.async_batch_get_cache.await_args.kwargs["key_list"] 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 b5e594db701..a5b2d8b0b8d 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 @@ -726,6 +726,7 @@ async def test_update_database_and_spend_counters_reconciles_reservation_before_ budget_reservation=budget_reservation, actual_cost=0.2, finalize=False, + apply_consistent=False, ) increment_spend_counters.assert_awaited_once() assert increment_spend_counters.await_args.kwargs["budget_reservation"] is budget_reservation @@ -771,6 +772,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u budget_reservation=budget_reservation, actual_cost=0.2, finalize=False, + apply_consistent=False, ) mock_release_budget_reservation.assert_awaited_once_with( budget_reservation=budget_reservation, 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 0731c233fef..ad86c3c5267 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -72,6 +72,7 @@ def _make_spend_counter_cache( def _make_user_api_key_cache(get_value=None, get_side_effect=None): cache = MagicMock() cache.async_get_cache = AsyncMock(return_value=get_value, side_effect=get_side_effect) + cache.async_batch_get_cache = AsyncMock(side_effect=lambda keys, **_: [get_value for _ in keys]) cache.async_set_cache_pipeline = AsyncMock() return cache @@ -633,7 +634,7 @@ async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch) 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()) + monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock(return_value=())) recorded: dict[str, float] = {} @@ -888,7 +889,8 @@ async def test_increment_spend_counters_pipeline_failure_invalidates_all_counter @pytest.mark.asyncio async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set_when_none(): result = await ps._reconcile_budget_reservation_for_counter_update(budget_reservation=None, response_cost=1.0) - assert result == set() + assert result.reserved_counter_keys == frozenset() + assert result.pending == () @pytest.mark.asyncio @@ -917,7 +919,8 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat budget_reservation={"foo": "bar"}, response_cost=1.0 ) - assert result == set() + assert result.reserved_counter_keys == frozenset() + assert result.pending == () assert fake_invalidate.called is True @@ -941,7 +944,8 @@ async def test_reconcile_budget_reservation_for_counter_update_finalized_reserva response_cost=1.0, ) - assert result == set() + assert result.reserved_counter_keys == frozenset() + assert result.pending == () fake_reconcile.assert_not_awaited() @@ -1531,16 +1535,15 @@ async def test_update_cache_no_cached_entities_schedules_pipeline_flush(monkeypa tags=["x"], ) - observed = { - "lookups": fake_user_cache.async_get_cache.call_count, - "got_user": True, - "got_team": True, - } - assert normalize(observed) == { - "lookups": 4, - "got_user": True, - "got_team": True, - } + assert fake_user_cache.async_get_cache.await_count == 0 + fake_user_cache.async_batch_get_cache.assert_awaited_once() + assert fake_user_cache.async_batch_get_cache.await_args.kwargs["keys"] == [ + "u1", + f"{ps.litellm_proxy_admin_name}:spend", + "end_user_id:eu1", + "team_id:t1", + "tag:x", + ] @pytest.mark.asyncio @@ -1548,7 +1551,7 @@ async def test_update_cache_user_cache_failure_invalid_state_is_swallowed(monkey """An inner _update_user_cache raising must not propagate — update_cache catches and logs, the public coroutine still completes normally.""" fake_user_cache = MagicMock() - fake_user_cache.async_get_cache = AsyncMock(side_effect=RuntimeError("cache down")) + fake_user_cache.async_batch_get_cache = AsyncMock(side_effect=RuntimeError("cache down")) fake_user_cache.async_set_cache_pipeline = AsyncMock() monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py index 6165af4920d..e0a74d50a6c 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py @@ -9,8 +9,10 @@ gives up, but ``increment_spend_counters`` still treats the counter as lands in the enforced counter, so budgets stop gating until the next cold reseed pulls a lagging value from the DB. -The fix makes the reconcile path fall back to the direct increment when it -fails, so the actual cost is always written to the shared counter. +The reconcile adjustment and the direct increment now leave in one pipeline, so +a failure either writes the actual cost or drops the counter (and surfaces the +error) for the next read to reseed from the DB; it never leaves the reserved +estimate in place as if it were reconciled. """ import pytest @@ -84,13 +86,14 @@ async def test_direct_increment_runs_when_reservation_reconcile_hits_redis_failu ], } - await proxy_server.increment_spend_counters( - token=hashed_token, - team_id=None, - user_id=None, - response_cost=response_cost, - budget_reservation=budget_reservation, - ) + with pytest.raises(Exception, match="Redis timeout"): + await proxy_server.increment_spend_counters( + token=hashed_token, + team_id=None, + user_id=None, + response_cost=response_cost, + budget_reservation=budget_reservation, + ) - enforced_spend = await flaky_redis.async_get_cache(key=counter_key) - assert enforced_spend == response_cost + assert await flaky_redis.async_get_cache(key=counter_key) is None + assert proxy_server.spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is None diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py index 3e4b817fab8..1fddfaaa766 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py @@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest +import litellm import litellm.proxy.proxy_server as ps from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth @@ -17,6 +18,7 @@ from litellm.proxy.spend_tracking.spend_counter_batch import ( active_spend_counter_batch, admission_counter_keys, bind_admission_counter_keys, + post_call_counter_keys, release_spend_counter_batch, spend_counter_batch_scope, ) @@ -86,6 +88,21 @@ def test_admission_counter_keys_cover_every_entity_the_checks_read(): ) +def test_post_call_counter_keys_skip_ids_that_are_not_strings(): + """A synthetic logging payload (batch cost polling, tests) can carry placeholders where the ids belong; those + have no counter, and deriving the key set must never raise inside the cost callback.""" + placeholder = object() + assert post_call_counter_keys( + token=placeholder, # pyright: ignore[reportArgumentType] # synthetic payload placeholder, not an id + team_id="team", + user_id=None, + org_id=placeholder, # pyright: ignore[reportArgumentType] # synthetic payload placeholder, not an id + end_user_id="eu", + tags=[placeholder, "t1"], + model_access_groups=None, + ) == {"spend:team:team", "spend:end_user:eu", "spend:tag:t1"} + + @pytest.mark.asyncio async def test_bound_counters_share_one_mget_and_a_clean_miss_is_authoritative(): redis = CountingRedis({"spend:key:hashed": 1.5, "spend:team:team": 2.5}) @@ -407,7 +424,7 @@ def _reservation(reserved_cost: float, counter_keys: frozenset[str] = RESERVED_K @pytest.mark.asyncio -async def test_post_call_with_a_reservation_costs_one_mget_one_reconcile_pipeline_one_increment_pipeline(monkeypatch): +async def test_post_call_with_a_reservation_costs_one_mget_and_one_pipeline_for_reconcile_and_increments(monkeypatch): redis = CountingRedis({key: 1.0 for key in POST_CALL_KEYS}) monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", None) @@ -425,10 +442,9 @@ async def test_post_call_with_a_reservation_costs_one_mget_one_reconcile_pipelin budget_reservation=reservation, ) - assert [c.split()[0] for c in redis.commands] == ["MGET", "PIPELINE", "PIPELINE"], redis.commands + assert [c.split()[0] for c in redis.commands] == ["MGET", "PIPELINE"], redis.commands assert set(redis.commands[0].split()[1:]) == POST_CALL_KEYS, "reconcile and warm checks share the MGET" - assert set(redis.commands[1].split()[1:]) == RESERVED_KEYS - assert set(redis.commands[2].split()[1:]) == POST_CALL_KEYS - RESERVED_KEYS + assert set(redis.commands[1].split()[1:]) == POST_CALL_KEYS, "reconcile adjustments ride the increment pipeline" assert {key: round(redis.store[key], 6) for key in POST_CALL_KEYS} == { key: (1.1 if key in RESERVED_KEYS else 1.5) for key in POST_CALL_KEYS } @@ -436,6 +452,27 @@ async def test_post_call_with_a_reservation_costs_one_mget_one_reconcile_pipelin assert reservation["finalized"] is True +@pytest.mark.asyncio +async def test_a_stale_counter_repair_updates_the_open_batch_instead_of_forcing_a_second_mget(monkeypatch): + redis = CountingRedis({"spend:key:hashed": 1.0, "spend:team:team": 1.0}) + + async def set_max(key: str, value: float, **kwargs: object) -> float: + redis.commands.append(f"SETMAX {key} {value}") + redis.store[key] = max(float(str(redis.store.get(key, 0.0))), value) + return float(str(redis.store[key])) + + redis.async_set_max = set_max + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + + with spend_counter_batch_scope(redis, counter_keys=frozenset({"spend:key:hashed", "spend:team:team"})): + assert await ps.read_spend_counter_cache_value(counter_key="spend:team:team") == (1.0, True) + await ps._repair_stale_spend_counter(counter_key="spend:team:team", db_spend=4.0) + assert await ps.read_spend_counter_cache_value(counter_key="spend:team:team") == (4.0, True) + assert await ps.read_spend_counter_cache_value(counter_key="spend:key:hashed") == (1.0, True) + + assert [c.split()[0] for c in redis.commands] == ["MGET", "SETMAX"], redis.commands + + @pytest.mark.asyncio async def test_reconcile_settles_a_flushed_counter_on_its_own_after_the_shared_pipeline(monkeypatch): from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation @@ -478,37 +515,35 @@ async def test_pre_call_resize_against_an_inconsistent_counter_writes_nothing_an @pytest.mark.asyncio -async def test_a_failed_reconcile_pipeline_invalidates_every_reserved_counter_and_falls_back(monkeypatch): +async def test_a_failed_post_call_pipeline_invalidates_every_counter_it_carried_and_stamps_nothing(monkeypatch): redis = CountingRedis({key: 1.0 for key in POST_CALL_KEYS}) redis.async_delete_cache = AsyncMock() - reconcile_pipeline_failed = False async def _pipeline(increment_list: Sequence[Mapping[str, object]], **kwargs: object) -> list[float]: - nonlocal reconcile_pipeline_failed - if not reconcile_pipeline_failed: - reconcile_pipeline_failed = True - raise ConnectionError("redis down") - return await CountingRedis.async_increment_pipeline(redis, increment_list, **kwargs) + raise ConnectionError("redis down") redis.async_increment_pipeline = _pipeline # pyright: ignore[reportAttributeAccessIssue] # instance override monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", None) reservation = _reservation(reserved_cost=0.4) - await ps.increment_spend_counters( - token="hashed", - team_id="team", - user_id="user", - org_id="org", - end_user_id="eu", - response_cost=0.5, - budget_reservation=reservation, - ) + with pytest.raises(ConnectionError): + await ps.increment_spend_counters( + token="hashed", + team_id="team", + user_id="user", + org_id="org", + end_user_id="eu", + response_cost=0.5, + budget_reservation=reservation, + ) - assert {call.kwargs["key"] for call in redis.async_delete_cache.await_args_list} == RESERVED_KEYS + assert [c.split()[0] for c in redis.commands] == ["MGET"], redis.commands + assert {call.kwargs["key"] for call in redis.async_delete_cache.await_args_list} == RESERVED_KEYS | { + "spend:user:user" + } assert all("applied_adjustment" not in entry for entry in reservation["entries"]) - assert redis.commands[-1].split()[0] == "PIPELINE" - assert set(redis.commands[-1].split()[1:]) == RESERVED_KEYS | {"spend:user:user"} + assert {key: redis.store[key] for key in POST_CALL_KEYS} == {key: 1.0 for key in POST_CALL_KEYS} def test_a_scope_opened_inside_an_open_scope_joins_its_batch_and_a_closed_one_gets_its_own(): @@ -525,3 +560,199 @@ def test_a_scope_opened_inside_an_open_scope_joins_its_batch_and_a_closed_one_ge assert inner is not outer assert inner is not None and inner.counter_keys == {"spend:key:c"} assert active_spend_counter_batch() is outer + + +@pytest.mark.asyncio +async def test_reservation_inside_the_admission_scope_reuses_its_mget_and_reserves_in_one_pipeline(monkeypatch): + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_request + + redis = CountingRedis({"spend:key:hashed": 1.0, "spend:team:team": 2.0}) + redis.default_ttl = 3600 + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr("litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", lambda **_: 0.5) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + with spend_counter_batch_scope(redis, counter_keys=admission_counter_keys(token, end_user_id=None)): + reservation = await reserve_budget_for_request( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + llm_router=None, + valid_token=token, + team_object=LiteLLM_TeamTableCachedObj(team_id="team", max_budget=20.0), + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + ) + + assert reservation is not None + assert [c.split()[0] for c in redis.commands] == ["MGET", "PIPELINE"], redis.commands + assert set(redis.commands[0].split()[1:]) == {"spend:key:hashed", "spend:team:team"} + assert redis.commands[1] == "PIPELINE spend:key:hashed spend:team:team" + assert redis.store == {"spend:key:hashed": 1.5, "spend:team:team": 2.5} + assert [entry["counter_key"] for entry in reservation["entries"]] == ["spend:key:hashed", "spend:team:team"] + + +@pytest.mark.asyncio +async def test_a_failed_reservation_pipeline_drops_every_counter_and_reserves_nothing(monkeypatch): + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_request + + redis = CountingRedis({"spend:key:hashed": 1.0, "spend:team:team": 2.0}) + redis.default_ttl = 3600 + redis.async_delete_cache = AsyncMock() + + async def _pipeline(increment_list: Sequence[Mapping[str, object]], **kwargs: object) -> list[float]: + raise ConnectionError("redis down") + + redis.async_increment_pipeline = _pipeline # pyright: ignore[reportAttributeAccessIssue] # instance override + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr("litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", lambda **_: 0.5) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + reservation = await reserve_budget_for_request( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + llm_router=None, + valid_token=token, + team_object=LiteLLM_TeamTableCachedObj(team_id="team", max_budget=20.0), + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + ) + + assert reservation is None + assert {call.kwargs["key"] for call in redis.async_delete_cache.await_args_list} == { + "spend:key:hashed", + "spend:team:team", + } + assert redis.store == {"spend:key:hashed": 1.0, "spend:team:team": 2.0} + + +@pytest.mark.asyncio +async def test_post_call_lifecycle_reads_the_counters_after_the_db_update_and_writes_one_pipeline(monkeypatch): + from litellm.proxy.hooks.proxy_track_cost_callback import _update_database_and_spend_counters + + redis = CountingRedis({key: 1.0 for key in POST_CALL_KEYS}) + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + proxy_logging_obj = MagicMock() + + async def _update_database(**kwargs: object) -> bool: + redis.commands.append("DB") + return True + + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=_update_database) + reservation = _reservation(reserved_cost=0.4) + + charged = await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=ps.increment_spend_counters, + user_api_key="hashed", + user_id="user", + end_user_id="eu", + team_id="team", + org_id="org", + kwargs={}, + completion_response=None, + start_time=None, + end_time=None, + response_cost=0.5, + budget_reservation=reservation, + request_tags=["prod"], + model_access_groups=["premium"], + ) + + assert charged is True + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + assert [c.split()[0] for c in redis.commands] == ["MGET", "DB", "MGET", "PIPELINE"], redis.commands + assert set(redis.commands[0].split()[1:]) == RESERVED_KEYS + assert set(redis.commands[2].split()[1:]) == POST_CALL_KEYS + assert set(redis.commands[3].split()[1:]) == POST_CALL_KEYS + assert {key: round(redis.store[key], 6) for key in POST_CALL_KEYS} == { + key: (1.1 if key in RESERVED_KEYS else 1.5) for key in POST_CALL_KEYS + } + assert [round(entry["applied_adjustment"], 6) for entry in reservation["entries"]] == [0.1] * len(RESERVED_KEYS) + assert reservation["finalized"] is True + assert active_spend_counter_batch() is None + + +def _reservation_fixture(monkeypatch, redis: CountingRedis) -> None: + redis.default_ttl = 3600 + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr("litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", lambda **_: 0.5) + + +async def _reserve(redis: CountingRedis, token: UserAPIKeyAuth, team_max_budget: float) -> dict | None: + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_request + + with spend_counter_batch_scope(redis, counter_keys=admission_counter_keys(token, end_user_id=None)): + return await reserve_budget_for_request( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + llm_router=None, + valid_token=token, + team_object=LiteLLM_TeamTableCachedObj(team_id="team", max_budget=team_max_budget), + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + ) + + +@pytest.mark.asyncio +async def test_a_rejected_counter_is_charged_alone_so_the_counters_after_it_are_never_touched(monkeypatch): + """Only counters the admission MGET says still fit the estimate share the reservation pipeline; a counter that + does not is charged on its own first, so its rejection never inflates a sibling counter, not even briefly.""" + redis = CountingRedis({"spend:key:hashed": 10.0, "spend:team:team": 2.0}) + _reservation_fixture(monkeypatch, redis) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + with pytest.raises(litellm.BudgetExceededError): + await _reserve(redis, token, team_max_budget=20.0) + + writes = [c for c in redis.commands if not c.startswith("MGET")] + assert writes and all("spend:team:team" not in c for c in writes), redis.commands + assert redis.store == {"spend:key:hashed": 10.0, "spend:team:team": 2.0} + + +@pytest.mark.asyncio +async def test_a_resized_reservation_is_carried_at_its_resized_cost_to_the_counters_charged_after_it(monkeypatch): + redis = CountingRedis({"spend:key:hashed": 9.8, "spend:team:team": 2.0}) + _reservation_fixture(monkeypatch, redis) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + reservation = await _reserve(redis, token, team_max_budget=2.1) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.1) + assert redis.store["spend:key:hashed"] == pytest.approx(9.9) + assert redis.store["spend:team:team"] == pytest.approx(2.1) + + +@pytest.mark.asyncio +async def test_update_cache_reads_an_object_redis_gained_right_after_a_batch_read_missed_it(monkeypatch): + """DualCache throttles repeated batch reads of a key that just missed; the per-object GET update_cache used to + issue never did, so its batched read must not either.""" + from litellm.caching.dual_cache import DualCache + + redis = CountingRedis() + cache = DualCache(redis_cache=redis) + monkeypatch.setattr(ps, "user_api_key_cache", cache) + assert await cache.async_batch_get_cache(keys=["team_id:team"]) == [None] + redis.store["team_id:team"] = {"spend": 1.0} + assert await cache.async_batch_get_cache(keys=["team_id:team"]) == [None] + + assert await ps._read_update_cache_values(keys=["team_id:team"], parent_otel_span=None) == { + "team_id:team": {"spend": 1.0} + } + assert redis.commands.count("MGET team_id:team") == 2 diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 18b046cd83c..c8e4df1030f 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -1880,8 +1880,8 @@ async def test_should_raise_503_when_counter_increment_fails_and_fail_closed( async def test_fail_closed_releases_earlier_counters_before_503( spend_counter_state, ): - """#33923: when a later counter's reservation write fails in strict mode, the - counters that already reserved must be released before the 503 propagates.""" + """#33923: when a later counter cannot be loaded in strict mode, the 503 is raised before any counter is + reserved.""" counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) valid_token = UserAPIKeyAuth( @@ -1915,12 +1915,8 @@ async def test_fail_closed_releases_earlier_counters_before_503( ) assert exc_info.value.status_code == 503 - assert ( - counter_cache.in_memory_cache.get_cache( - key="spend:key:key-budget-fail-closed-release" - ) - == 0.0 - ) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-fail-closed-release") is None + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-fail-closed-release:window:1h") is None @pytest.mark.asyncio @@ -1982,21 +1978,10 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme max_budget=1.0, ) - import litellm.proxy.proxy_server as ps - - original_increment_counter = ps._increment_spend_counter_cache - first_increment = True - - async def fail_after_increment(counter_key: str, increment: float): - nonlocal first_increment - if first_increment: - first_increment = False - await counter_cache.async_increment_cache(key=counter_key, value=increment) - raise RuntimeError("lost increment response") - return await original_increment_counter( - counter_key=counter_key, - increment=increment, - ) + async def fail_after_increment(pending): + for item in pending: + await counter_cache.async_increment_cache(key=item.counter_key, value=item.increment) + raise RuntimeError("lost increment response") with ( patch( @@ -2004,7 +1989,7 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme return_value=0.5, ), patch( - "litellm.proxy.proxy_server._increment_spend_counter_cache", + "litellm.proxy.proxy_server.run_spend_counter_pipeline", side_effect=fail_after_increment, ), patch( @@ -2596,6 +2581,72 @@ async def test_reconcile_before_db_update_does_not_double_count_when_flush_lands assert reservation["finalized"] is True +class _BatchReadingRedisCache(_ExpiringRedisCache): + async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, float | None]: + return {key: await self.async_get_cache(key) for key in key_list} + + +@pytest.mark.asyncio +async def test_reserved_counter_deleted_during_spend_write_is_reseeded_instead_of_going_negative( + spend_counter_state, +): + import litellm.proxy.proxy_server as ps + from litellm.proxy.hooks.proxy_track_cost_callback import _update_database_and_spend_counters + + counter_cache, _ = spend_counter_state + counter_key = "spend:key:key-deleted-mid-write" + redis_cache = _BatchReadingRedisCache() + counter_cache.redis_cache = redis_cache + await redis_cache.async_set_cache(counter_key, 0.6) + counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.6) + + async def _delete_counter_while_persisting(**kwargs: object) -> bool: + await redis_cache.async_delete_cache(counter_key) + counter_cache.in_memory_cache.delete_cache(key=counter_key) + return True + + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=_delete_counter_while_persisting) + reservation = { + "reserved_cost": 0.6, + "entries": [ + { + "counter_key": counter_key, + "entity_type": "Key", + "entity_id": "key-deleted-mid-write", + "reserved_cost": 0.6, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with ( + patch.object( # test-quality-ok: the reseed reads the DB floor through a Prisma client the test has no seam for + ps.SpendCounterReseed, "from_db", AsyncMock(return_value=0.3) + ) + ): + charged = await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=ps.increment_spend_counters, + user_api_key="key-deleted-mid-write", + user_id=None, + end_user_id=None, + team_id=None, + org_id=None, + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.05, + budget_reservation=reservation, + ) + + assert charged is True + assert redis_cache.store[counter_key] == pytest.approx(0.35), redis_cache.store + assert reservation["finalized"] is True + + @pytest.mark.asyncio async def test_should_invalidate_reserved_counters_after_persisted_spend_failure( spend_counter_state, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index df05cf0987e..23d319159ca 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -6189,7 +6189,7 @@ async def test_tag_cache_update_called(): "spend": 10.0, } - with patch.object(cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj)) as mock_get_cache: + with patch.object(cache, "async_batch_get_cache", new=AsyncMock(return_value=[mock_tag_obj])) as mock_get_cache: with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, @@ -6203,7 +6203,7 @@ async def test_tag_cache_update_called(): await asyncio.sleep(0.1) - mock_get_cache.assert_awaited_once_with(key="tag:test-tag") + mock_get_cache.assert_awaited_once_with(keys=["tag:test-tag"], parent_otel_span=None, throttle_redis=False) mock_set_cache.assert_awaited_once() call_args = mock_set_cache.call_args @@ -6234,15 +6234,11 @@ async def test_tag_cache_update_multiple_tags(): mock_tag1_obj = {"tag_name": "tag1", "spend": 10.0} mock_tag2_obj = {"tag_name": "tag2", "spend": 20.0} - async def mock_get_cache_side_effect(key): - if key == "tag:tag1": - return mock_tag1_obj - elif key == "tag:tag2": - return mock_tag2_obj - return None + async def mock_get_cache_side_effect(keys, **kwargs): + return [{"tag:tag1": mock_tag1_obj, "tag:tag2": mock_tag2_obj}.get(key) for key in keys] with patch.object( - cache, "async_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect) + cache, "async_batch_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect) ) as mock_get_cache: with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( @@ -6257,7 +6253,7 @@ async def test_tag_cache_update_multiple_tags(): await asyncio.sleep(0.1) - assert mock_get_cache.call_count == 2 + mock_get_cache.assert_awaited_once_with(keys=["tag:tag1", "tag:tag2"], parent_otel_span=None, throttle_redis=False) mock_set_cache.assert_awaited_once() call_args = mock_set_cache.call_args @@ -6288,8 +6284,8 @@ async def test_update_cache_pipeline_honors_user_api_key_cache_ttl(): try: with patch.object( cache, - "async_get_cache", - new=AsyncMock(return_value={"tag_name": "active-tag", "spend": 1.0}), + "async_batch_get_cache", + new=AsyncMock(return_value=[{"tag_name": "active-tag", "spend": 1.0}]), ): with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( @@ -6376,18 +6372,21 @@ async def test_update_cache_global_proxy_spend_scalar_stays_shared(): admin_name = litellm.proxy.proxy_server.litellm_proxy_admin_name global_key = "{}:spend".format(admin_name) - async def fake_get(key, **kwargs): + def fake_get(key): if key == "user-lit": return {"user_id": "user-lit", "spend": 1.0} if key == global_key: return 10.0 return None + async def fake_batch_get(keys, **kwargs): + return [fake_get(key) for key in keys] + original_cache = litellm.proxy.proxy_server.user_api_key_cache cache = DualCache(default_in_memory_ttl=300) setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache) try: - with patch.object(cache, "async_get_cache", new=AsyncMock(side_effect=fake_get)): + with patch.object(cache, "async_batch_get_cache", new=AsyncMock(side_effect=fake_batch_get)): with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, @@ -13868,7 +13867,7 @@ async def test_window_spend_row_is_enqueued_even_when_the_counter_was_reserved() } original_reconcile = br.reconcile_budget_reservation - br.reconcile_budget_reservation = AsyncMock(return_value=None) + br.reconcile_budget_reservation = AsyncMock(return_value=()) try: with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue: await increment_spend_counters( diff --git a/tests/unit/proxy/auth/test_jwt.py b/tests/unit/proxy/auth/test_jwt.py index 6ad253f33e8..fd1d8974b48 100644 --- a/tests/unit/proxy/auth/test_jwt.py +++ b/tests/unit/proxy/auth/test_jwt.py @@ -874,8 +874,7 @@ async def test_team_cache_update_called(): cache, ) - with patch.object(cache, "async_get_cache", new=AsyncMock()) as mock_call_cache: - cache.async_get_cache = mock_call_cache + with patch.object(cache, "async_batch_get_cache", new=AsyncMock(return_value=[None])) as mock_call_cache: # Call the function under test await litellm.proxy.proxy_server.update_cache( token=None, @@ -887,7 +886,7 @@ async def test_team_cache_update_called(): ) # type: ignore await asyncio.sleep(3) - mock_call_cache.assert_awaited_once() + mock_call_cache.assert_awaited_once_with(keys=["team_id:1234"], parent_otel_span=None, throttle_redis=False) @pytest.fixture