From 7d09375d52032ae46a3b097a674d0a056e4e0f18 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 19 Aug 2025 22:21:00 -0700 Subject: [PATCH] fix: fix gpt-5-chat mappings --- .../llms/azure/chat/gpt_5_transformation.py | 3 +- .../llms/openai/chat/gpt_5_transformation.py | 8 +- ...odel_prices_and_context_window_backup.json | 4 +- model_prices_and_context_window.json | 4 +- tests/test_litellm/test_utils.py | 88 ++++++++++--------- 5 files changed, 60 insertions(+), 47 deletions(-) diff --git a/litellm/llms/azure/chat/gpt_5_transformation.py b/litellm/llms/azure/chat/gpt_5_transformation.py index b8c032f4e9d..d563a2889ca 100644 --- a/litellm/llms/azure/chat/gpt_5_transformation.py +++ b/litellm/llms/azure/chat/gpt_5_transformation.py @@ -10,6 +10,7 @@ from .gpt_transformation import AzureOpenAIConfig class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): """Azure specific handling for gpt-5 models.""" + GPT5_SERIES_ROUTE = "gpt5_series/" @classmethod @@ -23,7 +24,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): def get_supported_openai_params(self, model: str) -> List[str]: return OpenAIGPT5Config.get_supported_openai_params(self, model=model) - + def map_openai_params( self, non_default_params: dict, diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index a0b76abb20b..9a8bb74d447 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -15,14 +15,19 @@ class OpenAIGPT5Config(OpenAIGPTConfig): - Mapping ``max_tokens`` -> ``max_completion_tokens``. - Dropping unsupported ``temperature`` values when requested. """ + @classmethod def is_model_gpt_5_model(cls, model: str) -> bool: return "gpt-5" in model - + def get_supported_openai_params(self, model: str) -> list: + from litellm.utils import supports_tool_choice + base_gpt_series_params = super().get_supported_openai_params(model=model) gpt_5_only_params = ["reasoning_effort"] base_gpt_series_params.extend(gpt_5_only_params) + if not supports_tool_choice(model=model): + base_gpt_series_params.remove("tool_choice") return base_gpt_series_params def map_openai_params( @@ -61,4 +66,3 @@ class OpenAIGPT5Config(OpenAIGPTConfig): model=model, drop_params=drop_params, ) - diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index a48c26bc87c..acca6184d36 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -2483,7 +2483,7 @@ "supports_vision": true, "supports_prompt_caching": true, "supports_system_messages": true, - "supports_tool_choice": true, + "supports_tool_choice": false, "supports_native_streaming": true, "supports_reasoning": true, "source": "https://azure.microsoft.com/en-us/blog/gpt-5-in-azure-ai-foundry-the-future-of-ai-apps-and-agents-starts-here/" @@ -2516,7 +2516,7 @@ "supports_vision": true, "supports_prompt_caching": true, "supports_system_messages": true, - "supports_tool_choice": true, + "supports_tool_choice": false, "supports_native_streaming": true, "supports_reasoning": true }, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index a48c26bc87c..acca6184d36 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -2483,7 +2483,7 @@ "supports_vision": true, "supports_prompt_caching": true, "supports_system_messages": true, - "supports_tool_choice": true, + "supports_tool_choice": false, "supports_native_streaming": true, "supports_reasoning": true, "source": "https://azure.microsoft.com/en-us/blog/gpt-5-in-azure-ai-foundry-the-future-of-ai-apps-and-agents-starts-here/" @@ -2516,7 +2516,7 @@ "supports_vision": true, "supports_prompt_caching": true, "supports_system_messages": true, - "supports_tool_choice": true, + "supports_tool_choice": false, "supports_native_streaming": true, "supports_reasoning": true }, diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index cb53638d83e..1a675a391d0 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -376,23 +376,23 @@ def test_cohere_embedding_optional_params(): def validate_model_cost_values(model_data, exceptions=None): """ Validates that cost values in model data do not exceed 1. - + Args: model_data (dict): The model data dictionary exceptions (list, optional): List of model IDs that are allowed to have costs > 1 - + Returns: tuple: (is_valid, violations) where is_valid is a boolean and violations is a list of error messages """ if exceptions is None: exceptions = [] - + violations = [] - + # Define all cost-related fields to check cost_fields = [ "input_cost_per_token", - "output_cost_per_token", + "output_cost_per_token", "input_cost_per_character", "output_cost_per_character", "input_cost_per_image", @@ -431,22 +431,22 @@ def validate_model_cost_values(model_data, exceptions=None): "output_cost_per_reasoning_token", "citation_cost_per_token", ] - + # Also check nested cost fields nested_cost_fields = [ "search_context_cost_per_query", ] - + for model_id, model_info in model_data.items(): # Skip if this model is in exceptions if model_id in exceptions: continue - + # Check direct cost fields for field in cost_fields: if field in model_info and model_info[field] is not None: cost_value = model_info[field] - + # Convert string values to float if needed if isinstance(cost_value, str): try: @@ -454,12 +454,12 @@ def validate_model_cost_values(model_data, exceptions=None): except (ValueError, TypeError): # Skip if we can't convert to float continue - + if isinstance(cost_value, (int, float)) and cost_value > 1: violations.append( f"Model '{model_id}' has {field} = {cost_value} which exceeds 1" ) - + # Check nested cost fields for field in nested_cost_fields: if field in model_info and model_info[field] is not None: @@ -473,12 +473,12 @@ def validate_model_cost_values(model_data, exceptions=None): except (ValueError, TypeError): # Skip if we can't convert to float continue - + if isinstance(nested_value, (int, float)) and nested_value > 1: violations.append( f"Model '{model_id}' has {field}.{nested_field} = {nested_value} which exceeds 1" ) - + return len(violations) == 0, violations @@ -653,7 +653,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): # Validate schema validate(actual_json, INTENDED_SCHEMA) - + # Validate cost values # Define exceptions for models that are allowed to have costs > 1 # Add model IDs here if they legitimately have costs > 1 @@ -661,9 +661,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): # Add any model IDs that should be exempt from the cost validation # Example: "expensive-model-id", ] - + is_valid, violations = validate_model_cost_values(actual_json, exceptions) - + if not is_valid: error_message = "Cost validation failed:\n" + "\n".join(violations) error_message += "\n\nTo add exceptions, add the model ID to the 'exceptions' list in the test function." @@ -2330,25 +2330,31 @@ def test_block_key_hashing_logic(): ("", False, ""), # Empty string should not be hashed ("sk-", True, hash_token("sk-")), # Edge case: just "sk-" ] - + for input_key, should_be_hashed, expected_output in test_cases: # Simulate the logic from block_key() function if input_key.startswith("sk-"): hashed_token = hash_token(token=input_key) else: hashed_token = input_key - + assert hashed_token == expected_output, f"Failed for input: {input_key}" - + # Additional verification: if it should be hashed, verify it's actually a hash if should_be_hashed: # SHA-256 hashes are 64 characters long and contain only hex digits - assert len(hashed_token) == 64, f"Hash length should be 64, got {len(hashed_token)} for {input_key}" - assert all(c in '0123456789abcdef' for c in hashed_token), f"Hash should contain only hex digits for {input_key}" + assert ( + len(hashed_token) == 64 + ), f"Hash length should be 64, got {len(hashed_token)} for {input_key}" + assert all( + c in "0123456789abcdef" for c in hashed_token + ), f"Hash should contain only hex digits for {input_key}" else: # If not hashed, it should be the original string - assert hashed_token == input_key, f"Non-hashed key should remain unchanged: {input_key}" - + assert ( + hashed_token == input_key + ), f"Non-hashed key should remain unchanged: {input_key}" + print("✅ All block_key hashing logic tests passed!") @@ -2357,36 +2363,38 @@ def test_generate_gcp_iam_access_token(): Test the _generate_gcp_iam_access_token function with mocked GCP IAM client. """ from unittest.mock import Mock, patch - + service_account = "projects/-/serviceAccounts/test@project.iam.gserviceaccount.com" expected_token = "test-access-token-12345" - + # Mock the GCP IAM client and its response mock_response = Mock() mock_response.access_token = expected_token - + mock_client = Mock() mock_client.generate_access_token.return_value = mock_response - + # Mock the iam_credentials_v1 module mock_iam_credentials_v1 = Mock() mock_iam_credentials_v1.IAMCredentialsClient = Mock(return_value=mock_client) mock_iam_credentials_v1.GenerateAccessTokenRequest = Mock() - + # Test successful token generation by mocking sys.modules - with patch.dict('sys.modules', {'google.cloud.iam_credentials_v1': mock_iam_credentials_v1}): + with patch.dict( + "sys.modules", {"google.cloud.iam_credentials_v1": mock_iam_credentials_v1} + ): from litellm._redis import _generate_gcp_iam_access_token - + result = _generate_gcp_iam_access_token(service_account) - + assert result == expected_token mock_iam_credentials_v1.IAMCredentialsClient.assert_called_once() mock_client.generate_access_token.assert_called_once() - + # Verify the request was created with correct parameters mock_iam_credentials_v1.GenerateAccessTokenRequest.assert_called_once_with( name=service_account, - scope=['https://www.googleapis.com/auth/cloud-platform'] + scope=["https://www.googleapis.com/auth/cloud-platform"], ) @@ -2398,17 +2406,17 @@ def test_generate_gcp_iam_access_token_import_error(): from litellm._redis import _generate_gcp_iam_access_token # Mock the import to fail when the function tries to import google.cloud.iam_credentials_v1 - original_import = __builtins__['__import__'] - + original_import = __builtins__["__import__"] + def mock_import(name, *args, **kwargs): - if name == 'google.cloud.iam_credentials_v1': + if name == "google.cloud.iam_credentials_v1": raise ImportError("No module named 'google.cloud.iam_credentials_v1'") return original_import(name, *args, **kwargs) - - with patch('builtins.__import__', side_effect=mock_import): + + with patch("builtins.__import__", side_effect=mock_import): with pytest.raises(ImportError) as exc_info: _generate_gcp_iam_access_token("test-service-account") - + assert "google-cloud-iam is required" in str(exc_info.value) assert "pip install google-cloud-iam" in str(exc_info.value) @@ -2428,4 +2436,4 @@ def test_model_info_for_vertex_ai_deepseek_model(): assert model_info["input_cost_per_token"] is not None assert model_info["output_cost_per_token"] is not None - print("vertex deepseek model info", model_info) \ No newline at end of file + print("vertex deepseek model info", model_info)