diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index c635cf828eb..77d3b255fd7 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -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": diff --git a/litellm/llms/bedrock_mantle/chat/claude_transformation.py b/litellm/llms/bedrock_mantle/chat/claude_transformation.py new file mode 100644 index 00000000000..883607d3b75 --- /dev/null +++ b/litellm/llms/bedrock_mantle/chat/claude_transformation.py @@ -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() diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py index 9892ef6224e..5a43da95604 100644 --- a/litellm/llms/bedrock_mantle/common_utils.py +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -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. diff --git a/litellm/main.py b/litellm/main.py index a818213b861..b24b35166ed 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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 diff --git a/litellm/utils.py b/litellm/utils.py index d72588e2b00..ac2879ca417 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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.""" diff --git a/tests/integration/providers/test_bedrock_mantle_claude_chat_chaos.py b/tests/integration/providers/test_bedrock_mantle_claude_chat_chaos.py new file mode 100644 index 00000000000..d8a3306ecd5 --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_claude_chat_chaos.py @@ -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) + } diff --git a/tests/integration/providers/test_bedrock_mantle_claude_chat_wire.py b/tests/integration/providers/test_bedrock_mantle_claude_chat_wire.py new file mode 100644 index 00000000000..65c536492a5 --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_claude_chat_wire.py @@ -0,0 +1,1200 @@ +import json +import re +from collections.abc import Callable, Mapping +from typing import Final +from uuid import uuid4 + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.sigv4 import signature +from integration._support.wire import Reply, Request, wire_server +from openai import AsyncOpenAI, OpenAI +from pydantic import JsonValue, TypeAdapter + +from litellm.constants import DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS + +_HAIKU: Final = "anthropic.claude-haiku-4-5" +_OPUS: Final = "anthropic.claude-opus-5-5" +_SAFEGUARD: Final = "openai.gpt-oss-safeguard-120b" +_API_KEY: Final = "synthetic-mantle-bearer" +_ACCESS_KEY: Final = "AKIAINTEGRATION000009" +_SECRET_KEY: Final = "synthetic-secret-key-for-testing" +_MESSAGES_PATH: Final = "/anthropic/v1/messages" +_BRIDGE_VERSION: Final = "bedrock-2023-05-31" +_INPUT_TOKENS: Final = 23 +_OUTPUT_TOKENS: Final = 7 +_USAGE: Final = (_INPUT_TOKENS, _OUTPUT_TOKENS, _INPUT_TOKENS + _OUTPUT_TOKENS) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})") +_CITY_SCHEMA: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} +_WEATHER_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_weather", "description": "Weather lookup", "parameters": _CITY_SCHEMA}, +} +_WEATHER_TOOL_OUTBOUND: Final[dict[str, JsonValue]] = { + "name": "get_weather", + "input_schema": _CITY_SCHEMA, + "type": "custom", + "description": "Weather lookup", +} +_SPEND_COLUMNS: Final = ( + "SELECT request_id, call_type, status, spend, prompt_tokens, completion_tokens, model, model_group," + ' custom_llm_provider, cache_hit, litellm_call_id FROM "LiteLLM_SpendLogs"' +) +_BY_REQUEST_ID: Final = _SPEND_COLUMNS + " WHERE request_id=%s" +_BY_CALL_ID: Final = _SPEND_COLUMNS + " WHERE litellm_call_id=%s" +_BY_MODEL_GROUP: Final = _SPEND_COLUMNS + " WHERE model_group=%s ORDER BY request_id" + + +def _prompt(marker: str) -> str: + return f"mantle claude chat marker-{marker}" + + +def _answer(marker: str) -> str: + return f"answer marker-{marker}" + + +def _user_turn(marker: str) -> dict[str, JsonValue]: + return {"role": "user", "content": [{"type": "text", "text": _prompt(marker)}]} + + +def _chat_messages(marker: str) -> list[JsonValue]: + return [{"role": "user", "content": _prompt(marker)}] + + +def _outbound(backend: str, messages: list[JsonValue], **fields: JsonValue) -> dict[str, JsonValue]: + return { + "messages": messages, + "max_tokens": DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS, + "anthropic_version": _BRIDGE_VERSION, + "model": backend, + **fields, + } + + +def _message_reply(marker: str, backend: str, content: list[JsonValue], stop_reason: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": f"msg_bdrk_{marker}", + "type": "message", + "role": "assistant", + "model": backend, + "content": content, + "stop_reason": stop_reason, + "stop_sequence": None, + "usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS}, + } + ).encode() + ) + + +def _text_reply(marker: str, backend: str = _HAIKU) -> Reply: + return _message_reply(marker, backend, [{"type": "text", "text": _answer(marker)}], "end_turn") + + +def _tool_reply(marker: str, backend: str, name: str) -> Reply: + block: Final[JsonValue] = {"type": "tool_use", "id": f"toolu_{marker}", "name": name, "input": {"city": "Paris"}} + return _message_reply(marker, backend, [block], "tool_use") + + +def _stream_reply( + marker: str, block: Mapping[str, JsonValue], deltas: tuple[Mapping[str, JsonValue], ...], stop_reason: str +) -> Reply: + opening: Final = { + "id": f"msg_bdrk_{marker}", + "type": "message", + "role": "assistant", + "model": _HAIKU, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": 1}, + } + events: Final = ( + ("message_start", {"message": opening}), + ("content_block_start", {"index": 0, "content_block": block}), + *(("content_block_delta", {"index": 0, "delta": delta}) for delta in deltas), + ("content_block_stop", {"index": 0}), + ( + "message_delta", + {"delta": {"stop_reason": stop_reason, "stop_sequence": None}, "usage": {"output_tokens": _OUTPUT_TOKENS}}, + ), + ("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 + ), + ) + + +def _text_stream(marker: str) -> Reply: + return _stream_reply( + marker, + {"type": "text", "text": ""}, + ({"type": "text_delta", "text": "answer "}, {"type": "text_delta", "text": f"marker-{marker}"}), + "end_turn", + ) + + +def _tool_stream(marker: str) -> Reply: + return _stream_reply( + marker, + {"type": "tool_use", "id": f"toolu_{marker}", "name": "get_weather", "input": {}}, + ( + {"type": "input_json_delta", "partial_json": '{"city": '}, + {"type": "input_json_delta", "partial_json": '"Paris"}'}, + ), + "tool_use", + ) + + +def _bearer_peer(expected: Mapping[str, JsonValue], reply: Reply) -> 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) + assert _JSON_OBJECT.validate_json(request.body) == expected, request.body + return reply + + return respond + + +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 _cost_map_row(gateway: Gateway, name: str) -> dict[str, JsonValue]: + return object_value(gateway.get("/public/litellm_model_cost_map")[name]) + + +def _number(value: JsonValue) -> float: + assert isinstance(value, (int, float)) and not isinstance(value, bool), value + return float(value) + + +def _mantle_cost( + gateway: Gateway, backend: str, input_tokens: int = _INPUT_TOKENS, output_tokens: int = _OUTPUT_TOKENS +) -> float: + row: Final = _cost_map_row(gateway, f"bedrock_mantle/{backend}") + cost: Final = input_tokens * _number(row["input_cost_per_token"]) + output_tokens * _number( + row["output_cost_per_token"] + ) + assert cost > 0, row + return cost + + +def _mantle_claude_where(gateway: Gateway, flag: str, disabled: bool) -> str: + rows: Final = gateway.get("/public/litellm_model_cost_map") + names: Final = sorted( + name + for name, row in rows.items() + if name.startswith("bedrock_mantle/anthropic.claude") and (object_value(row).get(flag) is False) is disabled + ) + assert names, f"The cost map has no Mantle Claude row with {flag} disabled={disabled}" + return names[0].removeprefix("bedrock_mantle/") + + +def _spend_row(query: str, identity: str) -> dict[str, JsonValue]: + rows: Final = eventually(lambda: read_rows(query, (identity,)), lambda found: len(found) == 1, seconds=70) + return rows[0] + + +def _assert_success_row( + row: Mapping[str, JsonValue], model_group: str, backend: str, call_type: str, cost: float +) -> None: + assert (row["status"], row["call_type"], row["cache_hit"] == "True") == ("success", call_type, False), row + assert (row["model_group"], row["model"], row["custom_llm_provider"]) == ( + model_group, + f"bedrock_mantle/{backend}", + "bedrock_mantle", + ), row + assert (row["prompt_tokens"], row["completion_tokens"]) == (_INPUT_TOKENS, _OUTPUT_TOKENS), row + assert _number(row["spend"]) == pytest.approx(cost), row + + +def _only_choice(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + choices: Final = body["choices"] + assert isinstance(choices, list) and len(choices) == 1, body + return object_value(choices[0]) + + +def _token_usage(body: Mapping[str, JsonValue]) -> tuple[JsonValue, JsonValue, JsonValue]: + usage: Final = object_value(body["usage"]) + return usage["prompt_tokens"], usage["completion_tokens"], usage["total_tokens"] + + +def _assert_text_completion(response: httpx.Response, model: str, marker: str) -> dict[str, JsonValue]: + assert response.status_code == 200, response.text + body: Final = _JSON_OBJECT.validate_json(response.content) + assert (body["object"], body["model"]) == ("chat.completion", model), response.text + choice: Final = _only_choice(body) + message: Final = object_value(choice["message"]) + assert (choice["finish_reason"], message["role"], message["content"], message.get("tool_calls")) == ( + "stop", + "assistant", + _answer(marker), + None, + ), response.text + assert _token_usage(body) == _USAGE, response.text + return body + + +def _sse_objects(text: str) -> tuple[dict[str, JsonValue], ...]: + data: Final = tuple(line.removeprefix("data: ") for line in text.splitlines() if line.startswith("data: ")) + assert data and data[-1] == "[DONE]", text + return tuple(_JSON_OBJECT.validate_json(item) for item in data[:-1]) + + +def _assert_chunk_envelope(chunks: tuple[dict[str, JsonValue], ...], model: str, text: str) -> str: + assert {(chunk["object"], chunk["model"]) for chunk in chunks} == {("chat.completion.chunk", model)}, text + identities: Final = {string_value(chunk["id"]) for chunk in chunks} + assert len(identities) == 1, text + return next(iter(identities)) + + +def _finish_reasons(chunks: tuple[dict[str, JsonValue], ...]) -> list[JsonValue]: + reasons: Final = (_only_choice(chunk).get("finish_reason") for chunk in chunks) + return [reason for reason in reasons if reason is not None] + + +def _deltas(chunks: tuple[dict[str, JsonValue], ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple(object_value(_only_choice(chunk)["delta"]) for chunk in chunks) + + +def _sdk_base(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + "/v1" + + +def _error_message(response: httpx.Response) -> str: + return string_value(object_value(_JSON_OBJECT.validate_json(response.content)["error"])["message"]) + + +@pytest.mark.parametrize("backend", [_HAIKU, _OPUS], ids=["haiku", "opus_with_a_cheaper_bare_twin"]) +def test_chat_completion_on_a_mantle_claude_id_posts_native_messages_and_spends_at_the_mantle_price( + gateway: Gateway, backend: str +) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(backend, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker, backend))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{backend}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + body: Final = _assert_text_completion(response, model, marker) + cost: Final = _mantle_cost(gateway, backend) + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(cost), response.text + _assert_success_row(_spend_row(_BY_REQUEST_ID, string_value(body["id"])), model, backend, "acompletion", cost) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +@pytest.mark.parametrize( + ("stream_options", "usage_chunks"), + [ + pytest.param({"stream_options": {"include_usage": True}}, [_USAGE], id="include_usage"), + pytest.param({}, [], id="no_usage"), + ], +) +def test_chat_stream_on_a_mantle_claude_id_relays_openai_chunks_from_the_native_stream( + gateway: Gateway, stream_options: dict[str, JsonValue], usage_chunks: list[tuple[int, int, int]] +) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)], stream=True) + with wire_server(_bearer_peer(expected, _text_stream(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "stream": True, "messages": _chat_messages(marker), **stream_options}, + ) + assert response.status_code == 200, response.text + chunks: Final = _sse_objects(response.text) + identity: Final = _assert_chunk_envelope(chunks, model, response.text) + assert "".join(str(delta.get("content") or "") for delta in _deltas(chunks)) == _answer(marker), response.text + assert _finish_reasons(chunks) == ["stop"], response.text + assert [_token_usage(chunk) for chunk in chunks if chunk.get("usage") is not None] == usage_chunks, ( + response.text + ) + _assert_success_row( + _spend_row(_BY_REQUEST_ID, identity), model, _HAIKU, "acompletion", _mantle_cost(gateway, _HAIKU) + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +@pytest.mark.parametrize( + ("leading_messages", "request_fields", "outbound_fields"), + [ + pytest.param( + [{"role": "system", "content": "be terse"}], + {}, + {"system": [{"type": "text", "text": "be terse"}]}, + id="system_prompt", + ), + pytest.param([], {"max_tokens": 77}, {"max_tokens": 77}, id="explicit_max_tokens"), + ], +) +def test_chat_request_fields_are_translated_to_the_native_messages_shape( + gateway: Gateway, + leading_messages: list[JsonValue], + request_fields: dict[str, JsonValue], + outbound_fields: dict[str, JsonValue], +) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)], **outbound_fields) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [*leading_messages, *_chat_messages(marker)], **request_fields}, + ) + _assert_text_completion(response, model, marker) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +@pytest.mark.parametrize( + ("tool_choice", "outbound_choice"), + [pytest.param("auto", {"type": "auto"}, id="auto"), pytest.param("required", {"type": "any"}, id="required")], +) +def test_chat_tools_reach_mantle_as_input_schema_and_the_tool_use_block_returns_as_a_tool_call( + gateway: Gateway, tool_choice: str, outbound_choice: dict[str, JsonValue] +) -> None: + marker: Final = uuid4().hex + backend: Final = _mantle_claude_where(gateway, "supports_forced_tool_use", disabled=False) + expected: Final = _outbound( + backend, [_user_turn(marker)], tools=[_WEATHER_TOOL_OUTBOUND], tool_choice=outbound_choice + ) + reply: Final = _tool_reply(marker, backend, "get_weather") + with wire_server(_bearer_peer(expected, reply)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{backend}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": _chat_messages(marker), "tools": [_WEATHER_TOOL], "tool_choice": tool_choice}, + ) + assert response.status_code == 200, response.text + choice: Final = _only_choice(_JSON_OBJECT.validate_json(response.content)) + message: Final = object_value(choice["message"]) + calls: Final = message["tool_calls"] + assert isinstance(calls, list) and len(calls) == 1, response.text + call: Final = object_value(calls[0]) + function: Final = object_value(call["function"]) + assert (choice["finish_reason"], message["content"], call["id"], call["type"], function["name"]) == ( + "tool_calls", + None, + f"toolu_{marker}", + "function", + "get_weather", + ), response.text + assert json.loads(string_value(function["arguments"])) == {"city": "Paris"}, response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_chat_stream_relays_a_native_tool_use_block_as_tool_call_chunks(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)], tools=[_WEATHER_TOOL_OUTBOUND], stream=True) + with wire_server(_bearer_peer(expected, _tool_stream(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "stream": True, "messages": _chat_messages(marker), "tools": [_WEATHER_TOOL]}, + ) + assert response.status_code == 200, response.text + chunks: Final = _sse_objects(response.text) + _assert_chunk_envelope(chunks, model, response.text) + fragments: Final = tuple( + object_value(object_value(calls[0])["function"]) + for calls in (delta.get("tool_calls") for delta in _deltas(chunks)) + if isinstance(calls, list) + ) + opening: Final = object_value(object_value(_deltas(chunks)[0])["tool_calls"][0]) + assert (opening["id"], opening["type"], fragments[0]["name"]) == ( + f"toolu_{marker}", + "function", + "get_weather", + ), response.text + arguments: Final = "".join(string_value(fragment["arguments"]) for fragment in fragments) + assert json.loads(arguments) == {"city": "Paris"}, response.text + assert _finish_reasons(chunks) == ["tool_calls"], response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_chat_tool_result_turn_reaches_mantle_as_tool_use_and_tool_result_blocks(gateway: Gateway) -> None: + marker: Final = uuid4().hex + tool_call: Final = f"toolu_{marker}" + expected: Final = _outbound( + _HAIKU, + [ + _user_turn(marker), + { + "role": "assistant", + "content": [{"type": "tool_use", "id": tool_call, "name": "get_weather", "input": {"city": "Paris"}}], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": tool_call, "content": "sunny"}]}, + ], + tools=[_WEATHER_TOOL_OUTBOUND], + ) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "tools": [_WEATHER_TOOL], + "messages": [ + *_chat_messages(marker), + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": tool_call, + "type": "function", + "function": {"name": "get_weather", "arguments": json.dumps({"city": "Paris"})}, + } + ], + }, + {"role": "tool", "tool_call_id": tool_call, "content": "sunny"}, + ], + }, + ) + _assert_text_completion(response, model, marker) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_chat_json_schema_response_format_is_sent_as_a_forced_tool_and_returned_as_json_content( + gateway: Gateway, +) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound( + _HAIKU, + [_user_turn(marker)], + tools=[{"name": "json_tool_call", "input_schema": _CITY_SCHEMA}], + tool_choice={"name": "json_tool_call", "type": "tool"}, + ) + reply: Final = _tool_reply(marker, _HAIKU, "json_tool_call") + with wire_server(_bearer_peer(expected, reply)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": _chat_messages(marker), + "response_format": {"type": "json_schema", "json_schema": {"name": "weather", "schema": _CITY_SCHEMA}}, + }, + ) + assert response.status_code == 200, response.text + choice: Final = _only_choice(_JSON_OBJECT.validate_json(response.content)) + message: Final = object_value(choice["message"]) + assert (choice["finish_reason"], message.get("tool_calls")) == ("stop", None), response.text + assert json.loads(string_value(message["content"])) == {"city": "Paris"}, response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def _authorization_field(part: str) -> tuple[str, str]: + name, _, value = part.partition("=") + return name, value + + +def _sigv4_peer(expected: Mapping[str, JsonValue], reply: Reply) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", _MESSAGES_PATH), request.target + authorization: Final = request.headers["authorization"] + assert authorization.startswith("AWS4-HMAC-SHA256 "), sorted(request.headers) + fields: Final = dict( + _authorization_field(part) for part in authorization.removeprefix("AWS4-HMAC-SHA256 ").split(", ") + ) + access_key, scope = fields["Credential"].split("/", 1) + assert access_key == _ACCESS_KEY, authorization + assert scope == f"{request.headers['x-amz-date'][:8]}/us-east-1/bedrock/aws4_request", authorization + signed: Final = signature( + "POST", _MESSAGES_PATH, request.headers, fields["SignedHeaders"], request.body, _SECRET_KEY, scope + ) + assert fields["Signature"] == signed[1], authorization + assert _JSON_OBJECT.validate_json(request.body) == expected, request.body + return reply + + return respond + + +def test_chat_completion_signs_the_native_messages_request_with_the_deployment_aws_keys(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_sigv4_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock_mantle/{_HAIKU}", + api_base=wire.url, + api_key=None, + aws_access_key_id=_ACCESS_KEY, + aws_secret_access_key=_SECRET_KEY, + aws_region_name="us-east-1", + ) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + _assert_text_completion(response, model, marker) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_region_prefix_in_a_mantle_claude_id_is_routing_only_and_keeps_the_mantle_price(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/us-east-2/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + body: Final = _assert_text_completion(response, model, marker) + cost: Final = _mantle_cost(gateway, _HAIKU) + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(cost), response.text + row: Final = _spend_row(_BY_REQUEST_ID, string_value(body["id"])) + assert (row["status"], row["model_group"], row["custom_llm_provider"]) == ( + "success", + model, + "bedrock_mantle", + ), row + assert _number(row["spend"]) == pytest.approx(cost), row + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_openai_sdk_chat_completion_on_a_mantle_claude_id(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + with OpenAI( + api_key=gateway.key, + base_url=_sdk_base(gateway), + max_retries=0, + http_client=httpx.Client(timeout=30, trust_env=False), + ) as client: + completion: Final = client.chat.completions.create( + model=model, messages=[{"role": "user", "content": _prompt(marker)}] + ) + assert (completion.object, completion.model) == ("chat.completion", model), completion + assert (completion.choices[0].message.content, completion.choices[0].finish_reason) == ( + _answer(marker), + "stop", + ), completion + assert completion.usage is not None + assert (completion.usage.prompt_tokens, completion.usage.completion_tokens, completion.usage.total_tokens) == ( + _USAGE + ), completion + _assert_success_row( + _spend_row(_BY_REQUEST_ID, completion.id), model, _HAIKU, "acompletion", _mantle_cost(gateway, _HAIKU) + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_openai_sdk_chat_stream_on_a_mantle_claude_id(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)], stream=True) + with wire_server(_bearer_peer(expected, _text_stream(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + with OpenAI( + api_key=gateway.key, + base_url=_sdk_base(gateway), + max_retries=0, + http_client=httpx.Client(timeout=30, trust_env=False), + ) as client: + chunks: Final = tuple( + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": _prompt(marker)}], + stream=True, + stream_options={"include_usage": True}, + ) + ) + identities: Final = {chunk.id for chunk in chunks} + assert len(identities) == 1 and {chunk.model for chunk in chunks} == {model}, chunks + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == _answer(marker), chunks + assert [chunk.choices[0].finish_reason for chunk in chunks if chunk.choices[0].finish_reason] == ["stop"], ( + chunks + ) + assert [ + (chunk.usage.prompt_tokens, chunk.usage.completion_tokens, chunk.usage.total_tokens) + for chunk in chunks + if chunk.usage is not None + ] == [_USAGE], chunks + _assert_success_row( + _spend_row(_BY_REQUEST_ID, next(iter(identities))), + model, + _HAIKU, + "acompletion", + _mantle_cost(gateway, _HAIKU), + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +async def test_async_openai_sdk_chat_completion_on_a_mantle_claude_id(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + async with AsyncOpenAI( + api_key=gateway.key, + base_url=_sdk_base(gateway), + max_retries=0, + http_client=httpx.AsyncClient(timeout=30, trust_env=False), + ) as client: + completion: Final = await client.chat.completions.create( + model=model, messages=[{"role": "user", "content": _prompt(marker)}] + ) + assert (completion.object, completion.model) == ("chat.completion", model), completion + assert (completion.choices[0].message.content, completion.choices[0].finish_reason) == ( + _answer(marker), + "stop", + ), completion + assert completion.usage is not None + assert (completion.usage.prompt_tokens, completion.usage.completion_tokens, completion.usage.total_tokens) == ( + _USAGE + ), completion + _assert_success_row( + _spend_row(_BY_REQUEST_ID, completion.id), model, _HAIKU, "acompletion", _mantle_cost(gateway, _HAIKU) + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +async def test_async_openai_sdk_chat_stream_on_a_mantle_claude_id(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)], stream=True) + with wire_server(_bearer_peer(expected, _text_stream(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + async with AsyncOpenAI( + api_key=gateway.key, + base_url=_sdk_base(gateway), + max_retries=0, + http_client=httpx.AsyncClient(timeout=30, trust_env=False), + ) as client: + stream: Final = await client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": _prompt(marker)}], + stream=True, + stream_options={"include_usage": True}, + ) + chunks: Final = tuple([chunk async for chunk in stream]) + identities: Final = {chunk.id for chunk in chunks} + assert len(identities) == 1 and {chunk.model for chunk in chunks} == {model}, chunks + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == _answer(marker), chunks + assert [chunk.choices[0].finish_reason for chunk in chunks if chunk.choices[0].finish_reason] == ["stop"], ( + chunks + ) + assert [ + (chunk.usage.prompt_tokens, chunk.usage.completion_tokens, chunk.usage.total_tokens) + for chunk in chunks + if chunk.usage is not None + ] == [_USAGE], chunks + _assert_success_row( + _spend_row(_BY_REQUEST_ID, next(iter(identities))), + model, + _HAIKU, + "acompletion", + _mantle_cost(gateway, _HAIKU), + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_identical_chat_request_is_served_from_the_response_cache_and_logged_as_a_free_cache_hit( + gateway: Gateway, +) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + request: Final[dict[str, JsonValue]] = {"model": model, "messages": _chat_messages(marker)} + first: Final = gateway.request("POST", "/v1/chat/completions", request) + identity: Final = string_value(_assert_text_completion(first, model, marker)["id"]) + assert "x-litellm-cache-key" not in first.headers, sorted(first.headers) + second: Final = gateway.request("POST", "/v1/chat/completions", request) + assert _assert_text_completion(second, model, marker)["id"] == identity, second.text + assert second.headers["x-litellm-cache-key"], sorted(second.headers) + rows: Final = eventually( + lambda: read_rows(_BY_MODEL_GROUP, (model,)), lambda found: len(found) == 2, seconds=70 + ) + priced, cached = rows + _assert_success_row(priced, model, _HAIKU, "acompletion", _mantle_cost(gateway, _HAIKU)) + assert priced["request_id"] == identity, rows + assert string_value(cached["request_id"]).startswith(f"{identity}_cache_hit"), rows + assert (cached["status"], cached["cache_hit"], _number(cached["spend"])) == ("success", "True", 0), rows + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def _open_weight_peer(prompt: str, identity: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", "/v1/chat/completions"), request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}", sorted(request.headers) + assert _JSON_OBJECT.validate_json(request.body) == { + "model": _SAFEGUARD, + "messages": [{"role": "user", "content": prompt}], + "temperature": 0.3, + "stream": False, + }, request.body + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": _SAFEGUARD, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "safeguard answer"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 3, "total_tokens": 12}, + } + ).encode() + ) + + return respond + + +def test_open_weight_mantle_id_keeps_the_openai_compatible_chat_route_and_its_sampling_params( + gateway: Gateway, +) -> None: + marker: Final = uuid4().hex + identity: Final = f"chatcmpl-safeguard-{marker}" + with wire_server(_open_weight_peer(_prompt(marker), identity)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_SAFEGUARD}", api_base=wire.url + "/v1", api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": _chat_messages(marker), "temperature": 0.3, "stream": False}, + ) + assert response.status_code == 200, response.text + body: Final = _JSON_OBJECT.validate_json(response.content) + message: Final = object_value(_only_choice(body)["message"]) + assert (body["id"], message["role"], message["content"]) == ( + identity, + "assistant", + "safeguard answer", + ), response.text + cost: Final = _mantle_cost(gateway, _SAFEGUARD, input_tokens=9, output_tokens=3) + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(cost), response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/v1/chat/completions")] + + +@pytest.mark.parametrize( + "suffix", + [ + pytest.param("/v1", id="v1"), + pytest.param("/anthropic/v1/messages", id="anthropic_v1_messages"), + pytest.param("/openai/v1", id="openai_v1"), + ], +) +def test_api_base_with_a_route_suffix_still_posts_to_the_native_messages_path(gateway: Gateway, suffix: str) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url + suffix, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + _assert_text_completion(response, model, marker) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_responses_api_on_a_mantle_claude_id_bridges_to_native_messages(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": _prompt(marker)}) + assert response.status_code == 200, response.text + body: Final = _JSON_OBJECT.validate_json(response.content) + assert (body["object"], body["status"], body["model"]) == ("response", "completed", model), response.text + output: Final = body["output"] + assert isinstance(output, list) and len(output) == 1, response.text + content: Final = object_value(output[0])["content"] + assert isinstance(content, list) and len(content) == 1, response.text + part: Final = object_value(content[0]) + assert (part["type"], part["text"]) == ("output_text", _answer(marker)), response.text + _assert_success_row( + _spend_row(_BY_CALL_ID, response.headers["x-litellm-call-id"]), + model, + _HAIKU, + "aresponses", + _mantle_cost(gateway, _HAIKU), + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_responses_api_stream_on_a_mantle_claude_id_bridges_to_the_native_stream(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final = _outbound(_HAIKU, [_user_turn(marker)], stream=True) + with wire_server(_bearer_peer(expected, _text_stream(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": _prompt(marker), "stream": True} + ) + assert response.status_code == 200, response.text + events: Final = _sse_objects(response.text) + kinds: Final = [event["type"] for event in events] + assert (kinds[0], kinds[-1]) == ("response.created", "response.completed"), kinds + assert "".join( + string_value(event["delta"]) for event in events if event["type"] == "response.output_text.delta" + ) == _answer(marker), response.text + assert {event["model"] for event in events} == {model}, response.text + completed: Final = object_value(events[-1]["response"]) + usage: Final = object_value(completed["usage"]) + assert (completed["status"], usage["input_tokens"], usage["output_tokens"]) == ( + "completed", + _INPUT_TOKENS, + _OUTPUT_TOKENS, + ), response.text + _assert_success_row( + _spend_row(_BY_CALL_ID, response.headers["x-litellm-call-id"]), + model, + _HAIKU, + "aresponses", + _mantle_cost(gateway, _HAIKU), + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def _health_peer(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", _MESSAGES_PATH), request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}", sorted(request.headers) + body: Final = _JSON_OBJECT.validate_json(request.body) + assert (body["model"], body["anthropic_version"]) == (_HAIKU, _BRIDGE_VERSION), request.body + return _text_reply(uuid4().hex) + + +def _health_counts(gateway: Gateway, model: str) -> tuple[int, JsonValue, JsonValue]: + response: Final = gateway.request("GET", "/health", params={"model": model}) + body: Final = _JSON_OBJECT.validate_json(response.content) + return response.status_code, body.get("healthy_count"), body.get("unhealthy_count") + + +def test_health_check_on_a_mantle_claude_deployment_probes_the_native_messages_route(gateway: Gateway) -> None: + with wire_server(_health_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + counts: Final = eventually( + lambda: _health_counts(gateway, model), + lambda found: found[1:] != (0, 0), + seconds=30, + return_last_on_timeout=True, + ) + assert counts == (200, 1, 0), counts + assert {(request.method, request.target) for request in wire.drain()} == {("POST", _MESSAGES_PATH)} + + +def _native_messages_peer(expected: Mapping[str, JsonValue], reply: Reply) -> 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) + assert request.headers["anthropic-version"] == "2023-06-01", sorted(request.headers) + assert _JSON_OBJECT.validate_json(request.body) == expected, request.body + return reply + + return respond + + +def test_messages_api_on_a_mantle_claude_id_keeps_forwarding_the_caller_shape(gateway: Gateway) -> None: + marker: Final = uuid4().hex + expected: Final[dict[str, JsonValue]] = {"messages": _chat_messages(marker), "max_tokens": 64, "model": _HAIKU} + with wire_server(_native_messages_peer(expected, _text_reply(marker))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 64, "messages": _chat_messages(marker)}, + headers={"anthropic-version": "2023-06-01"}, + ) + assert response.status_code == 200, response.text + body: Final = _JSON_OBJECT.validate_json(response.content) + assert (body["id"], body["type"], body["model"], body["stop_reason"]) == ( + f"msg_bdrk_{marker}", + "message", + model, + "end_turn", + ), response.text + assert body["content"] == [{"type": "text", "text": _answer(marker)}], response.text + _assert_success_row( + _spend_row(_BY_REQUEST_ID, f"msg_bdrk_{marker}"), + model, + _HAIKU, + "anthropic_messages", + _mantle_cost(gateway, _HAIKU), + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +_UNSUPPORTED: Final = ( + pytest.param(None, {"n": 2}, {}, "does not support parameters: ['n']", id="n"), + pytest.param( + "supports_sampling_params", + {"temperature": 0.2}, + {}, + "Only temperature=1 is supported", + id="temperature_on_a_fixed_sampling_model", + ), + pytest.param( + "supports_forced_tool_use", + {"tools": [_WEATHER_TOOL], "tool_choice": "required"}, + {"tools": [_WEATHER_TOOL_OUTBOUND], "tool_choice": {"type": "auto"}}, + "does not support forced tool use", + id="forced_tool_use_on_a_model_without_it", + ), +) + + +def _backend_without(gateway: Gateway, flag: str | None) -> str: + return _HAIKU if flag is None else _mantle_claude_where(gateway, flag, disabled=True) + + +@pytest.mark.parametrize(("flag", "request_fields", "outbound_fields", "reason"), _UNSUPPORTED) +def test_unsupported_claude_param_is_a_400_that_never_reaches_mantle( + gateway: Gateway, + flag: str | None, + request_fields: dict[str, JsonValue], + outbound_fields: dict[str, JsonValue], + reason: str, +) -> None: + marker: Final = uuid4().hex + backend: Final = _backend_without(gateway, flag) + with wire_server(lambda request: _text_reply(_marker_of(request), backend)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{backend}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker), **request_fields} + ) + assert response.status_code == 400, response.text + message: Final = _error_message(response) + assert "UnsupportedParamsError" in message and reason in message, response.text + assert wire.drain() == () + + +@pytest.mark.parametrize(("flag", "request_fields", "outbound_fields", "reason"), _UNSUPPORTED) +def test_unsupported_claude_param_is_dropped_under_drop_params( + gateway: Gateway, + flag: str | None, + request_fields: dict[str, JsonValue], + outbound_fields: dict[str, JsonValue], + reason: str, +) -> None: + marker: Final = uuid4().hex + backend: Final = _backend_without(gateway, flag) + expected: Final = _outbound(backend, [_user_turn(marker)], **outbound_fields) + with wire_server(_bearer_peer(expected, _text_reply(marker, backend))) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{backend}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": _chat_messages(marker), "drop_params": True, **request_fields}, + ) + _assert_text_completion(response, model, marker) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def _error_reply(status: int, kind: str, message: str) -> Reply: + return Reply( + status=status, body=json.dumps({"type": "error", "error": {"type": kind, "message": message}}).encode() + ) + + +def _failing_peer(failing: str, error: Reply) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", _MESSAGES_PATH), request.target + marker: Final = _marker_of(request) + return error if marker == failing else _text_reply(marker) + + return respond + + +@pytest.mark.parametrize( + ("upstream", "status", "error_class"), + [ + pytest.param( + _error_reply(400, "invalid_request_error", "messages: synthetic bad request"), + 400, + "BadRequestError", + id="bad_request", + ), + pytest.param( + _error_reply(401, "authentication_error", "synthetic invalid bearer"), + 401, + "AuthenticationError", + id="unauthorized", + ), + pytest.param( + _error_reply(429, "rate_limit_error", "synthetic throttle"), 429, "RateLimitError", id="rate_limited" + ), + pytest.param( + _error_reply(500, "api_error", "synthetic upstream fault"), + 503, + "ServiceUnavailableError", + id="upstream_fault", + ), + pytest.param( + _error_reply(400, "invalid_request_error", "prompt is too long: 250000 tokens > 200000 maximum"), + 400, + "ContextWindowExceededError", + id="context_overflow", + ), + ], +) +def test_mantle_error_on_the_native_route_reaches_the_caller_and_the_deployment_keeps_serving( + gateway: Gateway, upstream: Reply, status: int, error_class: str +) -> None: + failing: Final = uuid4().hex + following: Final = uuid4().hex + upstream_message: Final = string_value(object_value(_JSON_OBJECT.validate_json(upstream.body)["error"])["message"]) + with wire_server(_failing_peer(failing, upstream)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_HAIKU}", api_base=wire.url, api_key=_API_KEY) + failed: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(failing)} + ) + assert failed.status_code == status, failed.text + message: Final = _error_message(failed) + assert error_class in message and upstream_message in message, failed.text + row: Final = _spend_row(_BY_CALL_ID, failed.headers["x-litellm-call-id"]) + assert (row["status"], row["model_group"], _number(row["spend"])) == ("failure", model, 0), row + served: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(following)} + ) + _assert_text_completion(served, model, following) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] * 2 + + +@pytest.mark.parametrize( + "model", + [ + pytest.param(123, id="int"), + pytest.param([f"bedrock_mantle/{_HAIKU}"], id="list"), + pytest.param("", id="empty_string"), + ], +) +def test_malformed_model_field_is_a_400_with_an_error_body_and_the_proxy_keeps_serving( + gateway: Gateway, model: JsonValue +) -> None: + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(uuid4().hex)} + ) + assert response.status_code == 400, response.text + assert "model" in _error_message(response).lower(), response.text + assert gateway.request("GET", "/health/liveliness").status_code == 200 + + +def _wildcard_prefix(gateway: Gateway, scenario: Scenario, api_base: str) -> str: + prefix: Final = f"integration-{uuid4().hex}" + created: Final = gateway.post( + "/model/new", + { + "model_name": f"{prefix}/*", + "litellm_params": {"model": "bedrock_mantle/*", "api_base": api_base, "api_key": _API_KEY}, + "model_info": {}, + }, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return prefix + + +def test_uppercase_claude_id_through_a_mantle_wildcard_is_sent_to_native_messages(gateway: Gateway) -> None: + marker: Final = uuid4().hex + backend: Final = _HAIKU.upper() + expected: Final = _outbound(backend, [_user_turn(marker)]) + with wire_server(_bearer_peer(expected, _text_reply(marker, backend))) as wire, gateway.scenario() as scenario: + model: Final = f"{_wildcard_prefix(gateway, scenario, wire.url)}/{backend}" + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + _assert_text_completion(response, model, marker) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +def test_five_kilobyte_claude_id_through_a_mantle_wildcard_reaches_mantle_and_its_404_reaches_the_caller( + gateway: Gateway, +) -> None: + marker: Final = uuid4().hex + backend: Final = f"{_HAIKU}-{'x' * 5000}" + expected: Final = _outbound(backend, [_user_turn(marker)]) + unknown: Final = _error_reply(404, "not_found_error", "model: synthetic unknown model") + with wire_server(_bearer_peer(expected, unknown)) as wire, gateway.scenario() as scenario: + model: Final = f"{_wildcard_prefix(gateway, scenario, wire.url)}/{backend}" + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + assert response.status_code == 404, response.text[:600] + message: Final = _error_message(response) + assert "NotFoundError" in message and "model: synthetic unknown model" in message, message[:600] + row: Final = _spend_row(_BY_CALL_ID, response.headers["x-litellm-call-id"]) + assert (row["status"], row["model_group"]) == ("failure", model), row["status"] + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] + + +_LUNA: Final = "openai.gpt-6-luna" +_RESPONSES_PATH: Final = "/openai/v1/responses" + + +def _responses_bridge_peer(marker: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", _RESPONSES_PATH), request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}", sorted(request.headers) + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _LUNA, body + assert _prompt(marker) in json.dumps(body["input"]), body + return Reply( + body=json.dumps( + { + "id": f"resp-{marker}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": _LUNA, + "output": [ + { + "type": "message", + "id": f"msg-{marker}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": _answer(marker), "annotations": []}], + } + ], + "usage": { + "input_tokens": _INPUT_TOKENS, + "output_tokens": _OUTPUT_TOKENS, + "total_tokens": _INPUT_TOKENS + _OUTPUT_TOKENS, + }, + } + ).encode() + ) + + return respond + + +def test_non_claude_mantle_id_with_a_bare_bedrock_twin_keeps_its_route_and_spends_at_the_mantle_price( + gateway: Gateway, +) -> None: + marker: Final = uuid4().hex + with wire_server(_responses_bridge_peer(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"bedrock_mantle/{_LUNA}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _chat_messages(marker)} + ) + assert response.status_code == 200, response.text + body: Final = _JSON_OBJECT.validate_json(response.content) + message: Final = object_value(_only_choice(body)["message"]) + assert (body["model"], message["role"], message["content"]) == (model, "assistant", _answer(marker)), ( + response.text + ) + assert _token_usage(body) == _USAGE, response.text + cost: Final = _mantle_cost(gateway, _LUNA) + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(cost), response.text + _assert_success_row( + _spend_row(_BY_CALL_ID, response.headers["x-litellm-call-id"]), model, _LUNA, "responses", cost + ) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _RESPONSES_PATH)] diff --git a/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py index 4bf3dd11fa1..f9b183e10ec 100644 --- a/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py +++ b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py @@ -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.""" diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 36e188e82d6..9a7c22f6f82 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -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// 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