mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
d2a574b791
commit
2d034bb35b
14 changed files with 874 additions and 247 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue