Merge pull request #35278 from BerriAI/litellm_v3_limiter_contextvar_stash

refactor(rate-limits): move the v3 limiter per-request stash off request metadata onto a ContextVar
This commit is contained in:
Mateo Wang 2026-07-30 18:13:38 -07:00 • committed by GitHub
commit f0d13624be
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 516 additions and 743 deletions

View file

@ -2801,9 +2801,7 @@ class ProxyBaseLLMRequestProcessing:
and proxy_logging_obj is not None
and user_api_key_dict is not None
):
await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(
user_api_key_dict, request_data
)
await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(user_api_key_dict)
if hasattr(response, "aclose"):
try:

View file

@ -21,7 +21,10 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import (
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
RateLimitDescriptor,
RateLimitDescriptorRateLimitObject,
RateLimitResponse,
_PROXY_MaxParallelRequestsHandler_v3,
claim_request_stash_for_data,
get_or_create_request_stash,
)
from litellm.proxy.hooks.rate_limiter_utils import (
convert_priority_to_percent,
@ -373,7 +376,6 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
user_api_key_dict: UserAPIKeyAuth,
priority: Optional[str],
saturation: float,
data: dict,
) -> None:
"""
Check rate limits using THREE-PHASE approach to prevent partial increments.
@ -400,7 +402,6 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
user_api_key_dict: User authentication info
priority: User's priority level
saturation: Current saturation level
data: Request data dictionary
Raises:
HTTPException: If any limit is exceeded
@ -550,12 +551,12 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
parent_otel_span=user_api_key_dict.parent_otel_span,
read_only=False,
)
data["litellm_proxy_rate_limit_response"] = {
"overall_code": atomic_response["overall_code"],
"statuses": atomic_response["statuses"] + priority_tracking_response["statuses"],
}
get_or_create_request_stash().rate_limit_response = RateLimitResponse(
overall_code=atomic_response["overall_code"],
statuses=atomic_response["statuses"] + priority_tracking_response["statuses"],
)
else:
data["litellm_proxy_rate_limit_response"] = atomic_response
get_or_create_request_stash().rate_limit_response = atomic_response
async def async_pre_call_hook(
self,
@ -601,6 +602,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
if "model" not in data:
return None
claim_request_stash_for_data(data)
model = data["model"]
priority = self._get_priority_from_user_api_key_dict(user_api_key_dict=user_api_key_dict)
@ -632,7 +634,6 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
user_api_key_dict=user_api_key_dict,
priority=priority,
saturation=saturation,
data=data,
)
except HTTPException:

View file

@ -8,16 +8,18 @@ import asyncio
import binascii
import os
import uuid
from contextvars import ContextVar
from dataclasses import dataclass, field
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
FrozenSet,
List,
Literal,
Optional,
Set,
Tuple,
TypedDict,
Union,
@ -28,7 +30,6 @@ from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
@ -291,53 +292,11 @@ DEFAULT_CHARS_PER_TOKEN = 4
# (baseline floor) and to the smallest configured TPM limit (capped floor for
# small per-tenant TPM caps).
_TPM_FLOOR_FRACTION = 4
# Stash for the reserved-token count on the request data dict so success/
# failure callbacks can reconcile against the upfront reservation.
TPM_RESERVED_TOKENS_KEY = "_litellm_tpm_reserved_tokens"
# Stash for the model identifier the reservation was charged against.
# Reconciliation must target the same key that was incremented at reservation
TPM_RESERVED_MODEL_KEY = "_litellm_tpm_reserved_model"
# Stash for the (scope_key, scope_value) pairs whose :tokens counter the
# upfront reservation incremented. Reconciliation applies the delta to these
# scopes only; scopes without a configured TPM limit were never charged at
# pre-call and must receive the full actual usage instead of the delta —
# otherwise their counters drift negative whenever actual < reserved.
TPM_RESERVED_SCOPES_KEY = "_litellm_tpm_reserved_scopes"
# Idempotency marker for the reservation refund path. Set when any failure
# callback releases the reservation so the next callback in the same flow
# (e.g. async_log_failure_event firing after async_post_call_failure_hook)
# does not double-refund.
TPM_RESERVATION_RELEASED_KEY = "_litellm_tpm_reservation_released"
RATE_LIMIT_DESCRIPTORS_KEY = "_litellm_rate_limit_descriptors"
# Pre-call RateLimitResponse stashed here so streaming success logging can
# mirror ``x-ratelimit-*`` headers into the SLP. Streaming exits
# common_request_processing before ``async_post_call_success_hook`` runs.
RATE_LIMIT_RESPONSE_KEY = "_litellm_proxy_rate_limit_response"
# Holds the acquisition the pre-call hook made for this request: the slot id
# plus the gauge counter keys it was registered under. The success/failure
# callbacks release only this exact acquisition: those callbacks also fire
# for requests rejected at pre-call (which never acquired a slot), and an
# id-less release would free a slot still owned by another in-flight request
# — every rejection would then raise effective concurrency above the
# configured limit.
MAX_PARALLEL_SLOT_ACQUIRED_KEY = "_litellm_max_parallel_slot_acquired"
# How long an acquired slot counts toward the in-flight total before it is
# considered leaked (worker crashed without any release callback firing) and
# pruned. Also the longest request duration the gauge can track: a request
# running longer than this stops occupying its slot.
PARALLEL_REQUEST_SLOT_TTL_SECONDS = 3600
# Stash keys live ONLY in metadata channels — never at the top level of the
# request body. Top-level keys are forwarded as body params to upstream
# providers, which reject unknown fields with 400/429 errors.
_LITELLM_STASH_KEYS: Tuple[str, ...] = (
TPM_RESERVED_TOKENS_KEY,
TPM_RESERVED_MODEL_KEY,
TPM_RESERVED_SCOPES_KEY,
TPM_RESERVATION_RELEASED_KEY,
RATE_LIMIT_DESCRIPTORS_KEY,
RATE_LIMIT_RESPONSE_KEY,
MAX_PARALLEL_SLOT_ACQUIRED_KEY,
)
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
@ -382,6 +341,79 @@ class RateLimitResponseWithDescriptors(TypedDict):
response: RateLimitResponse
@dataclass(slots=True)
class RequestRateLimiterStash:
"""
Per-request bookkeeping the pre-call hook hands to the success/failure/
disconnect callbacks. Lives on a ContextVar instead of the request body so
it never reaches provider-facing ``metadata`` channels.
A single mutable instance is shared by every context forked from the
request task (the SDK call, streaming generators, and the logging worker's
captured context all see the same object), which is what makes the
``reservation_released`` flag and ``parallel_slot`` clearing effective
across sibling callbacks: the first release wins, later callbacks observe
the cleared state.
Because the stash is context-inherited, nested LiteLLM calls made inside
the request (LLM-judge guardrails, silent experiments) would also see it
from their own logging callbacks. ``owner_litellm_call_id`` pins the stash
to the proxy request's ``litellm_call_id`` so those callbacks can tell the
owning request's events apart from a nested call's: router retries and
fallbacks reuse the request's call id and keep access, while nested calls
mint fresh ids and are ignored.
"""
owner_litellm_call_id: Optional[str] = None
rate_limit_response: Optional[RateLimitResponse] = None
parallel_slot: Optional[ParallelSlotAcquisition] = None
reserved_tokens: int = 0
reserved_model: Optional[str] = None
reserved_scopes: FrozenSet[Tuple[str, str]] = field(default_factory=frozenset)
reservation_released: bool = False
_request_stash: ContextVar[Optional[RequestRateLimiterStash]] = ContextVar(
"litellm_v3_rate_limiter_request_stash", default=None
)
def get_request_stash() -> Optional[RequestRateLimiterStash]:
return _request_stash.get()
def get_or_create_request_stash() -> RequestRateLimiterStash:
stash = _request_stash.get()
if stash is None:
stash = RequestRateLimiterStash()
_request_stash.set(stash)
return stash
def claim_request_stash_for_data(data: dict) -> RequestRateLimiterStash:
stash = get_or_create_request_stash()
owner_call_id = data.get("litellm_call_id")
if isinstance(owner_call_id, str):
stash.owner_litellm_call_id = owner_call_id
return stash
def get_request_stash_for_call(litellm_call_id: Optional[str]) -> Optional[RequestRateLimiterStash]:
stash = _request_stash.get()
if stash is None:
return None
if stash.owner_litellm_call_id is None or litellm_call_id is None:
return stash
return stash if litellm_call_id == stash.owner_litellm_call_id else None
def _call_id_from_callback_kwargs(kwargs: object) -> Optional[str]:
if not isinstance(kwargs, dict):
return None
call_id = kwargs.get("litellm_call_id")
return call_id if isinstance(call_id, str) else None
class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
def __init__(
self,
@ -2343,12 +2375,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
verbose_proxy_logger.debug("Inside Rate Limit Pre-Call Hook")
# Reject caller-supplied stash values before any read/write. Otherwise
# a client can inject ``_litellm_rate_limit_descriptors`` /
# ``_litellm_tpm_reserved_tokens`` in body ``metadata`` and have
# ``async_post_call_failure_hook`` refund TPM counters against scopes
# they name (e.g. another tenant's api_key).
self._strip_stash_keys_from_all_channels(data)
stash = claim_request_stash_for_data(data)
#########################################################
# Check if the call type has a specific rate limiter
@ -2444,23 +2471,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
requested_model=requested_model,
)
else:
# add descriptors to request headers
data["litellm_proxy_rate_limit_response"] = response
# Mirror into metadata so streaming success logging can find
# it via ``kwargs["litellm_params"]["metadata"]``.
self._stash_value_in_internal_metadata(
data=data,
key=RATE_LIMIT_RESPONSE_KEY,
value=response,
)
stash.rate_limit_response = response
if parallel_slot_id is not None:
self._stash_value_in_internal_metadata(
data=data,
key=MAX_PARALLEL_SLOT_ACQUIRED_KEY,
value={
"slot_id": parallel_slot_id,
"counter_keys": parallel_counter_keys,
},
stash.parallel_slot = ParallelSlotAcquisition(
slot_id=parallel_slot_id,
counter_keys=parallel_counter_keys,
)
# ----------------------------------------------------------------
@ -2521,38 +2536,29 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
if tpm_response["overall_code"] == "OVER_LIMIT":
acquisition = self._get_parallel_slot_acquisition(kwargs=data)
acquisition = stash.parallel_slot
if acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
self._clear_parallel_slot_marker(data)
stash.parallel_slot = None
self._handle_rate_limit_error(
response=tpm_response,
descriptors=descriptors,
requested_model=requested_model,
)
else:
self._stash_value_in_internal_metadata(
data=data,
key=RATE_LIMIT_DESCRIPTORS_KEY,
value=descriptors,
)
# Capture the exact (key, value) scopes the reservation
# incremented so post-call reconciliation only applies
# the (actual - reserved) delta to those — unreserved
# scopes get charged the full actual usage instead.
reserved_scopes: List[Tuple[str, str]] = [
stash.reserved_tokens = estimated_tokens
stash.reserved_model = requested_model
stash.reserved_scopes = frozenset(
(d["key"], d["value"])
for d in descriptors
if (d.get("rate_limit") or {}).get("tokens_per_unit") is not None
]
self._stash_reservation_in_data(
data=data,
estimated_tokens=estimated_tokens,
reserved_model=requested_model,
reserved_scopes=reserved_scopes,
)
# Merge TPM statuses into the stored rate-limit response
@ -2560,44 +2566,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# headers reach the client. Without this, the RPM-only
# response from should_rate_limit (skip_tpm_check=True)
# silently drops all token headers.
stored_response = data.get("litellm_proxy_rate_limit_response")
if isinstance(stored_response, dict):
stored_response.setdefault("statuses", []).extend(tpm_response["statuses"])
elif tpm_response["statuses"]:
data["litellm_proxy_rate_limit_response"] = tpm_response
# Keep the metadata stash in sync when this is the
# first snapshot written.
self._stash_value_in_internal_metadata(
data=data,
key=RATE_LIMIT_RESPONSE_KEY,
value=tpm_response,
)
stored_response = stash.rate_limit_response
if stored_response is not None:
stored_response["statuses"].extend(tpm_response["statuses"])
verbose_proxy_logger.debug(f"TPM tokens reserved: {estimated_tokens} for model {requested_model}")
# Defense-in-depth: scrub any stash key that escaped onto data
# top-level (stale cache hit, router pass, test fixture) before the
# body is forwarded to the provider.
self._strip_stash_keys_from_top_level(data)
@staticmethod
def _strip_stash_keys_from_top_level(data: Any) -> None:
if not isinstance(data, dict):
return
for stash_key in _LITELLM_STASH_KEYS:
data.pop(stash_key, None)
@classmethod
def _strip_stash_keys_from_all_channels(cls, data: Any) -> None:
if not isinstance(data, dict):
return
cls._strip_stash_keys_from_top_level(data)
for channel in ("metadata", "litellm_metadata"):
channel_dict = data.get(channel)
if isinstance(channel_dict, dict):
for stash_key in _LITELLM_STASH_KEYS:
channel_dict.pop(stash_key, None)
def _create_pipeline_operations(
self,
key: str,
@ -2803,202 +2777,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
merged[f"{prefix}-limit-{status['rate_limit_type']}"] = status["current_limit"]
return merged
@staticmethod
def _stash_value_in_internal_metadata(
data: Dict[str, Any],
key: str,
value: Any,
) -> None:
# Writes only the proxy-internal bucket. Routes that own
# ``litellm_metadata`` (Responses, /v1/messages, batches, files) expose
# ``metadata`` as a provider request parameter, so creating or adding to
# it here would forward internal state upstream.
_, metadata_bucket = get_or_create_metadata_bucket(data)
metadata_bucket[key] = value
@classmethod
def _stash_reservation_in_data(
cls,
data: Dict[str, Any],
estimated_tokens: int,
reserved_model: Optional[str],
reserved_scopes: Optional[List[Tuple[str, str]]] = None,
) -> None:
"""
``reserved_scopes`` is serialized as a list of [key, value] pairs so
it round-trips through JSON-based metadata transports.
"""
scopes_payload: Optional[List[List[str]]] = [[k, v] for k, v in reserved_scopes] if reserved_scopes else None
cls._stash_value_in_internal_metadata(data=data, key=TPM_RESERVED_TOKENS_KEY, value=estimated_tokens)
if reserved_model:
cls._stash_value_in_internal_metadata(data=data, key=TPM_RESERVED_MODEL_KEY, value=reserved_model)
if scopes_payload is not None:
cls._stash_value_in_internal_metadata(data=data, key=TPM_RESERVED_SCOPES_KEY, value=scopes_payload)
@staticmethod
def _lookup_stashed_value(
kwargs: Any,
standard_logging_metadata: Optional[Dict[str, Any]],
key: str,
) -> Any:
"""
Resolve a stashed value from any metadata channel the request data
can flow through to a callback. Top-level ``kwargs`` is not checked
because stash keys must never live there.
"""
candidate: Any = None
if isinstance(kwargs, dict):
for channel in ("metadata", "litellm_metadata"):
channel_dict = kwargs.get(channel)
if isinstance(channel_dict, dict) and key in channel_dict:
candidate = channel_dict.get(key)
if candidate is not None:
return candidate
litellm_params = kwargs.get("litellm_params")
if isinstance(litellm_params, dict):
for channel in ("litellm_metadata", "metadata"):
lp_metadata = litellm_params.get(channel)
if isinstance(lp_metadata, dict) and lp_metadata.get(key) is not None:
return lp_metadata[key]
if candidate is None and isinstance(standard_logging_metadata, dict):
candidate = standard_logging_metadata.get(key)
return candidate
@classmethod
def _get_reserved_tokens_from_kwargs(
cls,
kwargs: Any,
standard_logging_metadata: Optional[Dict[str, Any]] = None,
) -> int:
candidate = cls._lookup_stashed_value(kwargs, standard_logging_metadata, TPM_RESERVED_TOKENS_KEY)
try:
return int(candidate or 0)
except (TypeError, ValueError):
return 0
@classmethod
def _get_reserved_model_from_kwargs(
cls,
kwargs: Any,
standard_logging_metadata: Optional[Dict[str, Any]] = None,
) -> Optional[str]:
"""
Resolve the model the upfront reservation was charged against. Used to
target reconciliation at the same key that was incremented, regardless
of whether the router later set a different ``model_group`` in
``litellm_params.metadata``.
"""
candidate = cls._lookup_stashed_value(kwargs, standard_logging_metadata, TPM_RESERVED_MODEL_KEY)
return candidate if isinstance(candidate, str) and candidate else None
@classmethod
def _get_reserved_scopes_from_kwargs(
cls,
kwargs: Any,
standard_logging_metadata: Optional[Dict[str, Any]] = None,
) -> Set[Tuple[str, str]]:
"""
Resolve the (scope_key, scope_value) pairs the upfront reservation
actually charged. Reconciliation distinguishes these from
unreserved scopes — applying the delta to reserved scopes (which
already carry +reserved on the counter) and the full actual to
unreserved ones (which were never charged).
"""
candidate = cls._lookup_stashed_value(kwargs, standard_logging_metadata, TPM_RESERVED_SCOPES_KEY)
if not isinstance(candidate, list):
return set()
scopes: Set[Tuple[str, str]] = set()
for entry in candidate:
if (
isinstance(entry, (list, tuple))
and len(entry) == 2
and isinstance(entry[0], str)
and isinstance(entry[1], str)
):
scopes.add((entry[0], entry[1]))
return scopes
@classmethod
def _is_reservation_released(
cls,
kwargs: Any,
standard_logging_metadata: Optional[Dict[str, Any]] = None,
) -> bool:
"""True if a prior callback already refunded this request's reservation."""
return bool(cls._lookup_stashed_value(kwargs, standard_logging_metadata, TPM_RESERVATION_RELEASED_KEY))
@classmethod
def _get_parallel_slot_acquisition(
cls,
kwargs: Any,
standard_logging_metadata: dict[str, Any] | None = None,
) -> ParallelSlotAcquisition | None:
"""The slot acquisition this request's pre-call hook made, if any."""
candidate = cls._lookup_stashed_value(kwargs, standard_logging_metadata, MAX_PARALLEL_SLOT_ACQUIRED_KEY)
if not isinstance(candidate, dict):
return None
slot_id = candidate.get("slot_id")
counter_keys = candidate.get("counter_keys")
if not isinstance(slot_id, str) or not slot_id:
return None
if not isinstance(counter_keys, list) or not counter_keys:
return None
if not all(isinstance(key, str) and key for key in counter_keys):
return None
return ParallelSlotAcquisition(slot_id=slot_id, counter_keys=counter_keys)
@staticmethod
def _clear_parallel_slot_marker(data: Any) -> None:
"""
Remove the acquired-slot marker from every metadata channel a sibling
callback might read, so one release per acquire is an invariant even
when multiple callbacks fire for the same request.
"""
if not isinstance(data, dict):
return
for channel in ("metadata", "litellm_metadata"):
channel_dict = data.get(channel)
if isinstance(channel_dict, dict):
channel_dict.pop(MAX_PARALLEL_SLOT_ACQUIRED_KEY, None)
litellm_params = data.get("litellm_params")
if isinstance(litellm_params, dict):
lp_metadata = litellm_params.get("metadata")
if isinstance(lp_metadata, dict):
lp_metadata.pop(MAX_PARALLEL_SLOT_ACQUIRED_KEY, None)
slo = data.get("standard_logging_object")
if isinstance(slo, dict):
slo_meta = slo.get("metadata")
if isinstance(slo_meta, dict):
slo_meta.pop(MAX_PARALLEL_SLOT_ACQUIRED_KEY, None)
@staticmethod
def _mark_reservation_released(data: Any) -> None:
"""
Stamp the released flag into every metadata channel a sibling
callback might read from. async_post_call_failure_hook receives the
request data dict; async_log_failure_event reads kwargs +
standard_logging_object.metadata. Same dict identity across
``request_data["metadata"]`` and ``kwargs["litellm_params"]["metadata"]``
means writes here propagate to the other hook.
"""
if not isinstance(data, dict):
return
for channel in ("metadata", "litellm_metadata"):
existing = data.get(channel)
if isinstance(existing, dict):
existing[TPM_RESERVATION_RELEASED_KEY] = True
litellm_params = data.get("litellm_params")
if isinstance(litellm_params, dict):
lp_metadata = litellm_params.get("metadata")
if isinstance(lp_metadata, dict):
lp_metadata[TPM_RESERVATION_RELEASED_KEY] = True
slo = data.get("standard_logging_object")
if isinstance(slo, dict):
slo_meta = slo.get("metadata")
if isinstance(slo_meta, dict):
slo_meta[TPM_RESERVATION_RELEASED_KEY] = True
def _collect_tpm_scope_targets(
self,
standard_logging_metadata: Dict[str, Any],
@ -3064,7 +2842,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
def _build_reservation_aware_tpm_ops(
self,
targets: List[Tuple[str, str]],
reserved_scopes: Set[Tuple[str, str]],
reserved_scopes: FrozenSet[Tuple[str, str]],
actual_tokens: int,
reserved_tokens: int,
) -> List[RedisPipelineIncrementOperation]:
@ -3139,18 +2917,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if total_tokens == 0:
total_tokens = self._aggregate_only_total_tokens(usage=_usage)
reserved_tokens = self._get_reserved_tokens_from_kwargs(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
reserved_model = self._get_reserved_model_from_kwargs(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
reserved_scopes = self._get_reserved_scopes_from_kwargs(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
reserved_tokens = stash.reserved_tokens if stash is not None else 0
reserved_model = stash.reserved_model if stash is not None else None
reserved_scopes: FrozenSet[Tuple[str, str]] = stash.reserved_scopes if stash is not None else frozenset()
# Reconciliation must target the same model-scoped counter that the
# pre-call reservation incremented. If a reservation was made,
# ``reserved_model`` is authoritative; otherwise fall back to the
@ -3206,18 +2976,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
try:
verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING")
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
acquisition = self._get_parallel_slot_acquisition(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
if acquisition is not None:
stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
acquisition = stash.parallel_slot if stash is not None else None
if stash is not None and acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=litellm_parent_otel_span,
)
self._clear_parallel_slot_marker(kwargs)
stash.parallel_slot = None
pipeline_operations = self._build_success_event_pipeline_operations(
kwargs=kwargs,
@ -3267,23 +3033,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if not isinstance(kwargs, dict):
return
standard_logging_object = kwargs.get("standard_logging_object")
standard_logging_metadata: Optional[Dict[str, Any]] = None
if isinstance(standard_logging_object, dict):
slp_metadata = standard_logging_object.get("metadata")
if isinstance(slp_metadata, dict):
standard_logging_metadata = slp_metadata
statuses = self._narrow_ratelimit_statuses(
self._lookup_stashed_value(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
key=RATE_LIMIT_RESPONSE_KEY,
)
)
stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
rate_limit_response = stash.rate_limit_response if stash is not None else None
statuses = rate_limit_response["statuses"] if rate_limit_response is not None else []
if not statuses:
return
standard_logging_object = kwargs.get("standard_logging_object")
if isinstance(standard_logging_object, dict):
hidden_params = standard_logging_object.get("hidden_params")
if not isinstance(hidden_params, dict):
@ -3303,43 +3059,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
statuses=statuses,
)
@staticmethod
def _narrow_ratelimit_statuses(stashed: Any) -> List[RateLimitStatus]:
"""
Narrow a stashed ``RateLimitResponse``-shaped dict to a typed
``statuses`` list. Entries missing any header-write field are dropped;
an empty list means "nothing to mirror".
"""
if not isinstance(stashed, dict):
return []
raw_statuses = stashed.get("statuses")
if not isinstance(raw_statuses, list):
return []
narrowed: List[RateLimitStatus] = []
for entry in raw_statuses:
if not isinstance(entry, dict):
continue
descriptor_key = entry.get("descriptor_key")
rate_limit_type = entry.get("rate_limit_type")
current_limit = entry.get("current_limit")
limit_remaining = entry.get("limit_remaining")
if (
isinstance(descriptor_key, str)
and rate_limit_type in ("requests", "tokens", "max_parallel_requests")
and isinstance(current_limit, int)
and isinstance(limit_remaining, int)
):
narrowed.append(
RateLimitStatus(
code=entry.get("code", "OK") if isinstance(entry.get("code"), str) else "OK",
current_limit=current_limit,
limit_remaining=limit_remaining,
rate_limit_type=rate_limit_type,
descriptor_key=descriptor_key,
)
)
return narrowed
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""
On failure: decrement max_parallel_requests and refund the upfront
@ -3353,55 +3072,36 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
try:
litellm_parent_otel_span: Union[Span, None] = _get_parent_otel_span_from_kwargs(kwargs)
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
pipeline_operations: List[RedisPipelineIncrementOperation] = []
acquisition = self._get_parallel_slot_acquisition(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
if acquisition is not None:
stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
acquisition = stash.parallel_slot if stash is not None else None
if stash is not None and acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=litellm_parent_otel_span,
)
self._clear_parallel_slot_marker(kwargs)
stash.parallel_slot = None
# Skip the reservation refund if async_post_call_failure_hook
# already released it (proxy-level rejection that also bubbles up
# here as an LLM-error callback). max_parallel_requests is its
# own counter and is always decremented per call.
already_released = self._is_reservation_released(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
reserved_tokens = (
0
if already_released
else self._get_reserved_tokens_from_kwargs(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
)
if reserved_tokens > 0:
reserved_tokens = 0
if stash is not None and not stash.reservation_released:
reserved_tokens = stash.reserved_tokens
if stash is not None and reserved_tokens > 0:
verbose_proxy_logger.debug(f"Releasing reserved TPM tokens on failure: {reserved_tokens}")
# Refund only against the scopes the reservation actually
# charged. _build_reservation_aware_tpm_ops with
# actual_tokens=0 emits -reserved on reserved scopes and 0
# on unreserved (skipped), so unreserved scopes can't drift
# negative. Targets are derived purely from the reserved
# set so we don't even need to re-collect them from
# metadata.
reserved_scopes = self._get_reserved_scopes_from_kwargs(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
# negative.
pipeline_operations.extend(
self._build_reservation_aware_tpm_ops(
targets=list(reserved_scopes),
reserved_scopes=reserved_scopes,
targets=list(stash.reserved_scopes),
reserved_scopes=stash.reserved_scopes,
actual_tokens=0,
reserved_tokens=reserved_tokens,
)
@ -3412,15 +3112,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
increment_list=pipeline_operations,
litellm_parent_otel_span=litellm_parent_otel_span,
)
if reserved_tokens > 0:
self._mark_reservation_released(kwargs)
if stash is not None and reserved_tokens > 0:
stash.reservation_released = True
except Exception as e:
verbose_proxy_logger.exception(f"Error in rate limit failure event: {str(e)}")
async def async_release_max_parallel_requests_on_disconnect(
self,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict | None = None,
) -> None:
"""
Release the api-key ``max_parallel_requests`` slot that
@ -3432,20 +3131,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
client cancels a stream mid-flight, the cancellation surfaces as
``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback
runs, so without this the slot leaks per cancelled stream until its
TTL prunes it. ``request_data`` carries the stashed acquisition;
its presence (not the key object's current max_parallel_requests
configuration, which can change mid-request) decides whether there
is anything to release.
TTL prunes it. The stashed acquisition's presence (not the key
object's current max_parallel_requests configuration, which can
change mid-request) decides whether there is anything to release.
"""
acquisition = self._get_parallel_slot_acquisition(kwargs=request_data)
if acquisition is None:
stash = get_request_stash()
if stash is None or stash.parallel_slot is None:
return
await self._release_parallel_request_slots(
acquisition=acquisition,
acquisition=stash.parallel_slot,
parent_otel_span=None,
)
self._clear_parallel_slot_marker(request_data)
stash.parallel_slot = None
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
"""
@ -3454,10 +3152,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
try:
from pydantic import BaseModel
litellm_proxy_rate_limit_response = cast(
Optional[RateLimitResponse],
data.get("litellm_proxy_rate_limit_response", None),
)
stash = get_request_stash()
litellm_proxy_rate_limit_response = stash.rate_limit_response if stash is not None else None
if litellm_proxy_rate_limit_response is not None:
# Update response headers
@ -3502,59 +3198,42 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
rejections, so a leaked slot would occupy the gauge for the full
PARALLEL_REQUEST_SLOT_TTL_SECONDS.
Idempotent: the slot release clears the acquisition marker (and slot
Idempotent: the slot release clears the stashed acquisition (and slot
removal is a no-op ZREM on a second run), and the TPM refund is
guarded by TPM_RESERVATION_RELEASED_KEY — if both this hook and
async_log_failure_event end up running in the same flow, only the
first release/refund applies.
guarded by the stash's ``reservation_released`` flag — if both this
hook and async_log_failure_event end up running in the same flow, only
the first release/refund applies.
"""
try:
acquisition = self._get_parallel_slot_acquisition(kwargs=request_data)
if acquisition is not None:
stash = get_request_stash()
if stash is None:
return
if stash.parallel_slot is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
acquisition=stash.parallel_slot,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
self._clear_parallel_slot_marker(request_data)
stash.parallel_slot = None
if self._is_reservation_released(kwargs=request_data):
if stash.reservation_released:
return
reserved_tokens = self._get_reserved_tokens_from_kwargs(kwargs=request_data)
reserved_tokens = stash.reserved_tokens
if reserved_tokens <= 0:
return
# Refund directly against the descriptors we reserved against —
# the pre-call hook stashes them in the request-data metadata
# channels before success/failure callbacks run.
stashed = self._lookup_stashed_value(
kwargs=request_data,
standard_logging_metadata=None,
key=RATE_LIMIT_DESCRIPTORS_KEY,
ops = self._build_reservation_aware_tpm_ops(
targets=list(stash.reserved_scopes),
reserved_scopes=stash.reserved_scopes,
actual_tokens=0,
reserved_tokens=reserved_tokens,
)
descriptors: List[RateLimitDescriptor] = stashed if isinstance(stashed, list) else []
ops: List[RedisPipelineIncrementOperation] = []
for descriptor in descriptors:
rate_limit = descriptor.get("rate_limit") or {}
if rate_limit.get("tokens_per_unit") is None:
continue
ops.append(
RedisPipelineIncrementOperation(
key=self.create_rate_limit_keys(
descriptor["key"],
descriptor["value"],
"tokens",
),
increment_value=-reserved_tokens,
ttl=self.window_size,
)
)
if ops:
verbose_proxy_logger.debug(f"Releasing reserved TPM tokens on proxy-level rejection: {reserved_tokens}")
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=ops,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
self._mark_reservation_released(request_data)
stash.reservation_released = True
except Exception as e:
verbose_proxy_logger.exception(f"Error releasing TPM reservation on post-call failure: {e}")
return None

View file

@ -2814,7 +2814,6 @@ class ProxyLogging:
async def _arelease_max_parallel_requests_on_disconnect(
self,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict | None = None,
) -> None:
"""
Release the api-key max_parallel_requests slot when a streaming
@ -2834,7 +2833,7 @@ class ProxyLogging:
limiter = self.get_proxy_hook("parallel_request_limiter")
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
return
await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict, request_data)
await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict)
def _init_response_taking_too_long_task(self, data: Optional[dict] = None):
"""

View file

@ -3280,7 +3280,6 @@ all_litellm_params = (
"mock_response",
"mock_timeout",
"disable_add_transform_inline_image_block",
"litellm_proxy_rate_limit_response",
"api_key",
"api_version",
"prompt_id",
@ -3375,11 +3374,6 @@ all_litellm_params = (
"enable_tag_filtering",
"enable_json_schema_validation",
"use_xai_oauth",
"_litellm_rate_limit_descriptors",
"_litellm_tpm_reserved_tokens",
"_litellm_tpm_reserved_model",
"_litellm_tpm_reserved_scopes",
"_litellm_tpm_reservation_released",
"auto_router_config_path",
"auto_router_config",
"auto_router_default_model",

View file

@ -1754,7 +1754,6 @@ async def test_priority_429_includes_model_name_and_configured_limits():
user_api_key_dict=user,
priority="prod",
saturation=0.95,
data={"model": model},
)
assert exc_info.value.status_code == 429

View file

@ -18,8 +18,12 @@ from litellm import Router
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
MAX_PARALLEL_SLOT_ACQUIRED_KEY,
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
ParallelSlotAcquisition,
RequestRateLimiterStash,
_request_stash,
get_or_create_request_stash,
get_request_stash,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
@ -52,6 +56,13 @@ def time_controller(monkeypatch):
return controller
@pytest.fixture(autouse=True)
def _isolated_request_stash():
token = _request_stash.set(None)
yield
_request_stash.reset(token)
@pytest.mark.parametrize(
"throttle_pct, expected_rpm, expected_tpm",
[
@ -673,35 +684,36 @@ async def test_async_log_failure_event_v3():
await _seed_max_parallel_requests_slots(local_cache, counter_key, ["slot-a", "slot-b"])
def kwargs_with_slot(slot_id):
return {
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": slot_id,
"counter_keys": [counter_key],
}
},
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
}
def seed_slot(slot_id):
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id=slot_id,
counter_keys=[counter_key],
)
kwargs = {"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}}
async def in_flight():
return parallel_request_handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=counter_key)
)
seed_slot("slot-a")
await parallel_request_handler.async_log_failure_event(
kwargs=kwargs_with_slot("slot-a"), response_obj=None, start_time=None, end_time=None
kwargs=kwargs, response_obj=None, start_time=None, end_time=None
)
assert get_request_stash().parallel_slot is None
assert await in_flight() == 1
for slot_id in ("slot-a", "slot-unknown", "slot-a"):
seed_slot(slot_id)
await parallel_request_handler.async_log_failure_event(
kwargs=kwargs_with_slot(slot_id), response_obj=None, start_time=None, end_time=None
kwargs=kwargs, response_obj=None, start_time=None, end_time=None
)
assert await in_flight() == 1
seed_slot("slot-b")
await parallel_request_handler.async_log_failure_event(
kwargs=kwargs_with_slot("slot-b"), response_obj=None, start_time=None, end_time=None
kwargs=kwargs, response_obj=None, start_time=None, end_time=None
)
assert await in_flight() == 0
@ -803,8 +815,9 @@ async def test_rejected_request_does_not_consume_parallel_slot_v3():
data=admitted_data,
call_type="",
)
acquisition = admitted_data["metadata"][MAX_PARALLEL_SLOT_ACQUIRED_KEY]
assert isinstance(acquisition, dict)
assert "metadata" not in admitted_data
acquisition = get_request_stash().parallel_slot
assert acquisition is not None
assert isinstance(acquisition["slot_id"], str) and acquisition["slot_id"]
assert acquisition["counter_keys"] == [f"{{api_key:{_api_key}}}:max_parallel_requests"]
@ -816,10 +829,10 @@ async def test_rejected_request_does_not_consume_parallel_slot_v3():
data={"model": "gpt-3.5-turbo"},
call_type="",
)
assert get_request_stash().parallel_slot == acquisition
await handler.async_log_failure_event(
kwargs={
"metadata": {MAX_PARALLEL_SLOT_ACQUIRED_KEY: acquisition},
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
},
response_obj=None,
@ -866,8 +879,8 @@ async def test_parallel_gauge_uses_atomic_redis_script_v3():
data=data,
call_type="",
)
stashed_acquisition = data["metadata"][MAX_PARALLEL_SLOT_ACQUIRED_KEY]
assert isinstance(stashed_acquisition, dict)
stashed_acquisition = get_request_stash().parallel_slot
assert stashed_acquisition is not None
stashed_slot_id = stashed_acquisition["slot_id"]
assert isinstance(stashed_slot_id, str) and stashed_slot_id
assert stashed_acquisition["counter_keys"] == [counter_key]
@ -882,7 +895,7 @@ async def test_parallel_gauge_uses_atomic_redis_script_v3():
)
gauge_statuses = [
s
for s in data["litellm_proxy_rate_limit_response"]["statuses"]
for s in get_request_stash().rate_limit_response["statuses"]
if s["rate_limit_type"] == "max_parallel_requests"
]
assert gauge_statuses == [
@ -3102,14 +3115,12 @@ async def test_project_model_rate_limits_not_triggered_for_other_model_v3():
@pytest.mark.asyncio
async def test_pre_call_hook_does_not_leak_internal_stash_to_request_body():
"""Regression for #27001: stash keys must stay in metadata, never on
the top level of ``data`` (which gets forwarded as the provider body)."""
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_LITELLM_STASH_KEYS,
RATE_LIMIT_DESCRIPTORS_KEY,
TPM_RESERVED_TOKENS_KEY,
)
async def test_pre_call_hook_keeps_internal_stash_out_of_request_body():
"""Regression for #27001 / #35197: the limiter's per-request bookkeeping
must never touch the outgoing request body — no top-level keys and no
created or mutated ``metadata`` / ``litellm_metadata`` buckets. The
reservation must land on the ContextVar stash instead."""
import copy
_api_key = hash_token("sk-leak-regression")
user_api_key_dict = UserAPIKeyAuth(
@ -3149,6 +3160,7 @@ async def test_pre_call_hook_does_not_leak_internal_stash_to_request_body():
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 10,
}
body_before = copy.deepcopy(data)
await parallel_request_handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
@ -3157,31 +3169,27 @@ async def test_pre_call_hook_does_not_leak_internal_stash_to_request_body():
call_type="completion",
)
leaked = [k for k in _LITELLM_STASH_KEYS if k in data]
assert not leaked, f"stash keys leaked to top level: {leaked}"
assert data == body_before
metadata = data.get("metadata") or {}
assert metadata.get(TPM_RESERVED_TOKENS_KEY)
assert isinstance(metadata.get(RATE_LIMIT_DESCRIPTORS_KEY), list)
stash = get_request_stash()
assert stash is not None
assert stash.reserved_tokens > 0
assert stash.reserved_model == "gpt-4o-mini"
assert stash.reserved_scopes == frozenset({("api_key", _api_key)})
@pytest.mark.asyncio
@pytest.mark.parametrize("caller_metadata", [None, {"user_tag": "abc"}])
async def test_pre_call_hook_does_not_touch_provider_metadata_on_litellm_metadata_routes(
caller_metadata,
):
"""Regression for #35197: routes that own ``litellm_metadata`` (Responses,
/v1/messages, batches, files) send ``metadata`` to the provider, so the
limiter must never create it or write stash keys into it."""
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_LITELLM_STASH_KEYS,
RATE_LIMIT_DESCRIPTORS_KEY,
RATE_LIMIT_RESPONSE_KEY,
TPM_RESERVED_TOKENS_KEY,
)
@pytest.mark.parametrize("caller_metadata", [None, {"user_tag": "campaign-42"}])
async def test_responses_route_body_untouched_by_pre_call_hook(caller_metadata):
"""Regression for #35197: on routes where ``metadata`` is a provider
request parameter (Responses API), the pre-call hook must forward the
body byte-identical — creating or adding to ``metadata`` /
``litellm_metadata`` produced upstream HTTP 400s."""
import copy
_api_key = hash_token("sk-responses-regression")
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token("sk-responses-metadata"),
api_key=_api_key,
tpm_limit=1000,
rpm_limit=5,
)
@ -3190,35 +3198,13 @@ async def test_pre_call_hook_does_not_touch_provider_metadata_on_litellm_metadat
internal_usage_cache=InternalUsageCache(local_cache),
)
async def mock_should_rate_limit(descriptors, **kwargs):
return {
"overall_code": "OK",
"statuses": [
{
"code": "OK",
"current_limit": 5,
"limit_remaining": 4,
"descriptor_key": d["key"],
"descriptor_value": d["value"],
"rate_limit_type": "requests",
}
for d in descriptors
],
}
async def mock_reserve_tpm_tokens(descriptors, estimated_tokens, **kwargs):
return {"overall_code": "OK", "statuses": []}
handler.should_rate_limit = mock_should_rate_limit
handler.reserve_tpm_tokens = mock_reserve_tpm_tokens
data: Dict[str, Any] = {
"model": "responses-model",
"model": "gpt-4o-mini",
"input": "hello",
"litellm_metadata": {},
}
if caller_metadata is not None:
data["metadata"] = dict(caller_metadata)
body_before = copy.deepcopy(data)
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
@ -3227,37 +3213,87 @@ async def test_pre_call_hook_does_not_touch_provider_metadata_on_litellm_metadat
call_type="aresponses",
)
assert data == body_before
if caller_metadata is None:
assert "metadata" not in data, f"limiter created provider metadata: {data.get('metadata')!r}"
assert "metadata" not in data
else:
assert data["metadata"] == caller_metadata
assert "litellm_metadata" not in data
litellm_metadata = data["litellm_metadata"]
assert litellm_metadata.get(TPM_RESERVED_TOKENS_KEY)
assert isinstance(litellm_metadata.get(RATE_LIMIT_DESCRIPTORS_KEY), list)
assert litellm_metadata.get(RATE_LIMIT_RESPONSE_KEY)
leaked = [k for k in _LITELLM_STASH_KEYS if k in data]
assert not leaked, f"stash keys leaked to top level: {leaked}"
for key in _LITELLM_STASH_KEYS:
assert handler._lookup_stashed_value(
kwargs={"litellm_params": {"litellm_metadata": litellm_metadata}},
standard_logging_metadata=None,
key=key,
) == litellm_metadata.get(key)
stash = get_request_stash()
assert stash is not None
assert stash.reserved_tokens > 0
assert stash.rate_limit_response is not None
@pytest.mark.asyncio
async def test_pre_call_hook_rejects_caller_supplied_stash_values():
"""Caller cannot pre-populate stash keys in body metadata to drive a
later TPM refund against an arbitrary scope."""
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_LITELLM_STASH_KEYS,
RATE_LIMIT_DESCRIPTORS_KEY,
TPM_RESERVED_TOKENS_KEY,
async def test_chat_tpm_refund_and_slot_release_via_context_stash(monkeypatch):
"""
Full chat lifecycle with no body stashing: pre-call reserves TPM tokens
and acquires a parallel slot on the ContextVar stash; the failure
callback refunds the reservation and frees the slot exactly once — a
second failure callback for the same request must not double-refund the
:tokens counter or double-release the gauge.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
_api_key = hash_token("sk-refund-lifecycle")
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
tpm_limit=10_000,
max_parallel_requests=2,
)
tokens_key = handler.create_rate_limit_keys(
key="api_key", value=_api_key, rate_limit_type="tokens"
)
parallel_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 50,
},
call_type="completion",
)
reserved = get_request_stash().reserved_tokens
assert reserved > 0
assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == reserved
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=parallel_key)
) == 1
kwargs = {"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}}
await handler.async_log_failure_event(
kwargs=kwargs, response_obj=None, start_time=None, end_time=None
)
assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=parallel_key)
) == 0
assert get_request_stash().reservation_released is True
await handler.async_log_failure_event(
kwargs=kwargs, response_obj=None, start_time=None, end_time=None
)
assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=parallel_key)
) == 0
@pytest.mark.asyncio
async def test_pre_call_hook_ignores_caller_supplied_stash_values():
"""Caller-supplied bookkeeping lookalikes in the body must not drive a
TPM refund against an arbitrary scope: the ContextVar stash is the only
source the refund path reads."""
user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-no-limits"))
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
@ -3271,19 +3307,15 @@ async def test_pre_call_hook_rejects_caller_supplied_stash_values():
"rate_limit": {"tokens_per_unit": 10000, "window_size": 60},
}
]
injected = {
"_litellm_tpm_reserved_tokens": 9999,
"_litellm_rate_limit_descriptors": victim_descriptors,
}
data: Dict[str, Any] = {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hi"}],
TPM_RESERVED_TOKENS_KEY: 9999,
RATE_LIMIT_DESCRIPTORS_KEY: victim_descriptors,
"metadata": {
TPM_RESERVED_TOKENS_KEY: 9999,
RATE_LIMIT_DESCRIPTORS_KEY: victim_descriptors,
},
"litellm_metadata": {
TPM_RESERVED_TOKENS_KEY: 9999,
RATE_LIMIT_DESCRIPTORS_KEY: victim_descriptors,
},
"metadata": dict(injected),
"litellm_metadata": dict(injected),
}
await handler.async_pre_call_hook(
@ -3293,13 +3325,139 @@ async def test_pre_call_hook_rejects_caller_supplied_stash_values():
call_type="completion",
)
for channel in (
data,
data.get("metadata") or {},
data.get("litellm_metadata") or {},
):
leaked = [k for k in _LITELLM_STASH_KEYS if k in channel]
assert not leaked, f"caller-supplied stash survived in {channel!r}: {leaked}"
refund_calls = []
async def spy_increment_pipeline(increment_list, **kwargs):
refund_calls.append(increment_list)
handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = (
spy_increment_pipeline
)
await handler.async_post_call_failure_hook(
request_data=data,
original_exception=Exception("boom"),
user_api_key_dict=user_api_key_dict,
)
assert refund_calls == []
stash = get_request_stash()
assert stash is not None
assert stash.reserved_tokens == 0
@pytest.mark.asyncio
async def test_log_events_from_nested_calls_leave_owner_stash_alone(monkeypatch):
"""
A nested LiteLLM call made inside the request (LLM-judge guardrail,
silent experiment) inherits the request context and fires the same global
logging callbacks with a fresh ``litellm_call_id``. Those callbacks must
not release the owning request's parallel slot or refund its TPM
reservation; only events carrying the owner's call id may.
"""
monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False)
_api_key = hash_token("sk-nested-guard")
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
tpm_limit=10_000,
max_parallel_requests=2,
)
tokens_key = handler.create_rate_limit_keys(
key="api_key", value=_api_key, rate_limit_type="tokens"
)
parallel_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 50,
"litellm_call_id": "owner-call-id",
},
call_type="completion",
)
stash = get_request_stash()
assert stash is not None
assert stash.owner_litellm_call_id == "owner-call-id"
reserved = stash.reserved_tokens
assert reserved > 0
nested_kwargs = {
"litellm_call_id": "nested-guardrail-call",
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
}
await handler.async_log_success_event(
kwargs=nested_kwargs, response_obj=None, start_time=None, end_time=None
)
await handler.async_log_failure_event(
kwargs=nested_kwargs, response_obj=None, start_time=None, end_time=None
)
assert stash.parallel_slot is not None
assert stash.reservation_released is False
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=parallel_key)
) == 1
assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == reserved
owner_kwargs = {
"litellm_call_id": "owner-call-id",
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
}
await handler.async_log_failure_event(
kwargs=owner_kwargs, response_obj=None, start_time=None, end_time=None
)
assert stash.parallel_slot is None
assert stash.reservation_released is True
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=parallel_key)
) == 0
assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0
@pytest.mark.asyncio
async def test_stash_applies_when_owner_or_callback_call_id_missing():
"""
The owner guard only rejects a positive mismatch. A stash never claimed
by a pre-call hook (no owner id) must stay visible to any callback, and a
claimed stash must stay visible to callbacks whose kwargs carry no call
id — otherwise reservations and slots would strand on request paths that
do not thread ``litellm_call_id`` into their logging kwargs.
"""
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
unclaimed = get_or_create_request_stash()
unclaimed.reserved_tokens = 42
await handler.async_log_failure_event(
kwargs={"litellm_call_id": "any-id", "standard_logging_object": {}},
response_obj=None,
start_time=None,
end_time=None,
)
assert unclaimed.reservation_released is True
claimed = RequestRateLimiterStash(
owner_litellm_call_id="owner-1", reserved_tokens=42
)
_request_stash.set(claimed)
await handler.async_log_failure_event(
kwargs={"standard_logging_object": {}},
response_obj=None,
start_time=None,
end_time=None,
)
assert claimed.reservation_released is True
# ----------------------- Per-MCP-server rate limiting (v3) -----------------------
@ -3594,18 +3752,13 @@ async def test_release_max_parallel_requests_on_disconnect_v3():
await local_cache.async_get_cache(key=counter_key)
) == 1
await handler.async_release_max_parallel_requests_on_disconnect(
user_api_key_dict,
request_data={
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": _TEST_SLOT_ID,
"counter_keys": [counter_key],
}
}
},
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id=_TEST_SLOT_ID,
counter_keys=[counter_key],
)
await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict)
assert get_request_stash().parallel_slot is None
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=counter_key)
) == 0
@ -3627,16 +3780,12 @@ async def test_release_on_disconnect_works_when_key_config_changed_v3():
counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests"
await _seed_max_parallel_requests_slots(local_cache, counter_key, [_TEST_SLOT_ID])
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id=_TEST_SLOT_ID,
counter_keys=[counter_key],
)
await handler.async_release_max_parallel_requests_on_disconnect(
UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None),
request_data={
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": _TEST_SLOT_ID,
"counter_keys": [counter_key],
}
}
},
UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None)
)
assert handler._gauge_in_flight_from_cache_value(
await local_cache.async_get_cache(key=counter_key)
@ -3684,7 +3833,6 @@ async def test_post_call_failure_hook_releases_parallel_slot_v3():
await handler.async_log_failure_event(
kwargs={
"metadata": admitted_data["metadata"],
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
},
response_obj=None,
@ -3732,7 +3880,6 @@ async def test_success_event_releases_parallel_slot_v3(monkeypatch):
await handler.async_log_success_event(
kwargs={
"metadata": admitted_data["metadata"],
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
},
response_obj=ModelResponse(
@ -3833,14 +3980,12 @@ async def test_redis_release_script_updates_local_mirror_v3():
handler.parallel_release_script = fake_release
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id="slot-redis-test",
counter_keys=[counter_key],
)
await handler.async_log_failure_event(
kwargs={
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": "slot-redis-test",
"counter_keys": [counter_key],
}
},
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
},
response_obj=None,
@ -3945,7 +4090,6 @@ async def test_in_memory_fallback_respects_mirrored_redis_count_v3():
await handler.async_log_failure_event(
kwargs={
"metadata": admitted_data["metadata"],
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
},
response_obj=None,
@ -4011,19 +4155,15 @@ async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3(
while True:
yield ModelResponse()
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id=_TEST_SLOT_ID,
counter_keys=[counter_key],
)
with _override_litellm_callbacks([]):
gen = ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={
"model": "claude-test",
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": _TEST_SLOT_ID,
"counter_keys": [counter_key],
}
},
},
request_data={"model": "claude-test"},
proxy_logging_obj=proxy_logging_obj,
)
await gen.__anext__()
@ -4064,21 +4204,17 @@ async def test_async_data_generator_releases_counter_on_disconnect_v3(disconnect
while True:
yield ModelResponse()
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id=_TEST_SLOT_ID,
counter_keys=[counter_key],
)
try:
with _override_litellm_callbacks([]):
assert proxy_logging_obj.needs_iterator_wrap() is False
gen = proxy_server.async_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={
"model": "gpt-test",
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": _TEST_SLOT_ID,
"counter_keys": [counter_key],
}
},
},
request_data={"model": "gpt-test"},
)
await gen.__anext__()
if disconnect == "cancel":
@ -4127,21 +4263,17 @@ async def test_async_data_generator_releases_counter_when_wrapped_v3():
while True:
yield ModelResponse()
get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition(
slot_id=_TEST_SLOT_ID,
counter_keys=[counter_key],
)
try:
with _override_litellm_callbacks([_PassthroughIteratorOverride()]):
assert proxy_logging_obj.needs_iterator_wrap() is True
gen = proxy_server.async_data_generator(
response=upstream(),
user_api_key_dict=user_api_key_dict,
request_data={
"model": "gpt-test",
"metadata": {
MAX_PARALLEL_SLOT_ACQUIRED_KEY: {
"slot_id": _TEST_SLOT_ID,
"counter_keys": [counter_key],
}
},
},
request_data={"model": "gpt-test"},
)
await gen.__anext__()
await gen.aclose()
@ -4258,12 +4390,7 @@ async def test_pre_call_hook_skips_reservation_when_disabled(monkeypatch):
assert reserve_calls == [], "reservation must be skipped when disabled"
assert should_rate_limit_calls[0]["skip_tpm_check"] is False
# No reservation stash leaks into the request metadata.
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
TPM_RESERVED_TOKENS_KEY,
)
assert TPM_RESERVED_TOKENS_KEY not in (data.get("metadata") or {})
assert get_request_stash().reserved_tokens == 0
@pytest.mark.asyncio

View file

@ -691,7 +691,6 @@ async def test_dynamic_rate_limiter_v3_model_capacity_path_populates_provider():
user_api_key_dict=user_api_key_dict,
priority="default",
saturation=1.0,
data={"model": "gpt-4o-mini"},
)
exc = exc_info.value
@ -741,7 +740,6 @@ async def test_dynamic_rate_limiter_v3_unknown_descriptor_path_populates_provide
user_api_key_dict=user_api_key_dict,
priority="default",
saturation=1.0,
data={"model": "gpt-4o-mini"},
)
assert exc_info.value.llm_provider == "openai"

View file

@ -253,7 +253,6 @@ async def test_dynamic_rate_limiter_v3_concurrent_bypasses_model_capacity():
user_api_key_dict=user,
priority="high",
saturation=0.0,
data={},
)
return "OK"
except Exception as e:
@ -332,7 +331,6 @@ async def test_dynamic_rate_limiter_v3_uses_atomic_check_and_increment():
user_api_key_dict=user,
priority="high",
saturation=0.0,
data={},
)
assert atomic_descriptors_observed, (
@ -482,7 +480,6 @@ async def test_dynamic_rate_limiter_v3_fails_closed_on_unknown_descriptor():
user_api_key_dict=user,
priority="high",
saturation=0.0,
data={},
)
assert (
exc.value.status_code == 429

View file

@ -23,13 +23,13 @@ import pytest
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
RATE_LIMIT_DESCRIPTORS_KEY,
TPM_RESERVATION_RELEASED_KEY,
TPM_RESERVED_MODEL_KEY,
TPM_RESERVED_SCOPES_KEY,
TPM_RESERVED_TOKENS_KEY,
_PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_request_stash,
get_or_create_request_stash,
get_request_stash,
)
from litellm.proxy.utils import InternalUsageCache, hash_token
from litellm.types.utils import ModelResponse, Usage
@ -41,6 +41,13 @@ def rate_limiter():
return handler, cache
@pytest.fixture(autouse=True)
def _isolated_request_stash():
token = _request_stash.set(None)
yield
_request_stash.reset(token)
@pytest.mark.asyncio
async def test_token_reservation_prevents_concurrent_bypass(rate_limiter):
"""
@ -79,7 +86,7 @@ async def test_token_reservation_prevents_concurrent_bypass(rate_limiter):
return {
"request_id": request_id,
"success": True,
"reserved_tokens": data.get(TPM_RESERVED_TOKENS_KEY, 0),
"reserved_tokens": get_request_stash().reserved_tokens,
}
except Exception as e:
return {
@ -167,12 +174,14 @@ async def test_token_adjustment_on_success(rate_limiter):
api_key = hash_token("sk-test-adjust")
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("api_key", api_key)})
mock_kwargs = {
"standard_logging_object": {
"metadata": {
"user_api_key_hash": api_key,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_SCOPES_KEY: [["api_key", api_key]],
}
},
"model": "gpt-3.5-turbo",
@ -227,12 +236,14 @@ async def test_token_release_on_failure(rate_limiter):
api_key = hash_token("sk-test-fail")
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("api_key", api_key)})
mock_kwargs = {
"standard_logging_object": {
"metadata": {
"user_api_key_hash": api_key,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_SCOPES_KEY: [["api_key", api_key]],
}
},
}
@ -285,6 +296,11 @@ async def test_model_scope_refund_targets_reserved_model(rate_limiter):
team_id = "team-abc"
reserved_model = "gpt-4o-mini"
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_model = reserved_model
stash.reserved_scopes = frozenset({("model_per_team", f"{team_id}:{reserved_model}")})
mock_kwargs = {
# NOTE: no litellm_params.metadata.model_group — get_model_group_from_litellm_kwargs
# returns None on this kwargs dict.
@ -292,11 +308,6 @@ async def test_model_scope_refund_targets_reserved_model(rate_limiter):
"metadata": {
"user_api_key_hash": api_key,
"user_api_key_team_id": team_id,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_MODEL_KEY: reserved_model,
TPM_RESERVED_SCOPES_KEY: [
["model_per_team", f"{team_id}:{reserved_model}"]
],
}
},
}
@ -446,13 +457,15 @@ async def test_org_scope_refund_on_failure(rate_limiter):
api_key = hash_token("sk-org-refund")
org_id = "org-acme"
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("organization", org_id)})
mock_kwargs = {
"standard_logging_object": {
"metadata": {
"user_api_key_hash": api_key,
"user_api_key_org_id": org_id,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_SCOPES_KEY: [["organization", org_id]],
}
},
}
@ -498,13 +511,15 @@ async def test_org_scope_reconciled_on_success(rate_limiter):
api_key = hash_token("sk-org-success")
org_id = "org-acme"
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("organization", org_id)})
mock_kwargs = {
"standard_logging_object": {
"metadata": {
"user_api_key_hash": api_key,
"user_api_key_org_id": org_id,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_SCOPES_KEY: [["organization", org_id]],
}
},
"model": "gpt-3.5-turbo",
@ -607,9 +622,9 @@ async def test_contentless_request_reserves_minimum(rate_limiter):
data=data,
call_type="",
)
assert (data.get("metadata") or {}).get(
TPM_RESERVED_TOKENS_KEY
) == 1, "Contentless request should reserve the floor of 1 token"
assert (
get_request_stash().reserved_tokens == 1
), "Contentless request should reserve the floor of 1 token"
counter_after_two = int(
await cache.async_get_cache(key=counter_key, local_only=True) or 0
@ -702,7 +717,7 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter):
data=data,
call_type="",
)
reserved = (data.get("metadata") or {})[TPM_RESERVED_TOKENS_KEY]
reserved = get_request_stash().reserved_tokens
assert reserved > 0
counter_key = handler.create_rate_limit_keys(
@ -727,8 +742,8 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter):
f"Reservation leaked: counter={counter_after_release} after "
f"proxy-level rejection refund (expected 0)."
)
assert (data.get("metadata") or {}).get(TPM_RESERVATION_RELEASED_KEY) is True, (
"Released marker must be stamped to prevent "
assert get_request_stash().reservation_released is True, (
"Released flag must be set to prevent "
"async_log_failure_event from double-refunding."
)
@ -754,28 +769,15 @@ async def test_reservation_release_idempotent(rate_limiter):
mock_increment
)
# Shared metadata dict simulates the propagation between
# request_data["metadata"] and kwargs["litellm_params"]["metadata"] —
# the post-call-failure-hook stamps the released marker there, and the
# log-failure-event reads it.
shared_metadata = {
"user_api_key_hash": api_key,
TPM_RESERVED_TOKENS_KEY: 100,
RATE_LIMIT_DESCRIPTORS_KEY: [
{
"key": "api_key",
"value": api_key,
"rate_limit": {"tokens_per_unit": 10000, "window_size": 60},
}
],
}
request_data = {
"metadata": shared_metadata,
}
# Both hooks read the same per-request ContextVar stash: the
# post-call-failure-hook flips reservation_released on it, and the
# log-failure-event observes the flip.
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("api_key", api_key)})
await handler.async_post_call_failure_hook(
request_data=request_data,
request_data={},
original_exception=Exception("rejected"),
user_api_key_dict=UserAPIKeyAuth(api_key=api_key),
)
@ -784,11 +786,10 @@ async def test_reservation_release_idempotent(rate_limiter):
assert first_refund_count > 0, "First refund should have applied"
# Now simulate async_log_failure_event firing afterwards. It must see
# the released marker (via shared metadata) and not double-refund.
# the released flag on the stash and not double-refund.
await handler.async_log_failure_event(
kwargs={
"litellm_params": {"metadata": shared_metadata},
"standard_logging_object": {"metadata": shared_metadata},
"standard_logging_object": {"metadata": {"user_api_key_hash": api_key}},
},
response_obj=None,
start_time=datetime.now(),
@ -818,13 +819,15 @@ async def test_unreserved_scopes_charged_actual_not_delta_on_success(rate_limite
team_id = "team-no-tpm-limit"
# Reservation ONLY hit api_key — team had no TPM limit configured.
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("api_key", api_key)})
mock_kwargs = {
"standard_logging_object": {
"metadata": {
"user_api_key_hash": api_key,
"user_api_key_team_id": team_id,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_SCOPES_KEY: [["api_key", api_key]],
}
},
"model": "gpt-3.5-turbo",
@ -888,13 +891,15 @@ async def test_unreserved_scopes_not_refunded_on_failure(rate_limiter):
api_key = hash_token("sk-mixed-fail")
team_id = "team-no-tpm"
stash = get_or_create_request_stash()
stash.reserved_tokens = 100
stash.reserved_scopes = frozenset({("api_key", api_key)})
mock_kwargs = {
"standard_logging_object": {
"metadata": {
"user_api_key_hash": api_key,
"user_api_key_team_id": team_id,
TPM_RESERVED_TOKENS_KEY: 100,
TPM_RESERVED_SCOPES_KEY: [["api_key", api_key]],
}
},
}
@ -939,10 +944,10 @@ async def test_unreserved_scopes_not_refunded_on_failure(rate_limiter):
async def test_token_rate_limit_headers_present_in_stored_response(rate_limiter):
"""
With `skip_tpm_check=True` on the RPM sliding-window pass, token statuses
only come from `reserve_tpm_tokens`. They must be merged into
`data["litellm_proxy_rate_limit_response"]` so the post-call hook can
emit `x-ratelimit-{key}-remaining-tokens` / `-limit-tokens` headers to
the client.
only come from `reserve_tpm_tokens`. They must be merged into the stashed
rate-limit response so the post-call hook can emit
`x-ratelimit-{key}-remaining-tokens` / `-limit-tokens` headers to the
client.
"""
handler, cache = rate_limiter
@ -966,10 +971,10 @@ async def test_token_rate_limit_headers_present_in_stored_response(rate_limiter)
call_type="",
)
response = data.get("litellm_proxy_rate_limit_response")
response = get_request_stash().rate_limit_response
assert isinstance(
response, dict
), "Expected litellm_proxy_rate_limit_response to be set after pre-call"
), "Expected the stashed rate-limit response to be set after pre-call"
statuses = response.get("statuses") or []
token_statuses = [s for s in statuses if s.get("rate_limit_type") == "tokens"]
@ -1080,8 +1085,8 @@ async def test_small_tpm_cap_admits_no_max_tokens_request(rate_limiter):
call_type="",
)
reserved = (data.get("metadata") or {}).get(TPM_RESERVED_TOKENS_KEY)
assert reserved is not None, "Reservation should have been stashed"
reserved = get_request_stash().reserved_tokens
assert reserved > 0, "Reservation should have been stashed"
assert reserved <= 1000 // 2, (
f"Capped floor must keep the reservation well under the 1000 TPM "
f"cap; got {reserved}"

View file

@ -881,7 +881,6 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError:
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test-v3"),
priority="default",
saturation=0.99,
data={},
)
e = exc_info.value
assert e.status_code == 429

View file

@ -321,29 +321,6 @@ class TestNativeFinishReason:
assert choice.provider_specific_fields["native_finish_reason"] == "MAX_TOKENS"
def test_parallel_request_limiter_internal_fields_in_all_litellm_params():
"""
Regression test: internal fields written by parallel_request_limiter_v3 must
be in all_litellm_params so they are stripped before forwarding to upstream
providers. If missing, they are sent as extra body parameters and providers
like OpenAI reject the request with a 400 invalid_request_error.
"""
from litellm.types.utils import all_litellm_params
internal_fields = [
"_litellm_rate_limit_descriptors",
"_litellm_tpm_reserved_tokens",
"_litellm_tpm_reserved_model",
"_litellm_tpm_reserved_scopes",
"_litellm_tpm_reservation_released",
]
for field in internal_fields:
assert field in all_litellm_params, (
f"{field!r} is not in all_litellm_params. "
"It will be forwarded to upstream providers and cause 400 errors."
)
def test_delta_maps_reasoning_to_reasoning_content():
"""
Test that Delta maps 'reasoning' field to 'reasoning_content'.