Resolve model group alias on Auth + /v1/messages Fallback support (#12440)

* fix(auth_checks.py): resolve a model group alias when key has access to underlying model

Fixes LIT-293

* feat(anthropic/): add mock_response to anthropic /v1/messages

makes it easy to test fallback logic

* fix(router.py): support fallbacks on /v1/messages

adds working fallbacks on generic api route

* refactor(router.py): point _ageneric_api_call_with_fallbacks to updated function

* test: add unit test for new helper on router

* fix(router.py): use correct metadata variable name

* fix(router.py): use correct metadata field

* docs(config_settings.md): document new param
This commit is contained in:
Krish Dholakia 2025-07-09 22:27:55 -07:00 • committed by GitHub
parent 0730f61127
commit 07e8609edb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 365 additions and 39 deletions

View file

@ -597,6 +597,7 @@ router_settings:
| OTEL_EXPORTER | Exporter type for OpenTelemetry
| OTEL_EXPORTER_OTLP_PROTOCOL | Exporter type for OpenTelemetry
| OTEL_HEADERS | Headers for OpenTelemetry requests
| OTEL_MODEL_ID | Model ID for OpenTelemetry tracing
| OTEL_EXPORTER_OTLP_HEADERS | Headers for OpenTelemetry requests
| OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry
| OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing

View file

@ -25,7 +25,7 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
from ..adapters.handler import LiteLLMMessagesToCompletionTransformationHandler
from .utils import AnthropicMessagesRequestUtils
from .utils import AnthropicMessagesRequestUtils, mock_response
####### ENVIRONMENT VARIABLES ###################
# Initialize any necessary instances or variables here
@ -92,6 +92,7 @@ async def anthropic_messages(
response = init_response
return response
def validate_anthropic_api_metadata(metadata: Optional[Dict] = None) -> Optional[Dict]:
"""
Validate Anthropic API metadata - This is done to ensure only allowed `metadata` fields are passed to Anthropic API
@ -103,6 +104,7 @@ def validate_anthropic_api_metadata(metadata: Optional[Dict] = None) -> Optional
anthropic_metadata_obj = AnthropicMetadata(**metadata)
return anthropic_metadata_obj.model_dump(exclude_none=True)
def anthropic_messages_handler(
max_tokens: int,
messages: List[Dict],
@ -131,12 +133,14 @@ def anthropic_messages_handler(
Makes Anthropic `/v1/messages` API calls In the Anthropic API Spec
"""
from litellm.types.utils import LlmProviders
metadata = validate_anthropic_api_metadata(metadata)
local_vars = locals()
is_async = kwargs.pop("is_async", False)
# Use provided client or create a new one
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
litellm_params = GenericLiteLLMParams(
**kwargs,
api_key=api_key,
@ -155,6 +159,15 @@ def anthropic_messages_handler(
api_key=litellm_params.api_key,
)
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
return mock_response(
model=model,
messages=messages,
max_tokens=max_tokens,
mock_response=litellm_params.mock_response,
)
anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None
if custom_llm_provider is not None and custom_llm_provider in [

View file

@ -1,6 +1,9 @@
from typing import Any, Dict, cast, get_type_hints
from typing import Any, Dict, List, cast, get_type_hints
from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
class AnthropicMessagesRequestUtils:
@ -22,3 +25,51 @@ class AnthropicMessagesRequestUtils:
k: v for k, v in params.items() if k in valid_keys and v is not None
}
return cast(AnthropicMessagesRequestOptionalParams, filtered_params)
def mock_response(
model: str,
messages: List[Dict],
max_tokens: int,
mock_response: str = "Hi! My name is Claude.",
**kwargs,
) -> AnthropicMessagesResponse:
"""
Mock response for Anthropic messages
"""
from litellm.exceptions import (
ContextWindowExceededError,
InternalServerError,
RateLimitError,
)
if mock_response == "litellm.InternalServerError":
raise InternalServerError(
message="this is a mock internal server error",
llm_provider="anthropic",
model=model,
)
elif mock_response == "litellm.ContextWindowExceededError":
raise ContextWindowExceededError(
message="this is a mock context window exceeded error",
llm_provider="anthropic",
model=model,
)
elif mock_response == "litellm.RateLimitError":
raise RateLimitError(
message="this is a mock rate limit error",
llm_provider="anthropic",
model=model,
)
return AnthropicMessagesResponse(
**{
"content": [{"text": mock_response, "type": "text"}],
"id": "msg_013Zva2CMHLNnXjNJJKqJ2EF",
"model": "claude-sonnet-4-20250514",
"role": "assistant",
"stop_reason": "end_turn",
"stop_sequence": None,
"type": "message",
"usage": {"input_tokens": 2095, "output_tokens": 503},
}
)

View file

@ -1191,6 +1191,10 @@ def _can_object_call_model(
if model in litellm.model_alias_map:
model = litellm.model_alias_map[model]
elif llm_router and model in llm_router.model_group_alias:
_model = llm_router._get_model_from_alias(model)
if _model:
model = _model
## check if model in allowed model names
from collections import defaultdict
@ -1199,6 +1203,7 @@ def _can_object_call_model(
if llm_router:
access_groups = llm_router.get_model_access_groups(model_name=model)
if (
len(access_groups) > 0 and llm_router is not None
): # check if token contains any model access groups
@ -1211,8 +1216,6 @@ def _can_object_call_model(
# Filter out models that are access_groups
filtered_models = [m for m in models if m not in access_groups]
verbose_proxy_logger.debug(f"model: {model}; allowed_models: {filtered_models}")
if _model_in_team_aliases(model=model, team_model_aliases=team_model_aliases):
return True

View file

@ -2468,30 +2468,44 @@ class Router:
self, model: str, original_function: Callable, **kwargs
):
"""
Make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
Args:
model: The model to use
handler_function: The handler function to call (e.g., litellm.anthropic_messages)
**kwargs: Additional arguments to pass to the handler function
Returns:
The response from the handler function
Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
"""
handler_name = original_function.__name__
function_name = "_ageneric_api_call_with_fallbacks"
passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False)
self._update_kwargs_before_fallbacks(
model=model,
kwargs=kwargs,
metadata_variable_name=_get_router_metadata_variable_name(
function_name=function_name
),
)
try:
verbose_router_logger.debug(
f"Inside _ageneric_api_call() - handler: {handler_name}, model: {model}; kwargs: {kwargs}"
kwargs["model"] = model
kwargs["original_generic_function"] = original_function
kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper
self._update_kwargs_before_fallbacks(
model=model, kwargs=kwargs, metadata_variable_name="litellm_metadata"
)
verbose_router_logger.debug(
f"Inside ageneric_api_call_with_fallbacks() - model: {model}; kwargs: {kwargs}"
)
response = await self.async_function_with_fallbacks(**kwargs)
return response
return response
except Exception as e:
asyncio.create_task(
send_llm_exception_alert(
litellm_router_instance=self,
request_kwargs=kwargs,
error_traceback_str=traceback.format_exc(),
original_exception=e,
)
)
raise e
async def _ageneric_api_call_with_fallbacks_helper(
self, model: str, original_generic_function: Callable, **kwargs
):
"""
Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
"""
passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False)
function_name = "_ageneric_api_call_with_fallbacks"
try:
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
try:
deployment = await self.async_get_available_deployment(
@ -2502,7 +2516,7 @@ class Router:
)
except Exception as e:
if passthrough_on_no_deployment:
return await original_function(model=model, **kwargs)
return await original_generic_function(model=model, **kwargs)
raise e
self._update_kwargs_with_deployment(
@ -2515,7 +2529,7 @@ class Router:
### get custom
response = original_function(
response = original_generic_function(
**{
**data,
"caching": self.cache_responses,
@ -2549,12 +2563,12 @@ class Router:
self.success_calls[model_name] += 1
verbose_router_logger.info(
f"{handler_name}(model={model_name})\033[32m 200 OK\033[0m"
f"ageneric_api_call_with_fallbacks(model={model_name})\033[32m 200 OK\033[0m"
)
return response
except Exception as e:
verbose_router_logger.info(
f"{handler_name}(model={model})\033[31m Exception {str(e)}\033[0m"
f"ageneric_api_call_with_fallbacks(model={model})\033[31m Exception {str(e)}\033[0m"
)
if model is not None:
self.fail_calls[model] += 1
@ -4243,6 +4257,9 @@ class Router:
When a retry or fallback happens, log the details of the just failed model call - similar to Sentry breadcrumbing
"""
try:
_metadata_var = (
"litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
)
# Log failed model as the previous model
previous_model = {
"exception_type": type(e).__name__,
@ -4254,11 +4271,11 @@ class Router:
) in (
kwargs.items()
): # log everything in kwargs except the old previous_models value - prevent nesting
if k not in ["metadata", "messages", "original_function"]:
if k not in [_metadata_var, "messages", "original_function"]:
previous_model[k] = v
elif k == "metadata" and isinstance(v, dict):
previous_model["metadata"] = {} # type: ignore
for metadata_k, metadata_v in kwargs["metadata"].items():
elif k == _metadata_var and isinstance(v, dict):
previous_model[_metadata_var] = {} # type: ignore
for metadata_k, metadata_v in kwargs[_metadata_var].items():
if metadata_k != "previous_models":
previous_model[k][metadata_k] = metadata_v # type: ignore
@ -4267,7 +4284,7 @@ class Router:
self.previous_models.pop(0)
self.previous_models.append(previous_model)
kwargs["metadata"]["previous_models"] = self.previous_models
kwargs[_metadata_var]["previous_models"] = self.previous_models
return kwargs
except Exception as e:
raise e

View file

@ -334,3 +334,39 @@ async def test_vector_store_access_check_with_permissions():
)
assert exc_info.value.type == ProxyErrorTypes.key_vector_store_access_denied
def test_can_object_call_model_with_alias():
"""Test that can_object_call_model works with model aliases"""
from litellm import Router
from litellm.proxy.auth.auth_checks import _can_object_call_model
model = "[ip-approved] gpt-4o"
llm_router = Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "test-api-key",
},
}
],
model_group_alias={
"[ip-approved] gpt-4o": {
"model": "gpt-3.5-turbo",
"hidden": True,
},
},
)
result = _can_object_call_model(
model=model,
llm_router=llm_router,
models=["gpt-3.5-turbo"],
team_model_aliases=None,
object_type="key",
fallback_depth=0,
)
print(result)

View file

@ -37,7 +37,7 @@ def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
"metadata": {"baz": 123},
}
original = copy.deepcopy(router.default_litellm_params)
kwargs = {}
kwargs: dict = {}
# invoke the helper
router._update_kwargs_with_default_litellm_params(
@ -342,7 +342,7 @@ def test_arouter_ignore_invalid_deployments():
router.upsert_deployment(
Deployment(
model_name="gpt-3.5-turbo",
litellm_params={"model": "my-bad-model"},
litellm_params={"model": "my-bad-model"}, # type: ignore
model_info={"tpm": 1000, "rpm": 1000},
)
)
@ -468,7 +468,7 @@ async def test_arouter_filter_team_based_models():
router.add_deployment(
Deployment(
model_name="gpt-3.5-turbo",
litellm_params={"model": "gpt-3.5-turbo"},
litellm_params={"model": "gpt-3.5-turbo"}, # type: ignore
model_info={"tpm": 1000, "rpm": 1000},
)
)
@ -661,16 +661,55 @@ def test_arouter_responses_api_bridge():
assert mock_post.call_args.kwargs["json"]["model"] == "webinterface-o3-pro"
@pytest.mark.asyncio
async def test_router_v1_messages_fallbacks():
"""
Test that router.v1_messages_fallbacks returns the correct response
"""
router = litellm.Router(
model_list=[
{
"model_name": "claude-3-5-sonnet-latest",
"litellm_params": {
"model": "anthropic/claude-3-5-sonnet-latest",
"mock_response": "litellm.InternalServerError",
},
},
{
"model_name": "bedrock-claude",
"litellm_params": {
"model": "anthropic.claude-3-5-sonnet-20240620-v1:0",
"mock_response": "Hello, world I am a fallback!",
},
},
],
fallbacks=[
{"claude-3-5-sonnet-latest": ["bedrock-claude"]},
],
)
result = await router.aanthropic_messages(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": "Hello, world!"}],
max_tokens=256,
)
assert result is not None
print(result)
assert result["content"][0]["text"] == "Hello, world I am a fallback!"
def test_add_invalid_provider_to_router():
"""
Test that router.add_deployment raises an error if the provider is invalid
"""
from litellm.types.router import Deployment
router = litellm.Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "gpt-3.5-turbo"},
}
],
@ -688,3 +727,169 @@ def test_add_invalid_provider_to_router():
)
assert router.pattern_router.patterns == {}
@pytest.mark.asyncio
async def test_router_ageneric_api_call_with_fallbacks_helper():
"""
Test the _ageneric_api_call_with_fallbacks_helper method with various scenarios
"""
from unittest.mock import AsyncMock, MagicMock, patch
router = litellm.Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "test-key",
"api_base": "https://api.openai.com/v1",
},
"model_info": {
"tpm": 1000,
"rpm": 1000,
},
},
],
)
# Test 1: Successful call
async def mock_generic_function(**kwargs):
return {"result": "success", "model": kwargs.get("model")}
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
mock_get_deployment.return_value = {
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "test-key",
"api_base": "https://api.openai.com/v1",
},
}
with patch.object(
router, "_update_kwargs_with_deployment"
) as mock_update_kwargs:
with patch.object(
router, "async_routing_strategy_pre_call_checks"
) as mock_pre_call_checks:
with patch.object(
router, "_get_client", return_value=None
) as mock_get_client:
result = await router._ageneric_api_call_with_fallbacks_helper(
model="gpt-3.5-turbo",
original_generic_function=mock_generic_function,
messages=[{"role": "user", "content": "test"}],
)
assert result is not None
assert result["result"] == "success"
mock_get_deployment.assert_called_once()
mock_update_kwargs.assert_called_once()
mock_pre_call_checks.assert_called_once()
# Test 2: Passthrough on no deployment (success case)
async def mock_passthrough_function(**kwargs):
return {"result": "passthrough", "model": kwargs.get("model")}
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
mock_get_deployment.side_effect = Exception("No deployment available")
result = await router._ageneric_api_call_with_fallbacks_helper(
model="gpt-3.5-turbo",
original_generic_function=mock_passthrough_function,
passthrough_on_no_deployment=True,
messages=[{"role": "user", "content": "test"}],
)
assert result is not None
assert result["result"] == "passthrough"
assert result["model"] == "gpt-3.5-turbo"
# Test 3: No deployment available and passthrough=False (should raise exception)
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
mock_get_deployment.side_effect = Exception("No deployment available")
with pytest.raises(Exception) as exc_info:
await router._ageneric_api_call_with_fallbacks_helper(
model="gpt-3.5-turbo",
original_generic_function=mock_generic_function,
passthrough_on_no_deployment=False,
messages=[{"role": "user", "content": "test"}],
)
assert "No deployment available" in str(exc_info.value)
# Test 4: Test with semaphore (rate limiting)
import asyncio
async def mock_semaphore_function(**kwargs):
return {"result": "semaphore_success", "model": kwargs.get("model")}
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
mock_get_deployment.return_value = {
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "test-key",
"api_base": "https://api.openai.com/v1",
},
}
mock_semaphore = asyncio.Semaphore(1)
with patch.object(
router, "_update_kwargs_with_deployment"
) as mock_update_kwargs:
with patch.object(
router, "_get_client", return_value=mock_semaphore
) as mock_get_client:
with patch.object(
router, "async_routing_strategy_pre_call_checks"
) as mock_pre_call_checks:
result = await router._ageneric_api_call_with_fallbacks_helper(
model="gpt-3.5-turbo",
original_generic_function=mock_semaphore_function,
messages=[{"role": "user", "content": "test"}],
)
assert result is not None
assert result["result"] == "semaphore_success"
mock_get_client.assert_called_once()
mock_pre_call_checks.assert_called_once()
# Test 5: Test call tracking (success and failure counts)
initial_success_count = router.success_calls.get("gpt-3.5-turbo", 0)
initial_fail_count = router.fail_calls.get("gpt-3.5-turbo", 0)
async def mock_failing_function(**kwargs):
raise Exception("Mock failure")
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
mock_get_deployment.return_value = {
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "test-key",
"api_base": "https://api.openai.com/v1",
},
}
with patch.object(
router, "_update_kwargs_with_deployment"
) as mock_update_kwargs:
with patch.object(
router, "_get_client", return_value=None
) as mock_get_client:
with patch.object(
router, "async_routing_strategy_pre_call_checks"
) as mock_pre_call_checks:
with pytest.raises(Exception) as exc_info:
await router._ageneric_api_call_with_fallbacks_helper(
model="gpt-3.5-turbo",
original_generic_function=mock_failing_function,
messages=[{"role": "user", "content": "test"}],
)
assert "Mock failure" in str(exc_info.value)
# Check that fail_calls was incremented
assert router.fail_calls["gpt-3.5-turbo"] == initial_fail_count + 1