perf(proxy): hold one spend counter batch across admission and across post-call accounting (#43369)

Auth's spend counter MGET scope spans common checks, model budget check and reservation;
reservation increments go out as one pipeline; post-call reconcile adjustments ride the ordinary
increment pipeline and update_cache uses one batched read. Over-budget reservation counters are
charged one at a time so a rejection never touches the counters after it; post-call counter keys are
derived from ids without validating a UserAPIKeyAuth.

Resolves LIT-8881

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 15:05:33 -07:00 • committed by GitHub
parent d2a574b791
commit 2d034bb35b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 874 additions and 247 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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