fix: fix gpt-5-chat mappings

This commit is contained in:
Krrish Dholakia 2025-08-19 22:21:00 -07:00
parent 1832c09d6b
commit 7d09375d52
5 changed files with 60 additions and 47 deletions

View file

@ -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,

View file

@ -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,
)

View file

@ -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
},

View file

@ -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
},

View file

@ -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)
print("vertex deepseek model info", model_info)