diff --git a/litellm/llms/azure/chat/gpt_transformation.py b/litellm/llms/azure/chat/gpt_transformation.py index 78fa91bbd7..0ae6fad730 100644 --- a/litellm/llms/azure/chat/gpt_transformation.py +++ b/litellm/llms/azure/chat/gpt_transformation.py @@ -159,9 +159,16 @@ class AzureOpenAIConfig(BaseConfig): supported_openai_params = self.get_supported_openai_params(model) api_version_times = api_version.split("-") - api_version_year = api_version_times[0] - api_version_month = api_version_times[1] - api_version_day = api_version_times[2] + + if len(api_version_times) >= 3: + api_version_year = api_version_times[0] + api_version_month = api_version_times[1] + api_version_day = api_version_times[2] + else: + api_version_year = None + api_version_month = None + api_version_day = None + for param, value in non_default_params.items(): if param == "tool_choice": """ @@ -171,47 +178,57 @@ class AzureOpenAIConfig(BaseConfig): """ ## check if api version supports this param ## if ( - api_version_year < "2023" - or (api_version_year == "2023" and api_version_month < "12") - or ( - api_version_year == "2023" - and api_version_month == "12" - and api_version_day < "01" - ) + api_version_year is None + or api_version_month is None + or api_version_day is None ): - if litellm.drop_params is True or ( - drop_params is not None and drop_params is True - ): - pass - else: - raise UnsupportedParamsError( - status_code=400, - message=f"""Azure does not support 'tool_choice', for api_version={api_version}. Bump your API version to '2023-12-01-preview' or later. This parameter requires 'api_version="2023-12-01-preview"' or later. Azure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions""", - ) - elif value == "required" and ( - api_version_year == "2024" and api_version_month <= "05" - ): ## check if tool_choice value is supported ## - if litellm.drop_params is True or ( - drop_params is not None and drop_params is True - ): - pass - else: - raise UnsupportedParamsError( - status_code=400, - message=f"Azure does not support '{value}' as a {param} param, for api_version={api_version}. To drop 'tool_choice=required' for calls with this Azure API version, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\nAzure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions", - ) - else: optional_params["tool_choice"] = value + else: + if ( + api_version_year < "2023" + or (api_version_year == "2023" and api_version_month < "12") + or ( + api_version_year == "2023" + and api_version_month == "12" + and api_version_day < "01" + ) + ): + if litellm.drop_params is True or ( + drop_params is not None and drop_params is True + ): + pass + else: + raise UnsupportedParamsError( + status_code=400, + message=f"""Azure does not support 'tool_choice', for api_version={api_version}. Bump your API version to '2023-12-01-preview' or later. This parameter requires 'api_version="2023-12-01-preview"' or later. Azure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions""", + ) + elif value == "required" and ( + api_version_year == "2024" and api_version_month <= "05" + ): ## check if tool_choice value is supported ## + if litellm.drop_params is True or ( + drop_params is not None and drop_params is True + ): + pass + else: + raise UnsupportedParamsError( + status_code=400, + message=f"Azure does not support '{value}' as a {param} param, for api_version={api_version}. To drop 'tool_choice=required' for calls with this Azure API version, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\nAzure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions", + ) + else: + optional_params["tool_choice"] = value elif param == "response_format" and isinstance(value, dict): _is_response_format_supported_model = ( self._is_response_format_supported_model(model) ) - is_response_format_supported_api_version = ( - self._is_response_format_supported_api_version( - api_version_year, api_version_month + if api_version_year is None or api_version_month is None: + is_response_format_supported_api_version = True + else: + is_response_format_supported_api_version = ( + self._is_response_format_supported_api_version( + api_version_year, api_version_month + ) ) - ) is_response_format_supported = ( is_response_format_supported_api_version and _is_response_format_supported_model diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index bd800bb0db..0000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index c9d8c4f7f5..af11bd8aa5 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,18 +1,5 @@ model_list: - - model_name: claude-sonnet-4 + - model_name: bedrock-claude-3.7-sonnet litellm_params: - model: anthropic/claude-sonnet-4-20250514 - api_key: os.environ/ANTHROPIC_API_KEY - api_base: https://api.anthropic.com/v1 - guardrails: ["azure-text-moderation"] - - model_name: openai-gpt-4o - litellm_params: - model: openai/gpt-4o + model: bedrock/us.anthropic.claude-3-7-sonnet-20250219-v1:0 -guardrails: - - guardrail_name: azure-text-moderation - litellm_params: - guardrail: azure/text_moderations - mode: "post_call" - api_key: os.environ/AZURE_GUARDRAIL_API_KEY - api_base: os.environ/AZURE_GUARDRAIL_API_BASE diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index 0f1f215139..a564419c79 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -1,6 +1,6 @@ import json from datetime import datetime -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import TYPE_CHECKING, Any, List, Optional, Union, cast import httpx @@ -96,10 +96,12 @@ class AnthropicPassthroughLoggingHandler: handles streaming and non-streaming responses """ try: + response_cost = litellm.completion_cost( completion_response=litellm_model_response, model=model, ) + kwargs["response_cost"] = response_cost kwargs["model"] = model passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( # type: ignore @@ -125,9 +127,10 @@ class AnthropicPassthroughLoggingHandler: litellm_model_response.id = logging_obj.litellm_call_id litellm_model_response.model = model logging_obj.model_call_details["model"] = model - logging_obj.model_call_details["custom_llm_provider"] = ( - litellm.LlmProviders.ANTHROPIC.value - ) + if not logging_obj.model_call_details.get("custom_llm_provider"): + logging_obj.model_call_details["custom_llm_provider"] = ( + litellm.LlmProviders.ANTHROPIC.value + ) return kwargs except Exception as e: verbose_proxy_logger.exception( @@ -155,6 +158,19 @@ class AnthropicPassthroughLoggingHandler: """ model = request_body.get("model", "") + # Dheck if it's available in the logging object + if ( + not model + and hasattr(litellm_logging_obj, "model_call_details") + and litellm_logging_obj.model_call_details.get("model") + ): + model = cast(str, litellm_logging_obj.model_call_details.get("model")) + custom_llm_provider = litellm_logging_obj.model_call_details.get( + "custom_llm_provider" + ) + + if custom_llm_provider and not model.startswith(custom_llm_provider): + model = f"{custom_llm_provider}/{model}" complete_streaming_response = ( AnthropicPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index 6dc0c4a83e..098daf7893 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -25,6 +25,7 @@ from litellm.router import Router import importlib from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from base_anthropic_unified_messages_test import BaseAnthropicMessagesTest + # Load environment variables load_dotenv() @@ -70,59 +71,58 @@ def _validate_anthropic_response(response: Dict[str, Any]): assert response["role"] == "assistant" - class TestAnthropicDirectAPI(BaseAnthropicMessagesTest): """Tests for direct Anthropic API calls""" + @property def model_config(self) -> Dict[str, Any]: return { "model": "claude-3-haiku-20240307", "api_key": os.getenv("ANTHROPIC_API_KEY"), } - + @property def expected_model_name_in_logging(self) -> str: """ This is the model name that is expected to be in the logging payload """ return "claude-3-haiku-20240307" - - + class TestAnthropicBedrockAPI(BaseAnthropicMessagesTest): """Tests for Anthropic via Bedrock""" + @property def model_config(self) -> Dict[str, Any]: return { "model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", } - @property def expected_model_name_in_logging(self) -> str: """ This is the model name that is expected to be in the logging payload """ - return "us.anthropic.claude-3-5-sonnet-20240620-v1:0" - + return "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" class TestAnthropicOpenAIAPI(BaseAnthropicMessagesTest): """Tests for OpenAI via Anthropic messages interface""" + @property def model_config(self) -> Dict[str, Any]: return { "model": "openai/gpt-4o-mini", "client": None, } - + @property def expected_model_name_in_logging(self) -> str: """ This is the model name that is expected to be in the logging payload """ return "gpt-4o-mini" - + @pytest.mark.asyncio async def test_anthropic_messages_litellm_router_streaming_with_logging(self): """ @@ -151,8 +151,8 @@ async def test_anthropic_messages_streaming_with_bad_request(): except Exception as e: print("got exception", e) print("vars", vars(e)) - if hasattr(e, 'status_code'): - assert getattr(e, 'status_code') == 400 + if hasattr(e, "status_code"): + assert getattr(e, "status_code") == 400 else: assert isinstance(e, Exception) @@ -188,8 +188,8 @@ async def test_anthropic_messages_router_streaming_with_bad_request(): except Exception as e: print("got exception", e) print("vars", vars(e)) - if hasattr(e, 'status_code'): - assert getattr(e, 'status_code') == 400 + if hasattr(e, "status_code"): + assert getattr(e, "status_code") == 400 else: assert isinstance(e, Exception) @@ -231,6 +231,7 @@ async def test_anthropic_messages_litellm_router_non_streaming(): print(f"Non-streaming response: {json.dumps(response, indent=2)}") return response + @pytest.mark.asyncio async def test_anthropic_messages_litellm_router_routing_strategy(): """ @@ -260,7 +261,7 @@ async def test_anthropic_messages_litellm_router_routing_strategy(): max_tokens=100, metadata={ "user_id": "hello", - } + }, ) # Verify response @@ -276,10 +277,10 @@ async def test_anthropic_messages_litellm_router_routing_strategy(): @pytest.mark.asyncio async def test_anthropic_messages_litellm_router_latency_metadata_tracking(): """ - Test the anthropic_messages with routing strategy and verify that _latency_per_deployment + Test the anthropic_messages with routing strategy and verify that _latency_per_deployment field is passed in litellm_metadata when calling litellm.anthropic_messages """ - with unittest.mock.patch('litellm.anthropic_messages') as mock_anthropic_messages: + with unittest.mock.patch("litellm.anthropic_messages") as mock_anthropic_messages: # Mock the return value mock_response = { "id": "msg_123456", @@ -293,7 +294,7 @@ async def test_anthropic_messages_litellm_router_latency_metadata_tracking(): mock_anthropic_messages.return_value = mock_response # Set the __name__ attribute that the router expects mock_anthropic_messages.__name__ = "anthropic_messages" - + MODEL_GROUP = "claude-special-alias" router = Router( model_list=[ @@ -318,51 +319,61 @@ async def test_anthropic_messages_litellm_router_latency_metadata_tracking(): max_tokens=100, metadata={ "user_id": "hello", - } + }, ) # Verify response assert response == mock_response - + # Verify that litellm.anthropic_messages was called mock_anthropic_messages.assert_called_once() - + # Get the call arguments call_args = mock_anthropic_messages.call_args call_kwargs = call_args.kwargs - + print("Call kwargs:", json.dumps(call_kwargs, indent=2, default=str)) - + # Verify that litellm_metadata was passed and contains _latency_per_deployment - assert "litellm_metadata" in call_kwargs, "litellm_metadata should be passed to anthropic_messages" - + assert ( + "litellm_metadata" in call_kwargs + ), "litellm_metadata should be passed to anthropic_messages" + litellm_metadata = call_kwargs["litellm_metadata"] assert litellm_metadata is not None, "litellm_metadata should not be None" - assert isinstance(litellm_metadata, dict), "litellm_metadata should be a dictionary" - + assert isinstance( + litellm_metadata, dict + ), "litellm_metadata should be a dictionary" + # Verify _latency_per_deployment is present - assert "_latency_per_deployment" in litellm_metadata, "litellm_metadata should contain _latency_per_deployment field" - + assert ( + "_latency_per_deployment" in litellm_metadata + ), "litellm_metadata should contain _latency_per_deployment field" + # Verify the structure of _latency_per_deployment latency_per_deployment = litellm_metadata["_latency_per_deployment"] - assert isinstance(latency_per_deployment, dict), "_latency_per_deployment should be a dictionary" - + assert isinstance( + latency_per_deployment, dict + ), "_latency_per_deployment should be a dictionary" + print(f"✅ Latency per deployment data: {latency_per_deployment}") - + # Verify other expected fields in litellm_metadata assert "model_group" in litellm_metadata assert litellm_metadata["model_group"] == MODEL_GROUP assert "deployment" in litellm_metadata assert "model_info" in litellm_metadata - + # Verify other call parameters assert call_kwargs["model"] == "claude-3-haiku-20240307" assert call_kwargs["messages"] == messages assert call_kwargs["max_tokens"] == 100 assert call_kwargs["metadata"] == {"user_id": "hello"} - - print("✅ Successfully verified that _latency_per_deployment is passed in litellm_metadata to anthropic_messages") - + + print( + "✅ Successfully verified that _latency_per_deployment is passed in litellm_metadata to anthropic_messages" + ) + return response @@ -417,9 +428,16 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging(): print(f"Non-streaming response: {json.dumps(response, indent=2)}") await asyncio.sleep(1) - - assert test_custom_logger.logged_standard_logging_payload is not None, "Logging payload should not be None" - print("tracked standard logging payload", json.dumps(test_custom_logger.logged_standard_logging_payload, indent=4, default=str)) + + assert ( + test_custom_logger.logged_standard_logging_payload is not None + ), "Logging payload should not be None" + print( + "tracked standard logging payload", + json.dumps( + test_custom_logger.logged_standard_logging_payload, indent=4, default=str + ), + ) assert test_custom_logger.logged_standard_logging_payload["messages"] == messages assert test_custom_logger.logged_standard_logging_payload["response"] is not None assert ( @@ -439,8 +457,9 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging(): ) # assert model_group - assert test_custom_logger.logged_standard_logging_payload["model_group"] == MODEL_GROUP - + assert ( + test_custom_logger.logged_standard_logging_payload["model_group"] == MODEL_GROUP + ) @pytest.mark.asyncio @@ -510,7 +529,6 @@ async def test_anthropic_messages_with_extra_headers(): return response - @pytest.mark.asyncio async def test_anthropic_messages_with_thinking(): """ @@ -522,7 +540,6 @@ async def test_anthropic_messages_with_thinking(): # Set up test parameters messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - # Create a mock response mock_response = MagicMock() mock_response.raise_for_status = MagicMock() @@ -568,7 +585,6 @@ async def test_anthropic_messages_with_thinking(): assert request_body["messages"] == messages assert request_body["thinking"] == {"budget_tokens": 100} - # Verify the response was processed correctly assert response == mock_response.json.return_value @@ -582,18 +598,22 @@ async def test_anthropic_messages_bedrock_credentials_passthrough(): when using anthropic.messages.acreate with a bedrock model """ # Mock the get_credentials method - with unittest.mock.patch.object(BaseAWSLLM, 'get_credentials') as mock_get_credentials: + with unittest.mock.patch.object( + BaseAWSLLM, "get_credentials" + ) as mock_get_credentials: # Create a proper mock for credentials with the necessary attributes mock_credentials = unittest.mock.MagicMock() mock_credentials.access_key = "mock_access_key" mock_credentials.secret_key = "mock_secret_key" mock_credentials.token = "mock_session_token" mock_get_credentials.return_value = mock_credentials - + # We also need to mock the actual AWS request signing to avoid real API calls - with unittest.mock.patch('botocore.auth.SigV4Auth.add_auth'): + with unittest.mock.patch("botocore.auth.SigV4Auth.add_auth"): # Set up mock for AsyncHTTPHandler.post to avoid actual API calls - with unittest.mock.patch('litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post') as mock_post: + with unittest.mock.patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post" + ) as mock_post: # Configure mock response mock_response = unittest.mock.MagicMock() mock_response.raise_for_status = unittest.mock.MagicMock() @@ -607,7 +627,7 @@ async def test_anthropic_messages_bedrock_credentials_passthrough(): "usage": {"input_tokens": 10, "output_tokens": 20}, } mock_post.return_value = mock_response - + # Test AWS credentials parameters - separate from function call parameters aws_params = { "aws_access_key_id": "test_access_key", @@ -620,7 +640,7 @@ async def test_anthropic_messages_bedrock_credentials_passthrough(): "aws_web_identity_token": "test_web_identity_token", "aws_sts_endpoint": "https://sts.test-region.amazonaws.com", } - + # Call the function with AWS credentials await litellm.anthropic.messages.acreate( messages=[{"role": "user", "content": "Hello, test credentials"}], @@ -628,15 +648,16 @@ async def test_anthropic_messages_bedrock_credentials_passthrough(): max_tokens=100, **aws_params, ) - + # Verify get_credentials was called with the correct parameters mock_get_credentials.assert_called_once() call_args = mock_get_credentials.call_args[1] - + # Assert that our test credentials were passed correctly for param_name, param_value in aws_params.items(): - assert call_args[param_name] == param_value, f"Parameter {param_name} was not passed correctly" - + assert ( + call_args[param_name] == param_value + ), f"Parameter {param_name} was not passed correctly" @pytest.mark.asyncio @@ -662,19 +683,22 @@ async def test_anthropic_messages_bedrock_dynamic_region(): mock_client.post = AsyncMock(return_value=mock_response) # Patch necessary AWS components - with unittest.mock.patch('botocore.auth.SigV4Auth.add_auth'), \ - unittest.mock.patch.object(BaseAWSLLM, 'get_credentials') as mock_get_credentials: - + with unittest.mock.patch( + "botocore.auth.SigV4Auth.add_auth" + ), unittest.mock.patch.object( + BaseAWSLLM, "get_credentials" + ) as mock_get_credentials: + # Setup mock credentials mock_credentials = unittest.mock.MagicMock() mock_credentials.access_key = "test_access_key" mock_credentials.secret_key = "test_secret_key" mock_credentials.token = "test_session_token" mock_get_credentials.return_value = mock_credentials - + # Test with specific region test_region = "us-east-1" - + # Call anthropic.messages.acreate with aws_region_name response = await litellm.anthropic.messages.acreate( messages=[{"role": "user", "content": "Hello, test region"}], @@ -683,22 +707,24 @@ async def test_anthropic_messages_bedrock_dynamic_region(): aws_region_name=test_region, client=mock_client, ) - + # Verify response assert response == mock_response.json.return_value - + # Verify the post method was called with the correct URL containing the region mock_client.post.assert_called_once() call_args = mock_client.post.call_args - + # Check that the URL contains the correct region - url = call_args.kwargs.get('url', '') - assert f"bedrock-runtime.{test_region}.amazonaws.com" in url, f"URL does not contain the correct region. URL: {url}" - + url = call_args.kwargs.get("url", "") + assert ( + f"bedrock-runtime.{test_region}.amazonaws.com" in url + ), f"URL does not contain the correct region. URL: {url}" + # Verify get_credentials was called with the correct region mock_get_credentials.assert_called_once() credentials_args = mock_get_credentials.call_args.kwargs - assert credentials_args.get('aws_region_name') == test_region + assert credentials_args.get("aws_region_name") == test_region def test_sync_openai_messages(): @@ -716,4 +742,3 @@ def test_sync_openai_messages(): assert response is not None assert isinstance(response, dict) assert response["content"][0].text is not None - diff --git a/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py b/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py index ac873710d4..8df35a3751 100644 --- a/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py +++ b/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py @@ -28,3 +28,17 @@ class TestAzureOpenAIConfig: assert not config._is_response_format_supported_model("gpt-3-5-turbo-suffix") assert not config._is_response_format_supported_model("gpt-35-turbo-suffix") assert not config._is_response_format_supported_model("gpt-35-turbo") + + +def test_map_openai_params_with_preview_api_version(): + config = AzureOpenAIConfig() + non_default_params = { + "response_format": {"type": "json_object"}, + } + optional_params = {} + model = "azure/gpt-4-1" + drop_params = False + api_version = "preview" + assert config.map_openai_params( + non_default_params, optional_params, model, drop_params, api_version + ) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py new file mode 100644 index 0000000000..e54e537eed --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -0,0 +1,154 @@ +import json +import os +import sys +from datetime import datetime +from typing import Any, Dict, List +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( + AnthropicPassthroughLoggingHandler, +) + + +class TestAnthropicLoggingHandlerModelFallback: + """Test the model fallback logic in the anthropic passthrough logging handler.""" + + def setup_method(self): + """Set up test fixtures""" + self.start_time = datetime.now() + self.end_time = datetime.now() + self.mock_chunks = [ + '{"type": "message_start", "message": {"id": "msg_123", "model": "claude-3-haiku-20240307"}}', + '{"type": "content_block_delta", "delta": {"text": "Hello"}}', + '{"type": "content_block_delta", "delta": {"text": " world"}}', + '{"type": "message_stop"}', + ] + + def _create_mock_logging_obj(self, model_in_details: str = None) -> LiteLLMLoggingObj: + """Create a mock logging object with optional model in model_call_details""" + mock_logging_obj = MagicMock() + + if model_in_details: + # Create a dict-like mock that returns the model for the 'model' key + mock_model_call_details = {'model': model_in_details} + mock_logging_obj.model_call_details = mock_model_call_details + else: + # Create empty dict or None + mock_logging_obj.model_call_details = {} + + return mock_logging_obj + + def _create_mock_passthrough_handler(self): + """Create a mock passthrough success handler""" + mock_handler = MagicMock() + return mock_handler + + + + @patch.object(AnthropicPassthroughLoggingHandler, '_build_complete_streaming_response') + @patch.object(AnthropicPassthroughLoggingHandler, '_create_anthropic_response_logging_payload') + def test_model_from_request_body_used_when_present(self, mock_create_payload, mock_build_response): + """Test that model from request_body is used when present""" + # Arrange + request_body = {"model": "claude-3-sonnet-20240229"} + logging_obj = self._create_mock_logging_obj(model_in_details="claude-3-haiku-20240307") + passthrough_handler = self._create_mock_passthrough_handler() + + # Mock successful response building + mock_build_response.return_value = MagicMock() + mock_create_payload.return_value = {"test": "payload"} + + # Act + result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks( + litellm_logging_obj=logging_obj, + passthrough_success_handler_obj=passthrough_handler, + url_route="/anthropic/v1/messages", + request_body=request_body, + endpoint_type="messages", + start_time=self.start_time, + all_chunks=self.mock_chunks, + end_time=self.end_time, + ) + + # Assert + assert result is not None + # Verify that _build_complete_streaming_response was called with the request_body model + mock_build_response.assert_called_once() + call_args = mock_build_response.call_args + assert call_args[1]['model'] == "claude-3-sonnet-20240229" # Should use request_body model + + def test_model_fallback_logic_isolated(self): + """Test just the model fallback logic in isolation""" + # Test case 1: Model from request body + request_body = {"model": "claude-3-sonnet-20240229"} + logging_obj = self._create_mock_logging_obj(model_in_details="claude-3-haiku-20240307") + + # Extract the logic directly from the function + model = request_body.get("model", "") + if not model and hasattr(logging_obj, 'model_call_details') and logging_obj.model_call_details.get('model'): + model = logging_obj.model_call_details.get('model') + + assert model == "claude-3-sonnet-20240229" # Should use request_body model + + # Test case 2: Fallback to logging obj + request_body = {} + logging_obj = self._create_mock_logging_obj(model_in_details="claude-3-haiku-20240307") + + model = request_body.get("model", "") + if not model and hasattr(logging_obj, 'model_call_details') and logging_obj.model_call_details.get('model'): + model = logging_obj.model_call_details.get('model') + + assert model == "claude-3-haiku-20240307" # Should use fallback model + + # Test case 3: Empty string in request body, fallback to logging obj + request_body = {"model": ""} + logging_obj = self._create_mock_logging_obj(model_in_details="claude-3-opus-20240229") + + model = request_body.get("model", "") + if not model and hasattr(logging_obj, 'model_call_details') and logging_obj.model_call_details.get('model'): + model = logging_obj.model_call_details.get('model') + + assert model == "claude-3-opus-20240229" # Should use fallback model + + # Test case 4: Both empty + request_body = {} + logging_obj = self._create_mock_logging_obj() + + model = request_body.get("model", "") + if not model and hasattr(logging_obj, 'model_call_details') and logging_obj.model_call_details.get('model'): + model = logging_obj.model_call_details.get('model') + + assert model == "" # Should be empty + + def test_edge_case_missing_model_call_details_attribute(self): + """Test fallback behavior when logging_obj doesn't have model_call_details attribute""" + # Case where logging_obj doesn't have the attribute at all + request_body = {"model": ""} # Empty model in request body + logging_obj = MagicMock() + # Remove the attribute to simulate it not existing + if hasattr(logging_obj, 'model_call_details'): + delattr(logging_obj, 'model_call_details') + + # Extract the logic directly from the function + model = request_body.get("model", "") + if not model and hasattr(logging_obj, 'model_call_details') and logging_obj.model_call_details.get('model'): + model = logging_obj.model_call_details.get('model') + + assert model == "" # Should remain empty since no fallback available + + # Case where model_call_details exists but get returns None + request_body = {"model": ""} + logging_obj = self._create_mock_logging_obj() # Empty dict + + model = request_body.get("model", "") + if not model and hasattr(logging_obj, 'model_call_details') and logging_obj.model_call_details.get('model'): + model = logging_obj.model_call_details.get('model') + + assert model == "" # Should remain empty \ No newline at end of file