From 12530b375fe969ae95972c1b6ef4cf5d6802d000 Mon Sep 17 00:00:00 2001 From: kothamah <104782493+kothamah@users.noreply.github.com> Date: Tue, 2 Dec 2025 12:19:53 -0500 Subject: [PATCH] Litellm bedrock OpenAI model support (#17368) * Update constants.py added constants * Update base_aws_llm.py added steps * Update invoke_handler.py added openai support * Update base_invoke_transformation.py added * Update test_bedrock_completion.py added --- litellm/constants.py | 1 + litellm/llms/bedrock/base_aws_llm.py | 4 + litellm/llms/bedrock/chat/invoke_handler.py | 45 +++ .../base_invoke_transformation.py | 9 + .../test_bedrock_completion.py | 333 ++++++++++++++++++ 5 files changed, 392 insertions(+) diff --git a/litellm/constants.py b/litellm/constants.py index e3de7368c8a..1d42ef9a910 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -859,6 +859,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[ "deepseek_r1", "qwen3", "twelvelabs", + "openai" ] BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[ diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index ed658c793af..816b93edd20 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -353,6 +353,10 @@ class BaseAWSLLM: model_id = BaseAWSLLM._get_model_id_from_model_with_spec( model_id, spec="deepseek_r1" ) + elif provider == "openai" and "openai/" in model_id: + model_id = BaseAWSLLM._get_model_id_from_model_with_spec( + model_id, spec="openai" + ) return model_id @staticmethod diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index b35e86cabd2..7a960fd45d2 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -73,6 +73,9 @@ bedrock_tool_name_mappings: InMemoryCache = InMemoryCache( max_size_in_memory=50, default_ttl=600 ) from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig +from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, +) converse_config = AmazonConverseConfig() @@ -401,6 +404,10 @@ class BedrockLLM(BaseAWSLLM): prompt = prompt_factory( model=model, messages=messages, custom_llm_provider="bedrock" ) + elif provider == "openai": + # OpenAI uses messages directly, no prompt conversion needed + # Return empty prompt as it won't be used + prompt = "" elif provider == "cohere": prompt, chat_history = cohere_message_pt(messages=messages) else: @@ -578,6 +585,30 @@ class BedrockLLM(BaseAWSLLM): ) elif provider == "meta" or provider == "llama": outputText = completion_response["generation"] + elif provider == "openai": + # OpenAI imported models use OpenAI Chat Completions format + if "choices" in completion_response and len(completion_response["choices"]) > 0: + choice = completion_response["choices"][0] + if "message" in choice: + outputText = choice["message"].get("content") + elif "text" in choice: # fallback for completion format + outputText = choice["text"] + + # Set finish reason + if "finish_reason" in choice: + model_response.choices[0].finish_reason = map_finish_reason( + choice["finish_reason"] + ) + + # Set usage if available + if "usage" in completion_response: + usage = completion_response["usage"] + _usage = litellm.Usage( + prompt_tokens=usage.get("prompt_tokens", 0), + completion_tokens=usage.get("completion_tokens", 0), + total_tokens=usage.get("total_tokens", 0), + ) + setattr(model_response, "usage", _usage) elif provider == "mistral": outputText = completion_response["outputs"][0]["text"] model_response.choices[0].finish_reason = completion_response[ @@ -895,6 +926,20 @@ class BedrockLLM(BaseAWSLLM): ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in inference_params[k] = v data = json.dumps({"prompt": prompt, **inference_params}) + elif provider == "openai": + ## OpenAI imported models use OpenAI Chat Completions format (messages-based) + # Use AmazonBedrockOpenAIConfig for proper OpenAI transformation + openai_config = AmazonBedrockOpenAIConfig() + supported_params = openai_config.get_supported_openai_params(model=model) + + # Filter to only supported OpenAI params + filtered_params = { + k: v for k, v in inference_params.items() + if k in supported_params + } + + # OpenAI uses messages format, not prompt + data = json.dumps({"messages": messages, **filtered_params}) else: ## LOGGING logging_obj.pre_call( diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 6c389ff3b7d..bcb4cae1c8b 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -258,6 +258,15 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): litellm_params=litellm_params, headers=headers, ) + elif provider == "openai": + # OpenAI imported models use OpenAI Chat Completions format + return litellm.AmazonBedrockOpenAIConfig().transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) else: raise BedrockError( status_code=404, diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index f43e939c681..bd08d4444f6 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3531,3 +3531,336 @@ def test_bedrock_openai_imported_model(): # Check max_tokens and temperature assert request_body["max_tokens"] == 300 assert request_body["temperature"] == 0.5 + +def test_bedrock_openai_provider_detection(): + """ + Test that the OpenAI provider is correctly detected from model strings. + """ + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + # Test various OpenAI model formats + test_cases = [ + "openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/abc123", + "bedrock/openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/xyz789", + ] + + for model in test_cases: + provider = BaseAWSLLM.get_bedrock_invoke_provider(model) + assert provider == "openai", f"Failed for model: {model}, got provider: {provider}" + print(f"✓ Provider detection works for: {model}") + + +def test_bedrock_openai_model_id_extraction(): + """ + Test that the model ID (ARN) is correctly extracted and encoded for OpenAI models. + """ + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + model = "openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-model-123" + provider = BaseAWSLLM.get_bedrock_invoke_provider(model) + + model_id = BaseAWSLLM.get_bedrock_model_id( + model=model, + provider=provider, + optional_params={} + ) + + # The ARN should be double URL encoded + assert "arn" in model_id + assert "imported-model" in model_id + print(f"✓ Model ID extracted and encoded: {model_id}") + + +def test_bedrock_openai_convert_messages_to_prompt(): + """ + Test that convert_messages_to_prompt returns empty string for OpenAI models. + """ + from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM + + bedrock_llm = BedrockLLM() + messages = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "Hello"} + ] + + prompt, chat_history = bedrock_llm.convert_messages_to_prompt( + model="test-model", + messages=messages, + provider="openai", + custom_prompt_dict={} + ) + + # OpenAI models use messages directly, no prompt conversion + assert prompt == "" + assert chat_history is None + print("✓ convert_messages_to_prompt returns empty for OpenAI") + + +def test_bedrock_openai_response_parsing(): + """ + Test that OpenAI responses are correctly parsed. + """ + from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM + from litellm import ModelResponse + from unittest.mock import Mock + import json + + bedrock_llm = BedrockLLM() + + # Mock OpenAI-style response + openai_response = { + "choices": [ + { + "message": { + "content": "The capital of France is Paris.", + "role": "assistant" + }, + "finish_reason": "stop", + "index": 0 + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 8, + "total_tokens": 18 + } + } + + mock_response = Mock() + mock_response.json.return_value = openai_response + mock_response.text = json.dumps(openai_response) + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ModelResponse() + mock_logging = Mock() + + result = bedrock_llm.process_response( + model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", + response=mock_response, + model_response=model_response, + stream=False, + logging_obj=mock_logging, + optional_params={}, + api_key="", + data={}, + messages=[{"role": "user", "content": "What is the capital of France?"}], + print_verbose=lambda x: None, + encoding=None + ) + + # Verify response content + assert result.choices[0].message.content == "The capital of France is Paris." + assert result.choices[0].finish_reason == "stop" + + # Verify usage + assert result.usage.prompt_tokens == 10 + assert result.usage.completion_tokens == 8 + assert result.usage.total_tokens == 18 + + print("✓ OpenAI response parsing works correctly") + + +def test_bedrock_openai_request_transformation(): + """ + Test that the request is correctly transformed for OpenAI models. + """ + from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import AmazonInvokeConfig + + config = AmazonInvokeConfig() + + model = "openai/arn:aws:bedrock:us-east-1:123:imported-model/test" + messages = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "Hello"} + ] + + optional_params = { + "max_tokens": 100, + "temperature": 0.7, + "top_p": 0.9, + "stream": False + } + + litellm_params = {} + headers = {} + + with patch.object(config, 'get_bedrock_invoke_provider', return_value="openai"): + result = config.transform_request( + model=model, + messages=messages, + optional_params=optional_params.copy(), + litellm_params=litellm_params, + headers=headers + ) + + # Verify the request uses messages format (not prompt) + assert "messages" in result + assert len(result["messages"]) == 2 + assert result["messages"][0]["role"] == "system" + assert result["messages"][1]["role"] == "user" + + # Verify parameters are included + assert "max_tokens" in result + assert "temperature" in result + + print("✓ Request transformation works correctly") + + +def test_bedrock_openai_parameter_filtering(): + """ + Test that only supported OpenAI parameters are included in the request. + """ + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import AmazonBedrockOpenAIConfig + + config = AmazonBedrockOpenAIConfig() + model = "test-model" + + supported_params = config.get_supported_openai_params(model=model) + + # Verify common OpenAI parameters are supported + assert "max_tokens" in supported_params + assert "temperature" in supported_params + assert "top_p" in supported_params + assert "stream" in supported_params + assert "stop" in supported_params + + print(f"✓ Parameter filtering supports: {len(supported_params)} parameters") + print(f" Supported params: {supported_params}") + + +def test_bedrock_openai_route_detection(): + """ + Test that the OpenAI route is correctly detected. + """ + from litellm.llms.bedrock.common_utils import BedrockModelInfo + + test_cases = [ + ("openai/arn:aws:bedrock:us-east-1:123:imported-model/test", "openai"), + ("bedrock/openai/arn:aws:bedrock:us-east-1:123:imported-model/test", "openai"), + ] + + for model, expected_route in test_cases: + route = BedrockModelInfo.get_bedrock_route(model) + assert route == expected_route, f"Failed for model: {model}, got route: {route}" + print(f"✓ Route detection works for: {model} -> {route}") + + +def test_bedrock_openai_explicit_route_check(): + """ + Test the explicit OpenAI route checker helper method. + """ + from litellm.llms.bedrock.common_utils import BedrockModelInfo + + # Test with openai/ prefix + assert BedrockModelInfo._explicit_openai_route("openai/arn:aws:bedrock:us-east-1:123:imported-model/test") is True + assert BedrockModelInfo._explicit_openai_route("bedrock/openai/arn:aws:bedrock:us-east-1:123:imported-model/test") is True + + # Test without openai/ prefix + assert BedrockModelInfo._explicit_openai_route("anthropic.claude-3-sonnet") is False + assert BedrockModelInfo._explicit_openai_route("arn:aws:bedrock:us-east-1:123:imported-model/test") is False + + print("✓ Explicit route check works correctly") + + +def test_bedrock_openai_config_initialization(): + """ + Test that AmazonBedrockOpenAIConfig can be properly initialized. + """ + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import AmazonBedrockOpenAIConfig + + config = AmazonBedrockOpenAIConfig() + + # Verify it has the necessary methods + assert hasattr(config, 'get_supported_openai_params') + assert hasattr(config, 'transform_request') + assert hasattr(config, 'transform_response') + assert hasattr(config, 'map_openai_params') + + print("✓ AmazonBedrockOpenAIConfig initializes correctly") + + +def test_bedrock_openai_multiple_message_types(): + """ + Test that various message content types are handled correctly. + """ + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + + # Test with mixed content types + messages = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "Simple text message"}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Complex message with text"}, + {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,iVBORw0KGg"}} + ] + } + ] + + with patch.object(client, "post") as mock_post: + try: + response = completion( + model="bedrock/openai/arn:aws:bedrock:us-east-1:123:imported-model/test", + messages=messages, + max_tokens=50, + client=client, + ) + except Exception as e: + pass + + # Verify the request was made + if mock_post.called: + request_body = json.loads(mock_post.call_args.kwargs["data"]) + + # Verify messages are preserved + assert "messages" in request_body + assert len(request_body["messages"]) == 3 + + # Verify mixed content is handled + assert isinstance(request_body["messages"][2]["content"], list) + + print("✓ Multiple message types handled correctly") + + +def test_bedrock_openai_error_handling(): + """ + Test that errors from OpenAI models are properly handled. + """ + from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM + from litellm import ModelResponse + from litellm.llms.bedrock.common_utils import BedrockError + from unittest.mock import Mock + import json + + bedrock_llm = BedrockLLM() + + # Mock error response + mock_response = Mock() + mock_response.json.side_effect = Exception("Invalid JSON") + mock_response.text = "Invalid response" + mock_response.status_code = 422 + + model_response = ModelResponse() + mock_logging = Mock() + + with pytest.raises(BedrockError) as exc_info: + bedrock_llm.process_response( + model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", + response=mock_response, + model_response=model_response, + stream=False, + logging_obj=mock_logging, + optional_params={}, + api_key="", + data={}, + messages=[], + print_verbose=lambda x: None, + encoding=None + ) + + assert exc_info.value.status_code == 422 + print("✓ Error handling works correctly")