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
This commit is contained in:
kothamah 2025-12-02 12:19:53 -05:00 • committed by GitHub
parent 81f4d863ca
commit 12530b375f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 392 additions and 0 deletions

View file

@ -859,6 +859,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
"deepseek_r1",
"qwen3",
"twelvelabs",
"openai"
]
BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[

View file

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

View file

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

View file

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

View file

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