From 4dc645fc334a2bce75e345433990ebaea42b1570 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 19 Mar 2026 13:59:59 +0530 Subject: [PATCH] feat(polling): check rate limits before creating polling ID Move pre-call checks (rate limits, guardrails, budget) to run BEFORE polling ID creation in the background streaming flow. This prevents the edge case where a rate-limited request receives a polling ID that immediately fails. Changes: - Add skip_pre_call_logic parameter to base_process_llm_request to allow skipping pre-call checks (avoiding double-counting of RPM/parallel requests) - Run common_processing_pre_call_logic before generating polling ID in the responses API endpoint. If rate limits/guardrails fail, return error immediately without creating a polling ID - Background streaming task passes skip_pre_call_logic=True to avoid re-running pre-call checks that were already done before polling ID creation - Add tests verifying skip_pre_call_logic parameter works correctly Fixes the edge case where polling_via_cache would return a polling ID for a request that immediately fails due to rate limiting. --- litellm/proxy/common_request_processing.py | 36 +++--- .../proxy/response_api_endpoints/endpoints.py | 32 +++++- .../response_polling/background_streaming.py | 5 +- .../test_response_polling_pre_call_checks.py | 104 ++++++++++++++++++ 4 files changed, 159 insertions(+), 18 deletions(-) create mode 100644 tests/proxy_unit_tests/test_response_polling_pre_call_checks.py diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 72765aab7d..84f9730a37 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -900,6 +900,7 @@ class ProxyBaseLLMRequestProcessing: version: Optional[str] = None, is_streaming_request: Optional[bool] = False, contents: Optional[list] = None, # Add contents parameter + skip_pre_call_logic: bool = False, ) -> Any: """ Common request processing logic for both chat completions and responses API endpoints @@ -909,22 +910,25 @@ class ProxyBaseLLMRequestProcessing: ) self._debug_log_request_payload() - self.data, logging_obj = await self.common_processing_pre_call_logic( - request=request, - general_settings=general_settings, - proxy_logging_obj=proxy_logging_obj, - user_api_key_dict=user_api_key_dict, - version=version, - proxy_config=proxy_config, - user_model=user_model, - user_temperature=user_temperature, - user_request_timeout=user_request_timeout, - user_max_tokens=user_max_tokens, - user_api_base=user_api_base, - model=model, - route_type=route_type, - llm_router=llm_router, - ) + if skip_pre_call_logic: + logging_obj = self.data.get("litellm_logging_obj") + else: + self.data, logging_obj = await self.common_processing_pre_call_logic( + request=request, + general_settings=general_settings, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + version=version, + proxy_config=proxy_config, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + model=model, + route_type=route_type, + llm_router=llm_router, + ) tasks = [] # Start the moderation check (during_call_hook) as early as possible diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index e9c7cce0d7..055fdeb84f 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -119,6 +119,34 @@ async def responses_api( f"Starting background response with polling for model={data.get('model')}" ) + # Run pre-call checks (rate limits, guardrails, budget) BEFORE creating + # polling ID. This ensures rate-limited requests get a synchronous 429 + # instead of a polling ID that immediately fails in the background task. + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + data, _logging_obj = await processor.common_processing_pre_call_logic( + request=request, + general_settings=general_settings, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + version=version, + proxy_config=proxy_config, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + route_type="aresponses", + llm_router=llm_router, + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) + # Initialize polling handler with configured TTL (from global config) polling_handler = ResponsePollingHandler( redis_cache=redis_usage_cache, @@ -134,7 +162,9 @@ async def responses_api( request_data=data, ) - # Start background task to stream and update cache + # Start background task to stream and update cache. + # Pass pre-processed data so the background task skips pre-call logic + # (rate limits, guardrails already checked above). asyncio.create_task( background_streaming_task( polling_id=polling_id, diff --git a/litellm/proxy/response_polling/background_streaming.py b/litellm/proxy/response_polling/background_streaming.py index 7583f30eb2..bcc9817577 100644 --- a/litellm/proxy/response_polling/background_streaming.py +++ b/litellm/proxy/response_polling/background_streaming.py @@ -65,7 +65,9 @@ async def background_streaming_task( # noqa: PLR0915 # Create processor processor = ProxyBaseLLMRequestProcessing(data=data) - # Make streaming request + # Make streaming request. + # Pre-call checks (rate limits, guardrails, budget) were already run + # before polling ID creation, so skip them here to avoid double-counting. response = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, @@ -83,6 +85,7 @@ async def background_streaming_task( # noqa: PLR0915 user_max_tokens=user_max_tokens, user_api_base=user_api_base, version=version, + skip_pre_call_logic=True, ) # Process streaming response following OpenAI events format diff --git a/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py b/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py new file mode 100644 index 0000000000..b39f1bf43d --- /dev/null +++ b/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py @@ -0,0 +1,104 @@ +""" +Unit tests for pre-call checks running before polling ID creation. + +Tests that rate limits, guardrails, and budget checks are enforced +BEFORE a polling ID is created, so rate-limited requests get a +synchronous error instead of a polling ID that immediately fails. +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import Request, Response + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + +class TestSkipPreCallLogic: + """Test that skip_pre_call_logic parameter works correctly""" + + @pytest.mark.asyncio + async def test_skip_pre_call_logic_skips_common_processing(self): + """When skip_pre_call_logic=True, common_processing_pre_call_logic should not be called""" + mock_logging_obj = MagicMock() + data = { + "model": "gpt-4", + "stream": True, + "litellm_logging_obj": mock_logging_obj, + } + processor = ProxyBaseLLMRequestProcessing(data=data) + + mock_proxy_logging = AsyncMock() + mock_proxy_logging.during_call_hook = AsyncMock() + + with ( + patch.object( + processor, "common_processing_pre_call_logic", new_callable=AsyncMock + ) as mock_pre_call, + patch( + "litellm.proxy.common_request_processing.route_request", + new_callable=AsyncMock, + return_value=MagicMock(), + ), + ): + try: + await processor.base_process_llm_request( + request=MagicMock(spec=Request), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + route_type="aresponses", + proxy_logging_obj=mock_proxy_logging, + llm_router=MagicMock(), + general_settings={}, + proxy_config=MagicMock(), + skip_pre_call_logic=True, + ) + except Exception: + pass # We only care that common_processing_pre_call_logic was not called + + mock_pre_call.assert_not_called() + + @pytest.mark.asyncio + async def test_without_skip_runs_common_processing(self): + """When skip_pre_call_logic=False (default), common_processing_pre_call_logic should be called""" + data = {"model": "gpt-4"} + processor = ProxyBaseLLMRequestProcessing(data=data) + + mock_logging_obj = MagicMock() + mock_proxy_logging = AsyncMock() + mock_proxy_logging.during_call_hook = AsyncMock() + + with ( + patch.object( + processor, + "common_processing_pre_call_logic", + new_callable=AsyncMock, + return_value=(data, mock_logging_obj), + ) as mock_pre_call, + patch( + "litellm.proxy.common_request_processing.route_request", + new_callable=AsyncMock, + ), + ): + try: + await processor.base_process_llm_request( + request=MagicMock(spec=Request), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + route_type="aresponses", + proxy_logging_obj=mock_proxy_logging, + llm_router=MagicMock(), + general_settings={}, + proxy_config=MagicMock(), + ) + except Exception: + pass + + mock_pre_call.assert_called_once() + +