mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(bedrock_mantle): route Claude chat completions to the native Messages endpoint (#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:
parent
0ea166c160
commit
01b4ffe16b
9 changed files with 1853 additions and 77 deletions
|
|
@ -3,6 +3,7 @@ from typing import Final, Literal
|
|||
import litellm
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
|
||||
from litellm.llms.bedrock_mantle.chat.claude_transformation import bedrock_mantle_chat_config
|
||||
from litellm.types.utils import LlmProviders, LlmProvidersSet
|
||||
|
||||
|
||||
|
|
@ -104,7 +105,7 @@ def get_supported_openai_params(
|
|||
elif custom_llm_provider == "groq":
|
||||
return litellm.GroqChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "bedrock_mantle":
|
||||
return litellm.BedrockMantleChatConfig().get_supported_openai_params(model=model)
|
||||
return bedrock_mantle_chat_config(model).get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "hosted_vllm":
|
||||
return litellm.HostedVLLMChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "vllm":
|
||||
|
|
|
|||
33
litellm/llms/bedrock_mantle/chat/claude_transformation.py
Normal file
33
litellm/llms/bedrock_mantle/chat/claude_transformation.py
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
from litellm.llms.base_llm.chat.transformation import BaseConfig
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.bedrock.chat.mantle.transformation import AmazonMantleConfig
|
||||
from litellm.llms.bedrock_mantle.chat.transformation import BedrockMantleChatConfig
|
||||
from litellm.llms.bedrock_mantle.common_utils import BedrockMantleAuthMixin, is_mantle_claude_model
|
||||
from litellm.llms.bedrock_mantle.messages.transformation import build_mantle_native_messages_url
|
||||
|
||||
|
||||
class BedrockMantleClaudeChatConfig(BedrockMantleAuthMixin, AmazonMantleConfig):
|
||||
def __init__(self, aws_signer: BaseAWSLLM | None = None) -> None:
|
||||
AmazonMantleConfig.__init__(self)
|
||||
self._aws_signer = aws_signer or self
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> str | None:
|
||||
return "bedrock_mantle"
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: dict[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()
|
||||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
124
litellm/utils.py
124
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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
1200
tests/integration/providers/test_bedrock_mantle_claude_chat_wire.py
Normal file
1200
tests/integration/providers/test_bedrock_mantle_claude_chat_wire.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue