mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(bedrock_mantle): route Claude chat completions to the native Messages endpoint
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
98c710c411
commit
9f05d84da5
6 changed files with 189 additions and 5 deletions
|
|
@ -3,6 +3,7 @@ from typing import Final, Literal
|
|||
import litellm
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
|
||||
from litellm.llms.bedrock_mantle.chat.claude_transformation import bedrock_mantle_chat_config
|
||||
from litellm.types.utils import LlmProviders, LlmProvidersSet
|
||||
|
||||
|
||||
|
|
@ -104,7 +105,7 @@ def get_supported_openai_params(
|
|||
elif custom_llm_provider == "groq":
|
||||
return litellm.GroqChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "bedrock_mantle":
|
||||
return litellm.BedrockMantleChatConfig().get_supported_openai_params(model=model)
|
||||
return bedrock_mantle_chat_config(model).get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "hosted_vllm":
|
||||
return litellm.HostedVLLMChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "vllm":
|
||||
|
|
|
|||
33
litellm/llms/bedrock_mantle/chat/claude_transformation.py
Normal file
33
litellm/llms/bedrock_mantle/chat/claude_transformation.py
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
from litellm.llms.base_llm.chat.transformation import BaseConfig
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.bedrock.chat.mantle.transformation import AmazonMantleConfig
|
||||
from litellm.llms.bedrock_mantle.chat.transformation import BedrockMantleChatConfig
|
||||
from litellm.llms.bedrock_mantle.common_utils import BedrockMantleAuthMixin, is_mantle_claude_model
|
||||
from litellm.llms.bedrock_mantle.messages.transformation import build_mantle_native_messages_url
|
||||
|
||||
|
||||
class BedrockMantleClaudeChatConfig(BedrockMantleAuthMixin, AmazonMantleConfig):
|
||||
def __init__(self, aws_signer: BaseAWSLLM | None = None) -> None:
|
||||
AmazonMantleConfig.__init__(self)
|
||||
self._aws_signer = aws_signer or self
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> str | None:
|
||||
return "bedrock_mantle"
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
return build_mantle_native_messages_url(api_base=api_base, litellm_params=litellm_params)
|
||||
|
||||
|
||||
def bedrock_mantle_chat_config(model: str) -> BaseConfig:
|
||||
if is_mantle_claude_model(model):
|
||||
return BedrockMantleClaudeChatConfig()
|
||||
return BedrockMantleChatConfig()
|
||||
|
|
@ -127,6 +127,10 @@ class BedrockMantleAuthMixin(SignsRequestsWithAWS):
|
|||
) from e
|
||||
|
||||
|
||||
def is_mantle_claude_model(model: str) -> bool:
|
||||
return "claude" in model.lower()
|
||||
|
||||
|
||||
def mantle_supports_responses(model: str | None, model_cost: dict) -> bool:
|
||||
"""Whether a Bedrock Mantle model can serve the native Responses API.
|
||||
|
||||
|
|
|
|||
|
|
@ -117,6 +117,7 @@ from litellm.llms.base_llm.base_model_iterator import (
|
|||
convert_model_response_to_streaming,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.llms.bedrock_mantle.chat.claude_transformation import bedrock_mantle_chat_config
|
||||
from litellm.llms.cohere.common_utils import CohereModelInfo
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler, http2_enabled
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
|
||||
|
|
@ -2188,7 +2189,7 @@ def _complete_bedrock_mantle(
|
|||
api_base = api_base or litellm.api_base or get_secret("BEDROCK_MANTLE_API_BASE")
|
||||
api_key = api_key or litellm.api_key or get_secret("BEDROCK_MANTLE_API_KEY")
|
||||
headers = headers or litellm.headers
|
||||
config: Final = litellm.BedrockMantleChatConfig.get_config()
|
||||
config: Final = bedrock_mantle_chat_config(model).get_config()
|
||||
for k, v in _provider_config_items(config):
|
||||
if k not in optional_params:
|
||||
optional_params[k] = v
|
||||
|
|
|
|||
|
|
@ -4857,7 +4857,7 @@ def get_optional_params(
|
|||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "bedrock_mantle":
|
||||
optional_params = litellm.BedrockMantleChatConfig().map_openai_params(
|
||||
optional_params = ProviderConfigManager._get_bedrock_mantle_config(model).map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
|
|
@ -8425,8 +8425,8 @@ class ProviderConfigManager:
|
|||
LlmProviders.TENCENT: (lambda: litellm.TencentChatConfig(), False),
|
||||
LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False),
|
||||
LlmProviders.BEDROCK_MANTLE: (
|
||||
lambda: litellm.BedrockMantleChatConfig(),
|
||||
False,
|
||||
lambda model: ProviderConfigManager._get_bedrock_mantle_config(model),
|
||||
True,
|
||||
),
|
||||
LlmProviders.A2A: (lambda: litellm.A2AConfig(), False),
|
||||
LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False),
|
||||
|
|
@ -8615,6 +8615,12 @@ class ProviderConfigManager:
|
|||
|
||||
return get_bedrock_chat_config(model=model)
|
||||
|
||||
@staticmethod
|
||||
def _get_bedrock_mantle_config(model: str) -> BaseConfig:
|
||||
from litellm.llms.bedrock_mantle.chat.claude_transformation import bedrock_mantle_chat_config
|
||||
|
||||
return bedrock_mantle_chat_config(model)
|
||||
|
||||
@staticmethod
|
||||
def _get_cohere_config(model: str) -> BaseConfig:
|
||||
"""Get Cohere config based on route."""
|
||||
|
|
|
|||
|
|
@ -818,6 +818,145 @@ class TestBedrockMantleProviderResolution:
|
|||
)
|
||||
|
||||
|
||||
def _anthropic_message(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "msg_test",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "anthropic.claude-opus-5-5",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
},
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_event_stream(request: httpx.Request) -> httpx.Response:
|
||||
events = (
|
||||
(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_test",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "anthropic.claude-opus-5-5",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 0},
|
||||
},
|
||||
},
|
||||
),
|
||||
(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
),
|
||||
(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "streamed"}},
|
||||
),
|
||||
(
|
||||
"message_delta",
|
||||
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 3}},
|
||||
),
|
||||
("message_stop", {"type": "message_stop"}),
|
||||
)
|
||||
body = "".join(f"event: {name}\ndata: {json.dumps(data)}\n\n" for name, data in events)
|
||||
return httpx.Response(
|
||||
status_code=200, content=body.encode(), headers={"content-type": "text/event-stream"}, request=request
|
||||
)
|
||||
|
||||
|
||||
class TestBedrockMantleClaudeChatRoute:
|
||||
def test_claude_completion_uses_native_messages_endpoint(self, monkeypatch, local_cost_map):
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "mantle-key")
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
handler = Mock(side_effect=_anthropic_message)
|
||||
|
||||
response = litellm.completion(
|
||||
model="bedrock_mantle/anthropic.claude-opus-5-5",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
max_tokens=64,
|
||||
aws_region_name="us-east-2",
|
||||
client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler))),
|
||||
)
|
||||
|
||||
sent = handler.call_args.args[0]
|
||||
opus = litellm.model_cost["anthropic.claude-opus-5-5"]
|
||||
assert str(sent.url) == "https://bedrock-mantle.us-east-2.api.aws/anthropic/v1/messages"
|
||||
assert sent.headers["Authorization"] == "Bearer mantle-key"
|
||||
assert json.loads(sent.content) == {
|
||||
"model": "anthropic.claude-opus-5-5",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hello"}]}],
|
||||
"max_tokens": 64,
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
}
|
||||
assert response.choices[0].message.content == "ok"
|
||||
assert response._hidden_params["response_cost"] == pytest.approx(
|
||||
10 * opus["input_cost_per_token"] + 5 * opus["output_cost_per_token"]
|
||||
)
|
||||
|
||||
def test_claude_streaming_completion_uses_native_messages_endpoint(self, monkeypatch, local_cost_map):
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "mantle-key")
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
handler = Mock(side_effect=_anthropic_event_stream)
|
||||
|
||||
stream = litellm.completion(
|
||||
model="bedrock_mantle/anthropic.claude-opus-5-5",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
max_tokens=64,
|
||||
stream=True,
|
||||
aws_region_name="us-east-2",
|
||||
client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler))),
|
||||
)
|
||||
text = "".join(chunk.choices[0].delta.content or "" for chunk in stream)
|
||||
|
||||
sent = handler.call_args.args[0]
|
||||
assert str(sent.url) == "https://bedrock-mantle.us-east-2.api.aws/anthropic/v1/messages"
|
||||
assert json.loads(sent.content)["stream"] is True
|
||||
assert text == "streamed"
|
||||
|
||||
def test_non_claude_completion_stays_on_chat_completions(self, monkeypatch, local_cost_map):
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "mantle-key")
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
handler = Mock(
|
||||
side_effect=lambda request: httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "chatcmpl-test",
|
||||
"object": "chat.completion",
|
||||
"created": 1733529600,
|
||||
"model": "openai.gpt-oss-120b",
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
},
|
||||
request=request,
|
||||
)
|
||||
)
|
||||
|
||||
response = litellm.completion(
|
||||
model="bedrock_mantle/openai.gpt-oss-120b",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
aws_region_name="us-east-2",
|
||||
client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler))),
|
||||
)
|
||||
|
||||
sent = handler.call_args.args[0]
|
||||
assert str(sent.url) == "https://bedrock-mantle.us-east-2.api.aws/v1/chat/completions"
|
||||
assert response.choices[0].message.content == "ok"
|
||||
|
||||
|
||||
class TestBedrockMantlePricing:
|
||||
"""Tests that verify Bedrock Mantle uses correct AWS Bedrock pricing, not OpenAI pricing."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue