diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index b749ddf2562..6e2c5624367 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -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 diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 58e40ce1b1b..cc9334ae68b 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -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 [ diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py index 29d00cd04cc..fa951ebd2e5 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py @@ -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}, + } + ) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 2cb60649890..cf637ce1087 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 5b64f8b1ee5..879a3a5c89b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 9b7af1859e8..635e39c37f4 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 8a15424d5be..e7201723df3 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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 \ No newline at end of file