This commit is contained in:
devin-ai-integration[bot] 2026-09-12 08:25:25 -04:00 committed by GitHub
commit f53ce4b4a0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 270 additions and 18 deletions

View file

@ -32,7 +32,11 @@ from litellm.constants import (
UNSAFE_PROXY_RESPONSE_HEADERS,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
from litellm.litellm_core_utils.core_helpers import (
get_or_create_metadata_bucket,
independent_snapshot,
is_expected_client_error,
)
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
@ -2034,6 +2038,21 @@ class ProxyBaseLLMRequestProcessing:
) -> tuple[dict, LiteLLMLoggingObj]:
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
original_model: Final = self.data.get("model")
fallback_models: Final = (
self._resolve_fallback_models(
model=original_model,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
if original_model
and isinstance(original_model, str)
and llm_router
and not self.data.get("disable_fallbacks")
else None
)
pristine: Final = independent_snapshot(self.data) if fallback_models else None
try:
return await self.common_processing_pre_call_logic(
request=request,
@ -2052,16 +2071,7 @@ class ProxyBaseLLMRequestProcessing:
llm_router=llm_router,
)
except ProxyRateLimitError as original_exc:
original_model: Final = self.data.get("model")
if not original_model or not llm_router or self.data.get("disable_fallbacks"):
raise
fallback_models: Final = self._resolve_fallback_models(
model=original_model,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
if not fallback_models:
if not fallback_models or pristine is None:
raise
verbose_proxy_logger.info(
@ -2074,7 +2084,7 @@ class ProxyBaseLLMRequestProcessing:
for fallback_model in fallback_models:
if fallback_model == original_model:
continue
self.data["model"] = fallback_model
self.data = {**independent_snapshot(pristine), "model": fallback_model}
try:
return await self.common_processing_pre_call_logic(
request=request,
@ -2095,10 +2105,10 @@ class ProxyBaseLLMRequestProcessing:
except ProxyRateLimitError:
continue
except BaseException:
self.data["model"] = original_model
self.data = pristine
raise
self.data["model"] = original_model
self.data = pristine
raise original_exc
def _resolve_fallback_models(

View file

@ -48,10 +48,10 @@ class TestSkipPreCallLogic:
await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth, router_settings=None),
route_type="aresponses",
proxy_logging_obj=mock_proxy_logging,
llm_router=MagicMock(),
llm_router=MagicMock(fallbacks=None),
general_settings={},
proxy_config=MagicMock(),
skip_pre_call_logic=True,
@ -87,10 +87,10 @@ class TestSkipPreCallLogic:
await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth, router_settings=None),
route_type="aresponses",
proxy_logging_obj=mock_proxy_logging,
llm_router=MagicMock(),
llm_router=MagicMock(fallbacks=None),
general_settings={},
proxy_config=MagicMock(),
)

View file

@ -6235,6 +6235,248 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
call_type="acompletion",
)
@pytest.mark.asyncio
async def test_fallback_retries_from_pristine_request_data(self):
import threading
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model = "gpt-4"
fallback_model = "gpt-3.5-turbo"
processor = ProxyBaseLLMRequestProcessing(
data={
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"tags": ["a"]},
}
)
metadata_at_entry = []
async def mock_pre_call_logic(**kwargs):
copy.deepcopy(processor.data["metadata"])
metadata_at_entry.append(dict(processor.data["metadata"]))
processor.data["metadata"]["litellm_parent_otel_span"] = threading.RLock()
processor.data["litellm_logging_obj"] = object()
if processor.data.get("model") == primary_model:
raise ProxyRateLimitError(
detail="TPM limit exceeded for gpt-4",
headers={"retry-after": "30"},
)
return processor.data, MagicMock()
mock_router = MagicMock()
mock_router.fallbacks = [{primary_model: [fallback_model]}]
with patch.object(
processor,
"common_processing_pre_call_logic",
side_effect=mock_pre_call_logic,
):
data, logging_obj = await processor._pre_call_with_fallbacks(
request=MagicMock(),
general_settings={},
proxy_logging_obj=MagicMock(),
user_api_key_dict=MagicMock(router_settings=None),
version=None,
proxy_config=MagicMock(),
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
model=primary_model,
route_type="acompletion",
llm_router=mock_router,
)
assert processor.data["model"] == fallback_model
assert metadata_at_entry[1] == {"tags": ["a"]}
@pytest.mark.asyncio
async def test_exhausted_fallbacks_restore_pristine_request_data(self):
import threading
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model = "gpt-4"
original_data = {
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"tags": ["a"]},
}
processor = ProxyBaseLLMRequestProcessing(data=copy.deepcopy(original_data))
async def mock_pre_call_logic(**kwargs):
processor.data["metadata"]["litellm_parent_otel_span"] = threading.RLock()
processor.data["litellm_logging_obj"] = object()
raise ProxyRateLimitError(
detail=f"TPM limit exceeded for {processor.data.get('model')}",
headers={"retry-after": "30"},
)
mock_router = MagicMock()
mock_router.fallbacks = [{primary_model: ["gpt-3.5-turbo"]}]
with patch.object(
processor,
"common_processing_pre_call_logic",
side_effect=mock_pre_call_logic,
):
with pytest.raises(ProxyRateLimitError, match="gpt-4"):
await processor._pre_call_with_fallbacks(
request=MagicMock(),
general_settings={},
proxy_logging_obj=MagicMock(),
user_api_key_dict=MagicMock(router_settings=None),
version=None,
proxy_config=MagicMock(),
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
model=primary_model,
route_type="acompletion",
llm_router=mock_router,
)
assert processor.data == original_data
@pytest.mark.asyncio
async def test_real_add_litellm_data_to_request_rerun_with_otel_span_falls_back(self):
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import ProxyConfig
trace.set_tracer_provider(TracerProvider())
primary_model = "gpt-4"
fallback_model = "gpt-3.5-turbo"
request_mock = MagicMock(spec=Request)
request_mock.url = MagicMock()
request_mock.url.path = "/v1/chat/completions"
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
user_api_key_dict = UserAPIKeyAuth(
parent_otel_span=trace.get_tracer("x").start_span("s"),
api_key="hashed-key",
user_id="u1",
team_id="t1",
metadata={},
team_metadata={},
team_member_tpm_limit=1000,
)
processor = ProxyBaseLLMRequestProcessing(
data={
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"tags": ["a"]},
}
)
async def real_add_litellm_data_pre_call(**kwargs):
await add_litellm_data_to_request(
data=processor.data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=ProxyConfig(),
general_settings={},
version="test",
)
if processor.data.get("model") == primary_model:
raise ProxyRateLimitError(
detail="TPM limit exceeded for gpt-4",
headers={"retry-after": "30"},
)
return processor.data, MagicMock()
mock_router = MagicMock()
mock_router.fallbacks = [{primary_model: [fallback_model]}]
with patch.object(
processor,
"common_processing_pre_call_logic",
side_effect=real_add_litellm_data_pre_call,
):
data, logging_obj = await processor._pre_call_with_fallbacks(
request=request_mock,
general_settings={},
proxy_logging_obj=MagicMock(),
user_api_key_dict=user_api_key_dict,
version=None,
proxy_config=MagicMock(),
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
model=primary_model,
route_type="acompletion",
llm_router=mock_router,
)
assert processor.data["model"] == fallback_model
@pytest.mark.asyncio
async def test_no_fallbacks_skips_snapshot(self):
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4"})
async def mock_pre_call_logic(**kwargs):
raise ProxyRateLimitError(
detail="TPM limit exceeded",
headers={"retry-after": "30"},
)
mock_router = MagicMock()
mock_router.fallbacks = None
with patch( # test-quality-ok: spying the snapshot seam is the only observable check that the no-fallback path skips it
"litellm.proxy.common_request_processing.independent_snapshot"
) as snapshot_mock:
with patch.object(
processor,
"common_processing_pre_call_logic",
side_effect=mock_pre_call_logic,
):
with pytest.raises(ProxyRateLimitError):
await processor._pre_call_with_fallbacks(
request=MagicMock(),
general_settings={},
proxy_logging_obj=MagicMock(),
user_api_key_dict=MagicMock(router_settings=None),
version=None,
proxy_config=MagicMock(),
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
model="gpt-4",
route_type="acompletion",
llm_router=mock_router,
)
snapshot_mock.assert_not_called()
class _RecordingSuccessLogger(CustomLogger):
def __init__(self):