From 2c1d62ce2b586e047307c83d3ed79c64857287e8 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Sat, 11 Jul 2026 12:28:38 -0700 Subject: [PATCH] fix(rate-limit-v3): populate x-ratelimit-* remaining/limit values in standard_logging_object for streaming (LIT-4333) (#32711) Streaming requests return from common_request_processing before async_post_call_success_hook runs, so response._hidden_params.additional_headers never gets the v3 x-ratelimit-{descriptor_key}-{remaining|limit}-{rate_limit_type} entries. Prometheus / logging callbacks that read those values from standard_logging_object.hidden_params.additional_headers then see nothing; combined with the pre-existing gap that Prometheus reads from that same slot (LIT-2577 / PR #28816), per-key remaining RPM/TPM cannot be monitored for streaming traffic at all. Fix in three parts: - Stash the pre-call RateLimitResponse in the metadata channels the async success-logging callback inherits, alongside the existing top-level entry the non-streaming path reads. - Add async_logging_hook to the v3 handler. It fires in a distinct earlier loop inside async_success_handler (all callbacks' async_logging_hook complete before any async_log_success_event starts), so mirroring the pre-call snapshot into standard_logging_object.hidden_params.additional_headers and response._hidden_params.additional_headers here guarantees every downstream success callback sees the values regardless of registration order. Non-streaming keeps the existing async_post_call_success_hook write and this hook re-populates the same values idempotently. - Extract the shared `_merge_ratelimit_statuses_into_additional_headers` helper the non-streaming path already had inlined so both callsites emit the identical key shape. --- .../hooks/parallel_request_limiter_v3.py | 155 +++++++- .../hooks/test_parallel_request_limiter_v3.py | 338 ++++++++++++++++++ 2 files changed, 484 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 7aedb74f2ea..d60c17c744f 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -244,6 +244,10 @@ TPM_RESERVED_SCOPES_KEY = "_litellm_tpm_reserved_scopes" # 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" # 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. @@ -253,6 +257,7 @@ _LITELLM_STASH_KEYS: Tuple[str, ...] = ( TPM_RESERVED_SCOPES_KEY, TPM_RESERVATION_RELEASED_KEY, RATE_LIMIT_DESCRIPTORS_KEY, + RATE_LIMIT_RESPONSE_KEY, ) @@ -2037,6 +2042,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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_metadata_channels( + data=data, + key=RATE_LIMIT_RESPONSE_KEY, + value=response, + ) # ---------------------------------------------------------------- # TPM token reservation @@ -2133,6 +2145,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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_metadata_channels( + data=data, + key=RATE_LIMIT_RESPONSE_KEY, + value=tpm_response, + ) verbose_proxy_logger.debug(f"TPM tokens reserved: {estimated_tokens} for model {requested_model}") @@ -2318,6 +2337,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return "total" # default to total return specified_rate_limit_type + @staticmethod + def _merge_ratelimit_statuses_into_additional_headers( + additional_headers: Dict[str, Any], + statuses: List[RateLimitStatus], + ) -> Dict[str, Any]: + """ + Return ``additional_headers`` extended with + ``x-ratelimit-{descriptor_key}-{remaining|limit}-{rate_limit_type}`` + entries. Non-mutating so callers pick their own target dict. + """ + merged: Dict[str, Any] = dict(additional_headers) + for status in statuses: + prefix = f"x-ratelimit-{status['descriptor_key']}" + merged[f"{prefix}-remaining-{status['rate_limit_type']}"] = status["limit_remaining"] + merged[f"{prefix}-limit-{status['rate_limit_type']}"] = status["current_limit"] + return merged + @staticmethod def _stash_value_in_metadata_channels( data: Dict[str, Any], @@ -2698,6 +2734,112 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): except Exception as e: verbose_proxy_logger.exception(f"Error in rate limit success event: {str(e)}") + async def async_logging_hook( + self, + kwargs: dict, + result: Any, + call_type: str, + ) -> Tuple[dict, Any]: + """ + Mirror the pre-call rate-limit snapshot into the SLP so streaming + success callbacks see the same ``x-ratelimit-*`` headers the + non-streaming path writes via ``async_post_call_success_hook``. + Runs in the earlier of the two callback loops inside + ``async_success_handler`` so downstream callbacks see the values + regardless of registration order. Idempotent for non-streaming. + """ + self._mirror_ratelimit_response_into_logging_payload( + kwargs=kwargs, + response_obj=result, + ) + return kwargs, result + + def _mirror_ratelimit_response_into_logging_payload( + self, + kwargs: Any, + response_obj: Any, + ) -> None: + """ + Copy the stashed ``RateLimitResponse`` into the SLP's + ``hidden_params.additional_headers`` and the response object's + ``_hidden_params.additional_headers`` (when the latter is a dict). + """ + 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, + ) + ) + if not statuses: + return + + if isinstance(standard_logging_object, dict): + hidden_params = standard_logging_object.get("hidden_params") + if not isinstance(hidden_params, dict): + hidden_params = {} + existing = hidden_params.get("additional_headers") + hidden_params["additional_headers"] = self._merge_ratelimit_statuses_into_additional_headers( + additional_headers=existing if isinstance(existing, dict) else {}, + statuses=statuses, + ) + standard_logging_object["hidden_params"] = hidden_params + + response_hidden = getattr(response_obj, "_hidden_params", None) + if isinstance(response_hidden, dict): + existing = response_hidden.get("additional_headers") + response_hidden["additional_headers"] = self._merge_ratelimit_statuses_into_additional_headers( + additional_headers=existing if isinstance(existing, dict) else {}, + 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 @@ -2838,15 +2980,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if isinstance(_hidden_params, BaseModel): _hidden_params = _hidden_params.model_dump() - _additional_headers = _hidden_params.get("additional_headers", {}) or {} - - # Add rate limit headers - for status in litellm_proxy_rate_limit_response["statuses"]: - prefix = f"x-ratelimit-{status['descriptor_key']}" - _additional_headers[f"{prefix}-remaining-{status['rate_limit_type']}"] = status[ - "limit_remaining" - ] - _additional_headers[f"{prefix}-limit-{status['rate_limit_type']}"] = status["current_limit"] + _additional_headers = self._merge_ratelimit_statuses_into_additional_headers( + additional_headers=_hidden_params.get("additional_headers", {}) or {}, + statuses=litellm_proxy_rate_limit_response["statuses"], + ) setattr( response, diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index d150591c8de..e7d2909263a 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3706,3 +3706,341 @@ async def test_per_tag_untagged_request_governed_by_key_limit_v3(monkeypatch): await call({"tags": ["cell-99"]}) assert exc_info.value.status_code == 429 assert "tag_per_key" not in str(exc_info.value.detail) + + +# -------------------------------------------------------------------------- +# Streaming success logging mirrors x-ratelimit-* remaining values into +# standard_logging_object.hidden_params.additional_headers so Prometheus / +# logging callbacks see them for streams too (non-streaming already gets +# them via async_post_call_success_hook, which the streaming path skips). +# -------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_streaming_end_to_end_populates_slp_ratelimit_headers(monkeypatch): + """ + End-to-end regression: on a streaming request, the same pre-call + + success-callback pair the proxy uses must land ``x-ratelimit-*`` + remaining/limit values in + ``kwargs["standard_logging_object"]["hidden_params"]["additional_headers"]``. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _api_key = hash_token("sk-stream-e2e") + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + rpm_limit=100, + tpm_limit=10000, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + # Real pre-call: populates data and stashes the response into metadata + # so the success callback can find it via litellm_params.metadata. + data: Dict[str, Any] = { + "model": "gpt-4o-mini", + "metadata": {}, + "stream": True, + } + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=data, + call_type="", + ) + + # Simulate the wrapper handing the pre-call metadata dict to the + # completion() call: it becomes kwargs["litellm_params"]["metadata"] by + # the time the success callback fires. + mock_response = ModelResponse( + id="mock-stream-e2e", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-4o-mini", + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + choices=[], + ) + mock_kwargs: Dict[str, Any] = { + "standard_logging_object": { + "metadata": { + "user_api_key_hash": _api_key, + "user_api_key_user_id": None, + "user_api_key_team_id": None, + "user_api_key_end_user_id": None, + } + }, + "litellm_params": {"metadata": data["metadata"]}, + "model": "gpt-4o-mini", + } + + async def _noop_increment(increment_list, **_): + return True + + monkeypatch.setattr( + handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + _noop_increment, + ) + + # async_logging_hook runs before async_log_success_event, so any + # downstream callback that reads the SLP sees the mirrored values. + await handler.async_logging_hook( + kwargs=mock_kwargs, + result=mock_response, + call_type="acompletion", + ) + + additional_headers = ( + mock_kwargs["standard_logging_object"] + .get("hidden_params", {}) + .get("additional_headers", {}) + ) + + # api_key-scoped remaining/limit values are the baseline every request + # emits and must always reach the SLP. + remaining_keys = [ + k for k in additional_headers if "-remaining-" in k + ] + assert ( + remaining_keys + ), f"streaming success must populate remaining values, got {additional_headers!r}" + limit_keys = [k for k in additional_headers if "-limit-" in k] + assert limit_keys, "streaming success must also populate limit values" + assert ( + additional_headers.get("x-ratelimit-api_key-remaining-requests") == 99 + ), ( + "api_key remaining requests should reflect the just-consumed slot;" + f" got {additional_headers!r}" + ) + + +@pytest.mark.asyncio +async def test_streaming_populates_model_per_key_ratelimit_headers(monkeypatch): + """ + Streaming must land the per-(key, model) remaining/limit values in the + SLP under ``x-ratelimit-model_per_key-{remaining|limit}-{requests,tokens}``. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _api_key = hash_token("sk-stream-mirror") + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + metadata={ + "model_rpm_limit": {"gpt-4o-mini": 100}, + "model_tpm_limit": {"gpt-4o-mini": 10000}, + }, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + async def _noop_increment(increment_list, **_): + return True + + monkeypatch.setattr( + handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + _noop_increment, + ) + + data: Dict[str, Any] = {"model": "gpt-4o-mini", "metadata": {}, "stream": True} + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=data, + call_type="", + ) + + mock_response = ModelResponse( + id="mock-stream", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-4o-mini", + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + choices=[], + ) + + mock_kwargs: Dict[str, Any] = { + "standard_logging_object": { + "metadata": { + "user_api_key_hash": _api_key, + "user_api_key_user_id": None, + "user_api_key_team_id": None, + "user_api_key_end_user_id": None, + } + }, + "litellm_params": {"metadata": data["metadata"]}, + "model": "gpt-4o-mini", + } + + await handler.async_logging_hook( + kwargs=mock_kwargs, + result=mock_response, + call_type="acompletion", + ) + + hidden_params = mock_kwargs["standard_logging_object"].get("hidden_params") or {} + additional_headers = hidden_params.get("additional_headers") or {} + + assert ( + additional_headers.get("x-ratelimit-model_per_key-remaining-requests") == 99 + ), f"got {additional_headers!r}" + assert additional_headers.get("x-ratelimit-model_per_key-limit-requests") == 100 + + # response._hidden_params is also updated for late readers. + response_hidden = getattr(mock_response, "_hidden_params", None) or {} + response_headers = response_hidden.get("additional_headers") or {} + assert response_headers.get("x-ratelimit-model_per_key-remaining-requests") == 99 + + +@pytest.mark.asyncio +async def test_async_log_success_event_no_mirror_when_no_snapshot(monkeypatch): + """ + No pre-call snapshot (no descriptors matched) -> no fabricated + ``x-ratelimit-*`` headers. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _api_key = hash_token("sk-stream-no-mirror") + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + + async def _noop_increment(increment_list, **_): + return True + + monkeypatch.setattr( + handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + _noop_increment, + ) + + mock_response = ModelResponse( + id="mock-stream-none", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-4o-mini", + usage=Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2), + choices=[], + ) + + mock_kwargs: Dict[str, Any] = { + "standard_logging_object": { + "metadata": { + "user_api_key_hash": _api_key, + "user_api_key_user_id": None, + "user_api_key_team_id": None, + "user_api_key_end_user_id": None, + } + }, + "litellm_params": {"metadata": {}}, + "model": "gpt-4o-mini", + } + + await handler.async_logging_hook( + kwargs=mock_kwargs, + result=mock_response, + call_type="acompletion", + ) + + hidden_params = mock_kwargs["standard_logging_object"].get("hidden_params") or {} + additional_headers = hidden_params.get("additional_headers") or {} + ratelimit_keys = [k for k in additional_headers if k.startswith("x-ratelimit-")] + assert ( + not ratelimit_keys + ), f"no snapshot must produce no rate-limit headers, got {ratelimit_keys}" + + +@pytest.mark.asyncio +async def test_streaming_mirror_matches_non_streaming_header_shape(monkeypatch): + """ + Given the same pre-call state, streaming and non-streaming must write + the identical ``x-ratelimit-*`` key/value shape to their respective + ``additional_headers`` slots. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _api_key = hash_token("sk-shape") + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + metadata={ + "model_rpm_limit": {"gpt-4o-mini": 50}, + "model_tpm_limit": {"gpt-4o-mini": 5000}, + }, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + async def _noop_increment(increment_list, **_): + return True + + monkeypatch.setattr( + handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + _noop_increment, + ) + + # Drive pre-call once so both paths have the same authoritative snapshot. + data: Dict[str, Any] = {"model": "gpt-4o-mini", "metadata": {}} + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=data, + call_type="", + ) + + # Non-streaming path: async_post_call_success_hook mutates response._hidden_params. + non_stream_response = ModelResponse( + id="mock-non-stream", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-4o-mini", + usage=Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2), + choices=[], + ) + non_stream_response._hidden_params = {} + await handler.async_post_call_success_hook( + data=data, + user_api_key_dict=user_api_key_dict, + response=non_stream_response, + ) + non_stream_headers = non_stream_response._hidden_params.get( + "additional_headers", {} + ) + + # Streaming path: async_logging_hook mirrors into standard_logging_object. + stream_kwargs: Dict[str, Any] = { + "standard_logging_object": { + "metadata": {"user_api_key_hash": _api_key} + }, + "litellm_params": {"metadata": data["metadata"]}, + "model": "gpt-4o-mini", + } + stream_response = ModelResponse( + id="mock-stream", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-4o-mini", + usage=Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2), + choices=[], + ) + await handler.async_logging_hook( + kwargs=stream_kwargs, + result=stream_response, + call_type="acompletion", + ) + stream_slp_headers = ( + stream_kwargs["standard_logging_object"] + .get("hidden_params", {}) + .get("additional_headers", {}) + ) + + def _rl_only(headers: Dict[str, Any]) -> Dict[str, Any]: + return {k: v for k, v in headers.items() if k.startswith("x-ratelimit-")} + + assert _rl_only(stream_slp_headers) == _rl_only(non_stream_headers), ( + f"streaming={_rl_only(stream_slp_headers)}" + f" non_streaming={_rl_only(non_stream_headers)}" + ) + assert "x-ratelimit-model_per_key-remaining-requests" in stream_slp_headers