mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-18 04:28:19 +00:00
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.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user