fix(bedrock_mantle): route Claude chat completions to the native Messages endpoint (#43646)

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

* fix(bedrock_mantle): price Claude chat on the Mantle row and route region-prefixed ids

A Mantle request always carries a region, so a Claude id with no bedrock_mantle/<region>/ row fell through the model-info lookup to the bare Bedrock row, which the bedrock provider family also matches, and billed about 10 percent under the Mantle price. The lookup now tries the provider's region-free row before the bare model. The Claude route test asserts the Mantle row, a region-prefixed Claude id is covered end to end, and the provider config map references the Mantle config directly.

* test(bedrock_mantle): cover supported params for Claude and open-weight Mantle ids

* test(bedrock_mantle): audit the Claude chat bridge on the integration rig

Adds the deterministic cells from the /audit of the Mantle Claude chat
bridge: wire-level translation on every chat, responses, and messages
route, SigV4 and bearer auth, region prefixes, api_base suffixes,
unsupported params with and without drop_params, malformed model ids,
upstream errors, the response-cache hit, the health check, pricing from
the Mantle row for Claude and non-Claude ids with a bare Bedrock twin,
and chaos cells for a mixed burst, an upstream outage, slow streams,
and a worker kill on an owned two-worker proxy

---------

Co-authored-by: jesus <jesus@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 15:24:41 -07:00 • committed by GitHub
parent 0ea166c160
commit 01b4ffe16b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 1853 additions and 77 deletions

View file

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

View 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[str, object],
litellm_params: dict[str, object],
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()

View file

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

View file

@ -121,6 +121,7 @@ from litellm.llms.bedrock.common_utils import (
bedrock_route_for_request,
without_bedrock_route_prefix,
)
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
@ -2192,7 +2193,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

View file

@ -4897,7 +4897,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,
@ -5669,7 +5669,7 @@ def _get_model_cost_key(potential_key: str) -> str | None:
return None
def _get_model_info_from_model_cost(key: str) -> dict:
def _get_model_info_from_model_cost(key: str) -> dict[str, Any]:
return litellm.model_cost[key]
@ -5724,12 +5724,26 @@ from typing_extensions import ReadOnly, TypedDict
class PotentialModelNamesAndCustomLLMProvider(TypedDict):
split_model: str
combined_model_name: str
region_free_combined_model_name: ReadOnly[str]
stripped_model_name: str
combined_stripped_model_name: str
provider_prefixed_model_name: ReadOnly[str]
custom_llm_provider: str
def _first_registered_match(
candidates: Sequence[str], custom_llm_provider: str | None
) -> tuple[str | None, dict[str, Any] | None]:
registered_keys: Final = (key for key in map(_get_model_cost_key, candidates) if key is not None)
entries: Final = ((key, _get_model_info_from_model_cost(key=key)) for key in registered_keys)
matches: Final = (
(key, info)
for key, info in entries
if _check_provider_match(model_info=info, custom_llm_provider=custom_llm_provider)
)
return next(matches, (None, None))
def _get_model_info_from_generalization(
model: str,
potential_model_names: PotentialModelNamesAndCustomLLMProvider,
@ -5751,6 +5765,7 @@ def _get_model_info_from_generalization(
candidates: Final = (
potential_model_names["combined_model_name"],
model,
potential_model_names["region_free_combined_model_name"],
potential_model_names["split_model"],
potential_model_names["combined_stripped_model_name"],
potential_model_names["stripped_model_name"],
@ -5828,6 +5843,11 @@ def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> P
return PotentialModelNamesAndCustomLLMProvider(
split_model=region_free_split_model,
combined_model_name=combined_model_name,
region_free_combined_model_name=(
f"bedrock_mantle/{region_free_split_model}"
if custom_llm_provider == "bedrock_mantle"
else combined_model_name
),
stripped_model_name=stripped_model_name,
combined_stripped_model_name=region_free_combined_stripped_model_name,
provider_prefixed_model_name=provider_cost_key or provider_prefixed_model_name,
@ -6007,78 +6027,29 @@ def _get_model_info_helper(
Check if: (in order of specificity)
1. 'custom_llm_provider/model' in litellm.model_cost. Checks "groq/llama3-8b-8192" if model="llama3-8b-8192" and custom_llm_provider="groq"
2. 'model' in litellm.model_cost. Checks "gemini-1.5-pro-002" in litellm.model_cost if model="gemini-1.5-pro-002" and custom_llm_provider=None
3. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8"
4. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given.
5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given.
6. 'provider_prefixed_model_name' in litellm.model_cost, for providers whose own model ids repeat the
3. 'region_free_combined_model_name' in litellm.model_cost. Checks "bedrock_mantle/anthropic.claude-opus-5-5" if
model="bedrock_mantle/us-east-1/anthropic.claude-opus-5-5", before 4 reaches the bare Bedrock row. Same as 1 for every other provider.
4. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8"
5. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given.
6. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given.
7. 'provider_prefixed_model_name' in litellm.model_cost, for providers whose own model ids repeat the
litellm provider name. Checks "perplexity/perplexity/glm-5.2" if model="perplexity/glm-5.2" and
custom_llm_provider="perplexity", where 1-5 all read the leading "perplexity/" as the litellm prefix
and strip it. Tried last so no model that already resolves through 1-5 can change.
custom_llm_provider="perplexity", where 1-6 all read the leading "perplexity/" as the litellm prefix
and strip it. Tried last so no model that already resolves through 1-6 can change.
"""
_model_info: dict[str, Any] | None = None
key: str | None = None
# Use case-insensitive lookup for all model name checks
_matched_key = _get_model_cost_key(combined_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(split_model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(combined_stripped_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(stripped_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(provider_prefixed_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
lookup_order: Final = (
combined_model_name,
model,
potential_model_names["region_free_combined_model_name"],
split_model,
combined_stripped_model_name,
stripped_model_name,
provider_prefixed_model_name,
)
lookup: Final = _first_registered_match(lookup_order, model_cost_custom_llm_provider)
key: str | None = lookup[0]
_model_info: dict[str, Any] | None = lookup[1]
if _model_info is not None and key is not None and _model_info.get("mode", "chat") in _BACKFILL_MODES:
fill_missing: Final = match_fill_missing_generalizations(key, _model_info.get("litellm_provider", ""))
@ -8478,10 +8449,7 @@ class ProviderConfigManager:
LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False),
LlmProviders.TENCENT: (lambda: litellm.TencentChatConfig(), False),
LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False),
LlmProviders.BEDROCK_MANTLE: (
lambda: litellm.BedrockMantleChatConfig(),
False,
),
LlmProviders.BEDROCK_MANTLE: (ProviderConfigManager._get_bedrock_mantle_config, True),
LlmProviders.A2A: (lambda: litellm.A2AConfig(), False),
LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False),
LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False),
@ -8669,6 +8637,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."""

View file

@ -0,0 +1,340 @@
import asyncio
import json
import re
import signal
import threading
import uuid
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final, Literal
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
import yaml
from integration._support.client import Gateway, eventually, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "anthropic.claude-haiku-4-5"
_API_KEY: Final = "synthetic-mantle-bearer"
_CONFIG_MODEL: Final = "bedrock-mantle-claude-chat-chaos"
_MESSAGES_PATH: Final = "/anthropic/v1/messages"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})")
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
_ROWS_BY_CALL: Final = (
"SELECT litellm_call_id, status FROM \"LiteLLM_SpendLogs\" WHERE litellm_call_id = ANY(string_to_array(%s, ','))"
)
Endpoint = Literal["chat", "messages", "responses"]
@dataclass(frozen=True, slots=True)
class _Call:
endpoint: Endpoint
stream: bool
marker: str
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
text: str
call_id: str
def _answer(marker: str) -> str:
return f"answer marker-{marker}"
def _path(endpoint: Endpoint) -> str:
match endpoint:
case "chat":
return "/v1/chat/completions"
case "messages":
return "/v1/messages"
case "responses":
return "/v1/responses"
def _body(model: str, call: _Call) -> dict[str, JsonValue]:
question: Final = f"Question marker-{call.marker}"
common: Final[dict[str, JsonValue]] = {"model": model, "stream": call.stream}
match call.endpoint:
case "chat":
return {**common, "messages": [{"role": "user", "content": question}]}
case "messages":
return {**common, "max_tokens": 64, "messages": [{"role": "user", "content": question}]}
case "responses":
return {**common, "input": question}
def _mantle_reply(marker: str, stream: bool, *, abort: bool = False, pause: float = 0) -> Reply:
message: Final = {
"id": f"msg_bdrk_{marker}",
"type": "message",
"role": "assistant",
"model": _BACKEND,
"stop_sequence": None,
}
if not stream:
payload: Final = json.dumps(
{
**message,
"content": [{"type": "text", "text": _answer(marker)}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 23, "output_tokens": 7},
}
).encode()
return Reply(chunks=(payload,), abort_after=0) if abort else Reply(body=payload)
opening: Final = {**message, "content": [], "stop_reason": None, "usage": {"input_tokens": 23, "output_tokens": 1}}
events: Final = (
("message_start", {"message": opening}),
("content_block_start", {"index": 0, "content_block": {"type": "text", "text": ""}}),
("content_block_delta", {"index": 0, "delta": {"type": "text_delta", "text": "answer "}}),
("content_block_delta", {"index": 0, "delta": {"type": "text_delta", "text": f"marker-{marker}"}}),
("content_block_stop", {"index": 0}),
("message_delta", {"delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 7}}),
("message_stop", {}),
)
return Reply(
content_type="text/event-stream",
chunks=tuple(
f"event: {kind}\ndata: {json.dumps({'type': kind, **payload})}\n\n".encode() for kind, payload in events
),
abort_after=0 if abort else None,
pause_between_chunks=pause,
)
def _marker_of(request: Request) -> str:
found: Final = _MARKER.search(request.body.decode())
assert found is not None, request.body
return found.group(1)
def _native_peer(aborted: frozenset[str] = frozenset(), pause: float = 0) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", _MESSAGES_PATH), request.target
assert request.headers["authorization"] == f"Bearer {_API_KEY}", sorted(request.headers)
marker: Final = _marker_of(request)
stream: Final = _JSON_OBJECT.validate_json(request.body).get("stream") is True
return _mantle_reply(marker, stream, abort=marker in aborted, pause=pause)
return respond
def _assert_native_requests_for(received: tuple[Request, ...], calls: tuple[_Call, ...]) -> None:
assert {(request.method, request.target) for request in received} == {("POST", _MESSAGES_PATH)}
assert sorted(_marker_of(request) for request in received) == sorted(call.marker for call in calls)
def _spend_status_by_call(served: tuple[_Served, ...]) -> dict[str, JsonValue]:
wanted: Final = sorted(item.call_id for item in served)
rows: Final = eventually(
lambda: read_rows(_ROWS_BY_CALL, (",".join(wanted),)), lambda found: len(found) >= len(wanted), seconds=70
)
assert sorted(string_value(row["litellm_call_id"]) for row in rows) == wanted, rows
return {string_value(row["litellm_call_id"]): row["status"] for row in rows}
async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served:
async with client.stream(
"POST",
_path(call.endpoint),
json=_body(model, call),
headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"},
) as response:
raw: Final = await response.aread()
return _Served(
call=call,
status=response.status_code,
text=raw.decode(),
call_id=response.headers.get("x-litellm-call-id", ""),
)
async def _burst(
base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client:
results: Final = await asyncio.gather(
*(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Served))
def _calls(count: int, endpoints: tuple[Endpoint, ...], stream: Callable[[int], bool]) -> tuple[_Call, ...]:
return tuple(
_Call(endpoint=endpoints[index % len(endpoints)], stream=stream(index), marker=uuid.uuid4().hex)
for index in range(count)
)
def _assert_answered_with_its_own_marker(served: _Served) -> None:
assert served.status == 200, served.text
assert set(_MARKER.findall(served.text)) == {served.call.marker}, served.text
async def test_concurrent_burst_across_endpoints_is_answered_from_the_native_route_and_logged_once_per_call(
gateway: Gateway,
) -> None:
calls: Final = _calls(30, ("chat", "messages", "responses"), lambda index: index % 2 == 0)
with wire_server(_native_peer()) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"bedrock_mantle/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
assert len(served) == 30
for item in served:
_assert_answered_with_its_own_marker(item)
_assert_native_requests_for(wire.drain(), calls)
assert _spend_status_by_call(served) == {item.call_id: "success" for item in served}
def _unhealthy_count(gateway: Gateway, model: str) -> JsonValue:
response: Final = gateway.request("GET", "/health", params={"model": model})
return _JSON_OBJECT.validate_json(response.content).get("unhealthy_count")
async def test_upstream_aborts_then_an_outage_fail_each_call_once_and_the_native_route_recovers(
gateway: Gateway,
) -> None:
calls: Final = _calls(21, ("chat", "messages", "responses"), lambda index: index % 2 == 0)
aborted: Final = frozenset(call.marker for call in calls if call.endpoint == "chat")
during_outage: Final = (
_Call(endpoint="chat", stream=True, marker=uuid.uuid4().hex),
_Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex),
_Call(endpoint="messages", stream=False, marker=uuid.uuid4().hex),
_Call(endpoint="responses", stream=False, marker=uuid.uuid4().hex),
)
after_restart: Final = _calls(3, ("chat", "messages", "responses"), lambda index: index == 0)
proxy: Final = str(gateway.client.base_url)
with gateway.scenario() as scenario:
with wire_server(_native_peer(aborted)) as wire:
upstream: Final = wire.url
model: Final = scenario.model(model=f"bedrock_mantle/{_BACKEND}", api_base=upstream, api_key=_API_KEY)
burst: Final = await _burst(proxy, gateway.key, model, calls)
assert len(burst) == 21
for item in burst:
if item.call.marker in aborted:
assert item.status == (500 if item.call.stream else 503), item.text
assert "Response payload is not completed" in item.text and "answer marker-" not in item.text
else:
_assert_answered_with_its_own_marker(item)
_assert_native_requests_for(wire.drain(), calls)
refused: Final = await _burst(proxy, gateway.key, model, during_outage)
assert len(refused) == 4
for item in refused:
assert item.status == 503, item.text
assert "Cannot connect to host" in item.text, item.text
assert await asyncio.to_thread(
eventually, lambda: _unhealthy_count(gateway, model), lambda count: count == 1, 30
)
with wire_server(_native_peer(), port=urlsplit(upstream).port or 0) as restarted:
recovered: Final = await _burst(proxy, gateway.key, model, after_restart)
assert len(recovered) == 3
for item in recovered:
_assert_answered_with_its_own_marker(item)
_assert_native_requests_for(restarted.drain(), after_restart)
failed: Final = frozenset(item.call_id for item in (*burst, *refused) if item.status != 200)
assert len(failed) == len(aborted) + 4
assert _spend_status_by_call((*burst, *refused, *recovered)) == {
item.call_id: "failure" if item.call_id in failed else "success" for item in (*burst, *refused, *recovered)
}
async def test_slow_native_streams_reach_every_caller_whole_and_are_logged_once(gateway: Gateway) -> None:
calls: Final = _calls(10, ("chat", "responses"), lambda _: True)
with wire_server(_native_peer(pause=0.3)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"bedrock_mantle/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
assert len(served) == 10
for item in served:
_assert_answered_with_its_own_marker(item)
assert item.text.rstrip().endswith("data: [DONE]"), item.text
_assert_native_requests_for(wire.drain(), calls)
assert _spend_status_by_call(served) == {item.call_id: "success" for item in served}
def _chaos_config(wire: Wire, tmp_path: Path) -> Path:
config: Final = {
**yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()),
"model_list": [
{
"model_name": _CONFIG_MODEL,
"litellm_params": {"model": f"bedrock_mantle/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY},
}
],
}
target: Final = tmp_path / "bedrock-mantle-claude-chat-chaos.yaml"
target.write_text(yaml.safe_dump(config))
return target
def _live_workers(log: Path) -> tuple[int, ...]:
started: Final = (int(pid) for pid in _STARTED_WORKER.findall(log.read_text()))
return tuple(pid for pid in started if psutil.pid_exists(pid))
def _open_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
@pytest.mark.timeout(300)
async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_the_native_route(
gateway: Gateway, tmp_path: Path
) -> None:
calls: Final = _calls(20, ("chat",), lambda _: False)
release: Final = threading.Event()
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
answer: Final = _native_peer()
def held(request: Request) -> Reply:
held_markers.put(_marker_of(request))
assert release.wait(timeout=60), "The burst was never released"
return answer(request)
with wire_server(held) as wire:
config: Final = _chaos_config(wire, tmp_path)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
candidate: Final = owned.gateway
workers: Final = eventually(lambda: _live_workers(owned.log), lambda pids: len(pids) == 2, seconds=30)
burst: Final = asyncio.create_task(
_burst(
str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, calls, tolerate_transport_errors=True
)
)
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60)
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers})
assert sum(held_by.values()) == 20, held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
served: Final = await burst
assert held_by[survivor_pid] >= 10, held_by
assert len(served) == held_by[survivor_pid], (held_by, len(served))
for item in served:
_assert_answered_with_its_own_marker(item)
follow_up: Final = _Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex)
(answered,) = await _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, (follow_up,))
_assert_answered_with_its_own_marker(answered)
_assert_native_requests_for(wire.drain(), (*calls, follow_up))
assert _spend_status_by_call((*served, answered)) == {
item.call_id: "success" for item in (*served, answered)
}

File diff suppressed because it is too large Load diff

View file

@ -7,6 +7,8 @@ API docs: https://docs.aws.amazon.com/bedrock/latest/userguide/bedrock-mantle.ht
import json
import asyncio
from collections.abc import Mapping
from typing import Final
from unittest.mock import Mock, patch
@ -16,6 +18,7 @@ from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
import litellm
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.bedrock_mantle.chat.transformation import BedrockMantleChatConfig
from litellm.llms.bedrock.base_aws_llm import sign_request_off_loop_if_aws
from litellm.types.utils import LlmProviders
@ -818,6 +821,196 @@ class TestBedrockMantleProviderResolution:
)
def _row_cost(key: str, input_tokens: int, output_tokens: int) -> float:
row: Final[Mapping[str, float]] = litellm.model_cost[key]
return input_tokens * row["input_cost_per_token"] + output_tokens * row["output_cost_per_token"]
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]
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(
_row_cost("bedrock_mantle/anthropic.claude-opus-5-5", 10, 5)
)
assert response._hidden_params["response_cost"] != pytest.approx(_row_cost("anthropic.claude-opus-5-5", 10, 5))
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))),
)
assert isinstance(stream, CustomStreamWrapper)
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_claude_region_prefixed_model_sends_bare_model_to_that_region(self, monkeypatch, local_cost_map):
from litellm.llms.custom_httpx.http_handler import HTTPHandler
for var in (
"BEDROCK_MANTLE_API_KEY",
"AWS_BEARER_TOKEN_BEDROCK",
"BEDROCK_MANTLE_API_BASE",
"BEDROCK_MANTLE_REGION",
"AWS_REGION_NAME",
"AWS_REGION",
"AWS_PROFILE",
):
monkeypatch.delenv(var, raising=False)
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIAEXAMPLE")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0")
handler = Mock(side_effect=_anthropic_message)
response = litellm.completion(
model="bedrock_mantle/us-gov-west-1/anthropic.claude-opus-5-5",
messages=[{"role": "user", "content": "hello"}],
max_tokens=64,
client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler))),
)
sent = handler.call_args.args[0]
assert str(sent.url) == "https://bedrock-mantle.us-gov-west-1.api.aws/anthropic/v1/messages"
assert json.loads(sent.content)["model"] == "anthropic.claude-opus-5-5"
assert "/us-gov-west-1/bedrock/aws4_request" in sent.headers["Authorization"]
assert response._hidden_params["response_cost"] == pytest.approx(
_row_cost("bedrock_mantle/us-gov-west-1/anthropic.claude-opus-5-5", 10, 5)
)
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)
def respond(request: httpx.Request) -> httpx.Response:
return 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,
)
handler = Mock(side_effect=respond)
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"
@pytest.mark.parametrize("request_type", ["chat_completion", "embeddings"])
def test_supported_openai_params_follow_the_route_the_model_takes(self, request_type):
claude_params = litellm.get_supported_openai_params(
model="anthropic.claude-opus-5-5", custom_llm_provider="bedrock_mantle", request_type=request_type
)
open_weight_params = litellm.get_supported_openai_params(
model="openai.gpt-oss-120b", custom_llm_provider="bedrock_mantle", request_type=request_type
)
assert claude_params is not None and open_weight_params is not None
assert "thinking" in claude_params
assert "thinking" not in open_weight_params
class TestBedrockMantlePricing:
"""Tests that verify Bedrock Mantle uses correct AWS Bedrock pricing, not OpenAI pricing."""

View file

@ -1,5 +1,6 @@
import datetime
import time
from collections.abc import Mapping
from pathlib import Path
from types import MappingProxyType, SimpleNamespace
from typing import Final, cast
@ -3864,6 +3865,35 @@ def test_completion_cost_mantle_native_messages_prices_unversioned_claude_from_t
) == pytest.approx(expected), model
@pytest.mark.parametrize("model", ["anthropic.claude-opus-5-5", "anthropic.claude-sonnet-5-5"])
def test_completion_cost_region_without_its_own_row_prices_mantle_claude_from_the_mantle_row(
_local_model_cost_map, model: str
):
"""The proxy resolves a Mantle region for every call. A region with no
bedrock_mantle/<region>/<model> row must fall back to the model's own bedrock_mantle/ row, not to the
bare Bedrock row that the bedrock provider family also matches."""
response = litellm.ModelResponse(
id="msg_x",
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
model=model,
usage={"prompt_tokens": 100, "completion_tokens": 10, "total_tokens": 110},
)
mantle: Final[Mapping[str, float]] = litellm.model_cost[f"bedrock_mantle/{model}"]
bedrock: Final[Mapping[str, float]] = litellm.model_cost[model]
expected: Final = 100 * mantle["input_cost_per_token"] + 10 * mantle["output_cost_per_token"]
assert expected != 100 * bedrock["input_cost_per_token"] + 10 * bedrock["output_cost_per_token"]
for deployment in (model, f"bedrock_mantle/{model}", f"bedrock_mantle/us-east-1/{model}"):
assert litellm.completion_cost(
completion_response=response,
model=deployment,
custom_llm_provider="bedrock_mantle",
region_name="us-east-1",
) == pytest.approx(expected), deployment
assert litellm.get_model_info(f"bedrock_mantle/us-east-1/{model}", "bedrock_mantle")["key"] == f"bedrock_mantle/{model}"
@pytest.mark.parametrize("model", ["anthropic.claude-opus-5-5", "anthropic.claude-sonnet-5-5"])
def test_cost_per_token_gov_region_prices_mantle_claude_on_the_gov_row(_local_model_cost_map, model):
"""A bedrock_mantle/ deployment in us-gov-west-1 must price from the