This commit is contained in:
Mateo Wang 2026-06-22 23:09:37 +08:00 • committed by GitHub
commit d22ab81b5f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 409 additions and 1232 deletions

1
.gitignore vendored
View file

@ -101,6 +101,7 @@ STABILIZATION_TODO.md
**/playwright-report
**/*.storageState.json
**/coverage
.coverage
test-config
# ---------- Terraform ----------

View file

@ -14,6 +14,8 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS
from litellm.types.llms.bedrock_invoke import assert_no_control_params_in_payload
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -37,6 +39,7 @@ def make_sync_call(
if client is None:
client = _get_httpx_client() # Create a new client if none provided
assert_no_control_params_in_payload(data)
response = client.post(
api_base,
headers=headers,
@ -268,7 +271,9 @@ class BedrockConverseLLM(BaseAWSLLM):
):
## SETUP ##
stream = optional_params.pop("stream", None)
stream_chunk_size = optional_params.pop("stream_chunk_size", None)
stream_chunk_size = optional_params.get("stream_chunk_size")
for _control_key in LITELLM_CONTROL_PARAM_KEYS:
optional_params.pop(_control_key, None)
unencoded_model_id = optional_params.pop("model_id", None)
fake_stream = optional_params.pop("fake_stream", False)
json_mode = optional_params.get("json_mode", False)

File diff suppressed because it is too large Load diff

View file

@ -154,34 +154,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
return _anthropic_request
async def async_transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
_anthropic_request = self._build_bedrock_anthropic_request_base(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
await self._async_convert_document_url_sources_to_base64(_anthropic_request)
beta_list = self._compute_bedrock_invoke_beta_headers(
model=model,
messages=messages,
optional_params=optional_params,
headers=headers,
)
if beta_list:
_anthropic_request["anthropic_beta"] = beta_list
return _anthropic_request
def _build_bedrock_anthropic_request_base(
self,
model: str,
@ -190,11 +162,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
litellm_params: dict,
headers: dict,
) -> dict:
filtered_params = {
k: v
for k, v in optional_params.items()
if k not in self.aws_authentication_params
}
filtered_params = self.filter_invoke_request_params(optional_params)
output_config = filtered_params.get("output_config")
if isinstance(output_config, dict):
filtered_params["output_config"] = dict(output_config)
@ -215,7 +183,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
anthropic_request.pop("model", None)
anthropic_request.pop("stream", None)
anthropic_request.pop("stream_chunk_size", None)
output_format = anthropic_request.pop("output_format", None)
output_config_format = pop_bedrock_invoke_output_config_format(
anthropic_request

View file

@ -24,6 +24,8 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
_get_httpx_client,
)
from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS
from litellm.types.llms.bedrock_invoke import parse_invoke_inference_params
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, Usage
from litellm.utils import CustomStreamWrapper
@ -140,6 +142,14 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
if k not in inference_params:
inference_params[k] = v
def filter_invoke_request_params(self, optional_params: dict) -> dict:
return {
k: v
for k, v in optional_params.items()
if k not in self.aws_authentication_params
and k not in LITELLM_CONTROL_PARAM_KEYS
}
def transform_request(
self,
model: str,
@ -150,7 +160,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
) -> dict:
## SETUP ##
stream = optional_params.pop("stream", None)
optional_params.pop("stream_chunk_size", None)
custom_prompt_dict: dict = litellm_params.pop("custom_prompt_dict", None) or {}
hf_model_name = litellm_params.get("hf_model_name", None)
@ -162,12 +171,13 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
provider=provider,
custom_prompt_dict=custom_prompt_dict,
)
inference_params = copy.deepcopy(optional_params)
inference_params = {
k: v
for k, v in inference_params.items()
if k not in self.aws_authentication_params
}
drop_params = bool(litellm_params.get("drop_params") or litellm.drop_params)
inference_params = parse_invoke_inference_params(
provider=provider,
model=model,
params=self.filter_invoke_request_params(copy.deepcopy(optional_params)),
drop_params=drop_params,
)
request_data: dict = {}
if provider == "cohere":
if model.startswith("cohere.command-r"):

View file

@ -1,11 +1,24 @@
import json
from typing import Any, Dict, List, Literal, Optional, Union
from pydantic import BaseModel
from typing_extensions import TYPE_CHECKING, Required, TypedDict, override
from .openai import ChatCompletionToolCallChunk
class LiteLLMControlParams(BaseModel):
"""LiteLLM-internal control parameters that must never be serialized into a
provider request body. They live in optional_params for convenience but
govern client-side behavior (e.g. how the HTTP response stream is
re-chunked), so Bedrock rejects them as unknown fields."""
stream_chunk_size: int | None = None
LITELLM_CONTROL_PARAM_KEYS = frozenset(LiteLLMControlParams.model_fields)
class CachePointBlock(TypedDict, total=False):
type: Literal["default"]
ttl: str

View file

@ -0,0 +1,160 @@
"""Typed request bodies for the Bedrock Invoke sub-providers.
Each sub-provider declares the exact wire keys it accepts in its inference
params. Parsing `optional_params` into one of these models splits the payload
into known fields (kept) and unknown ones captured as `extra_body` passthrough.
The passthrough is forwarded by default and stripped when `drop_params` is set,
so a strict caller never ships keys the provider would reject.
"""
from __future__ import annotations
import json
from typing import Dict, List, Optional, Type
from pydantic import BaseModel, ConfigDict
from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS
BEDROCK_INVOKE_PROVIDER = str
class BedrockInvokeInferenceParams(BaseModel):
"""Base for a sub-provider's inference-param body.
Declared fields are the provider's wire keys; anything else is captured as
passthrough (`model_extra`) and dropped only when `drop_params` is set.
"""
model_config = ConfigDict(extra="allow")
def to_body(self, *, drop_params: bool) -> Dict[str, object]:
known = self.model_dump(
exclude_none=True, exclude=set((self.model_extra or {}).keys())
)
if drop_params:
return known
return {**known, **(self.model_extra or {})}
class _CohereCommandRBody(BedrockInvokeInferenceParams):
max_tokens: Optional[int] = None
stream: Optional[bool] = None
temperature: Optional[float] = None
p: Optional[float] = None
k: Optional[float] = None
seed: Optional[int] = None
frequency_penalty: Optional[float] = None
presence_penalty: Optional[float] = None
stop_sequences: Optional[List[str]] = None
preamble: Optional[str] = None
prompt_truncation: Optional[str] = None
return_prompt: Optional[bool] = None
raw_prompting: Optional[bool] = None
search_queries_only: Optional[bool] = None
class _CohereLegacyBody(BedrockInvokeInferenceParams):
max_tokens: Optional[int] = None
stream: Optional[bool] = None
temperature: Optional[float] = None
p: Optional[float] = None
seed: Optional[int] = None
frequency_penalty: Optional[float] = None
presence_penalty: Optional[float] = None
stop_sequences: Optional[List[str]] = None
num_generations: Optional[int] = None
return_likelihood: Optional[str] = None
class _AI21Body(BedrockInvokeInferenceParams):
maxTokens: Optional[int] = None
temperature: Optional[float] = None
topP: Optional[float] = None
stream: Optional[bool] = None
stopSequences: Optional[List[str]] = None
class _MistralBody(BedrockInvokeInferenceParams):
max_tokens: Optional[int] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
top_k: Optional[float] = None
stream: Optional[bool] = None
stop: Optional[List[str]] = None
class _TitanTextGenerationConfig(BedrockInvokeInferenceParams):
maxTokenCount: Optional[int] = None
temperature: Optional[float] = None
topP: Optional[float] = None
stopSequences: Optional[List[str]] = None
class _LlamaBody(BedrockInvokeInferenceParams):
max_gen_len: Optional[int] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
topP: Optional[float] = None
stream: Optional[bool] = None
_COHERE_COMMAND_R_PREFIX = "cohere.command-r"
_INVOKE_BODY_MODELS: Dict[str, Type[BedrockInvokeInferenceParams]] = {
"cohere_command_r": _CohereCommandRBody,
"cohere": _CohereLegacyBody,
"ai21": _AI21Body,
"mistral": _MistralBody,
"amazon": _TitanTextGenerationConfig,
"meta": _LlamaBody,
"llama": _LlamaBody,
"deepseek_r1": _LlamaBody,
}
def _resolve_invoke_body_model(
provider: Optional[str], model: str
) -> Optional[Type[BedrockInvokeInferenceParams]]:
if provider == "cohere" and model.startswith(_COHERE_COMMAND_R_PREFIX):
return _CohereCommandRBody
if provider is None:
return None
return _INVOKE_BODY_MODELS.get(provider)
def parse_invoke_inference_params(
provider: Optional[str],
model: str,
params: Dict[str, object],
drop_params: bool,
) -> Dict[str, object]:
"""Validate inference params against the sub-provider's typed body.
Providers without a typed body (e.g. delegating ones) pass through
unchanged, preserving today's behavior.
"""
body_model = _resolve_invoke_body_model(provider, model)
if body_model is None:
return params
return body_model.model_validate(params).to_body(drop_params=drop_params)
def assert_no_control_params(body: Dict[str, object]) -> None:
"""Guard run right before dispatch: control keys must never reach the wire."""
leaked = LITELLM_CONTROL_PARAM_KEYS & body.keys()
if leaked:
raise ValueError(
"litellm control params leaked into the Bedrock request body: "
f"{sorted(leaked)}"
)
def assert_no_control_params_in_payload(data: str) -> None:
"""Dispatch-time guard over the serialized request payload."""
try:
parsed = json.loads(data)
except (TypeError, ValueError):
return
if isinstance(parsed, dict):
assert_no_control_params(parsed)

View file

@ -3559,89 +3559,6 @@ def test_bedrock_openai_model_id_extraction():
print(f"✓ Model ID extracted and encoded: {model_id}")
def test_bedrock_openai_convert_messages_to_prompt():
"""
Test that convert_messages_to_prompt returns empty string for OpenAI models.
"""
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
bedrock_llm = BedrockLLM()
messages = [
{"role": "system", "content": "You are helpful"},
{"role": "user", "content": "Hello"},
]
prompt, chat_history = bedrock_llm.convert_messages_to_prompt(
model="test-model", messages=messages, provider="openai", custom_prompt_dict={}
)
# OpenAI models use messages directly, no prompt conversion
assert prompt == ""
assert chat_history is None
print("✓ convert_messages_to_prompt returns empty for OpenAI")
def test_bedrock_openai_response_parsing():
"""
Test that OpenAI responses are correctly parsed.
"""
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm import ModelResponse
from unittest.mock import Mock
import json
bedrock_llm = BedrockLLM()
# Mock OpenAI-style response
openai_response = {
"choices": [
{
"message": {
"content": "The capital of France is Paris.",
"role": "assistant",
},
"finish_reason": "stop",
"index": 0,
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
}
mock_response = Mock()
mock_response.json.return_value = openai_response
mock_response.text = json.dumps(openai_response)
mock_response.status_code = 200
mock_response.headers = {}
model_response = ModelResponse()
mock_logging = Mock()
result = bedrock_llm.process_response(
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
response=mock_response,
model_response=model_response,
stream=False,
logging_obj=mock_logging,
optional_params={},
api_key="",
data={},
messages=[{"role": "user", "content": "What is the capital of France?"}],
print_verbose=lambda x: None,
encoding=None,
)
# Verify response content
assert result.choices[0].message.content == "The capital of France is Paris."
assert result.choices[0].finish_reason == "stop"
# Verify usage
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 8
assert result.usage.total_tokens == 18
print("✓ OpenAI response parsing works correctly")
def test_bedrock_openai_request_transformation():
"""
Test that the request is correctly transformed for OpenAI models.
@ -3831,46 +3748,6 @@ def test_bedrock_openai_multiple_message_types():
print("✓ Multiple message types handled correctly")
def test_bedrock_openai_error_handling():
"""
Test that errors from OpenAI models are properly handled.
"""
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm import ModelResponse
from litellm.llms.bedrock.common_utils import BedrockError
from unittest.mock import Mock
import json
bedrock_llm = BedrockLLM()
# Mock error response
mock_response = Mock()
mock_response.json.side_effect = Exception("Invalid JSON")
mock_response.text = "Invalid response"
mock_response.status_code = 422
model_response = ModelResponse()
mock_logging = Mock()
with pytest.raises(BedrockError) as exc_info:
bedrock_llm.process_response(
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
response=mock_response,
model_response=model_response,
stream=False,
logging_obj=mock_logging,
optional_params={},
api_key="",
data={},
messages=[],
print_verbose=lambda x: None,
encoding=None,
)
assert exc_info.value.status_code == 422
print("✓ Error handling works correctly")
# ============================================================================
# Nova Grounding (web_search_options) Unit Tests (Mocked)
# ============================================================================

View file

@ -1,4 +1,3 @@
import json
import os
import sys
@ -25,12 +24,15 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
(AmazonAnthropicClaudeConfig, "anthropic.claude-sonnet-4-6"),
],
)
def test_transform_request_drops_stream_chunk_size(config, model):
def test_signed_invoke_body_drops_stream_chunk_size(config, model):
"""stream_chunk_size is a LiteLLM-internal knob for re-chunking the HTTP
response stream. Leaking it into the provider request body makes Bedrock
reject the whole request: ValidationException 'stream_chunk_size: Extra
inputs are not permitted'."""
request_body = config().transform_request(
inputs are not permitted'. The two invoke transform entry points build the
body differently, so this asserts on the actual signed wire bytes that both
funnel through, regardless of which transform produced them."""
cfg = config()
request_body = cfg.transform_request(
model=model,
messages=[{"role": "user", "content": "hi"}],
optional_params={"stream": True, "stream_chunk_size": 2048, "max_tokens": 10},
@ -38,4 +40,50 @@ def test_transform_request_drops_stream_chunk_size(config, model):
headers={},
)
assert "stream_chunk_size" not in json.dumps(request_body)
_, signed_body = cfg.sign_request(
headers={},
optional_params={},
request_data=request_body,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/{}/invoke".format(
model
),
api_key="test-bearer-token",
model=model,
stream=True,
)
assert signed_body is not None
assert "stream_chunk_size" not in signed_body.decode()
assert "max_tokens" in signed_body.decode()
def test_extra_body_passthrough_by_default():
"""Unknown body keys are forwarded verbatim when drop_params is off, so the
soft-allowlist escape hatch keeps working."""
cfg = AmazonInvokeConfig()
request_body = cfg.transform_request(
model="mistral.mistral-7b-instruct-v0:2",
messages=[{"role": "user", "content": "hi"}],
optional_params={"temperature": 0.5, "made_up_param": "x"},
litellm_params={},
headers={},
)
assert request_body["temperature"] == 0.5
assert request_body["made_up_param"] == "x"
def test_drop_params_strips_extra_body_but_keeps_known_params():
"""drop_params selects the strict body: typed provider keys survive, unknown
passthrough keys are dropped instead of being shipped to the provider."""
cfg = AmazonInvokeConfig()
request_body = cfg.transform_request(
model="mistral.mistral-7b-instruct-v0:2",
messages=[{"role": "user", "content": "hi"}],
optional_params={"temperature": 0.5, "made_up_param": "x"},
litellm_params={"drop_params": True},
headers={},
)
assert request_body["temperature"] == 0.5
assert "made_up_param" not in request_body

View file

@ -8,14 +8,11 @@ sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
BedrockLLM,
make_call,
make_sync_call,
)
from litellm.llms.custom_httpx.http_handler import HTTPHandler
def test_transform_thinking_blocks_with_redacted_content():
@ -297,32 +294,42 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.iter_bytes = MagicMock(return_value=iter([]))
client = HTTPHandler()
client.post = MagicMock(return_value=mock_response)
@pytest.mark.asyncio
async def test_make_call_guards_against_leaked_control_param():
"""The dispatch wrapper validates the payload right before sending, so a
control param that slipped into the body aborts the request instead of
being rejected downstream by Bedrock."""
client = MagicMock()
client.post = AsyncMock()
BedrockLLM().completion(
model="cohere.command-text-v14",
messages=[{"role": "user", "content": "hi"}],
api_base=None,
custom_prompt_dict={},
model_response=litellm.ModelResponse(),
print_verbose=lambda *args, **kwargs: None,
encoding=litellm.encoding,
logging_obj=MagicMock(),
optional_params={
"stream": True,
"aws_access_key_id": "fake",
"aws_secret_access_key": "fake",
"aws_region_name": "us-east-1",
},
acompletion=False,
timeout=None,
litellm_params={},
client=client,
)
with pytest.raises(ValueError, match="stream_chunk_size"):
await make_call(
client=client,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/cohere.command-text-v14/invoke-with-response-stream",
headers={},
data='{"prompt": "hi", "stream_chunk_size": 2048}',
model="cohere.command-text-v14",
messages=[],
logging_obj=MagicMock(),
)
mock_response.iter_bytes.assert_called_once_with(chunk_size=None)
client.post.assert_not_called()
def test_make_sync_call_guards_against_leaked_control_param():
client = MagicMock()
client.post = MagicMock()
with pytest.raises(ValueError, match="stream_chunk_size"):
make_sync_call(
client=client,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/cohere.command-text-v14/invoke-with-response-stream",
headers={},
data='{"prompt": "hi", "stream_chunk_size": 2048}',
signed_json_body=None,
model="cohere.command-text-v14",
messages=[],
logging_obj=MagicMock(),
)
client.post.assert_not_called()

View file

@ -202,6 +202,27 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
def test_make_sync_call_guards_against_leaked_control_param():
"""The converse dispatch wrapper validates the payload right before sending,
so a control param that slipped into the body aborts the request instead of
being rejected downstream by Bedrock."""
client = MagicMock()
client.post = MagicMock()
with pytest.raises(ValueError, match="stream_chunk_size"):
make_sync_call(
client=client,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-sonnet-4-6/converse-stream",
headers={},
data='{"messages": [], "stream_chunk_size": 2048}',
model="anthropic.claude-sonnet-4-6",
messages=[],
logging_obj=MagicMock(),
)
client.post.assert_not_called()
def test_completion_plumbs_stream_chunk_size_through_converse():
iter_bytes_spy = _stream_completion_with_spied_iter_bytes(
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0"

View file

@ -0,0 +1,91 @@
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.types.llms.bedrock_invoke import (
assert_no_control_params,
assert_no_control_params_in_payload,
parse_invoke_inference_params,
)
def test_passthrough_keeps_unknown_keys():
body = parse_invoke_inference_params(
provider="mistral",
model="mistral.mistral-7b-instruct-v0:2",
params={"temperature": 0.5, "unknown_key": "x"},
drop_params=False,
)
assert body == {"temperature": 0.5, "unknown_key": "x"}
def test_drop_params_strips_only_unknown_keys():
body = parse_invoke_inference_params(
provider="mistral",
model="mistral.mistral-7b-instruct-v0:2",
params={"temperature": 0.5, "unknown_key": "x"},
drop_params=True,
)
assert body == {"temperature": 0.5}
def test_command_r_and_legacy_resolve_to_different_bodies():
"""command-r exposes k/p; the legacy text model does not, so the same key
must be classified differently per model id."""
command_r = parse_invoke_inference_params(
provider="cohere",
model="cohere.command-r-v1:0",
params={"k": 2},
drop_params=True,
)
legacy = parse_invoke_inference_params(
provider="cohere",
model="cohere.command-text-v14",
params={"k": 2},
drop_params=True,
)
assert command_r == {"k": 2.0}
assert legacy == {}
def test_none_provider_passes_through_untouched():
params = {"anything": 1}
assert (
parse_invoke_inference_params(
provider=None,
model="some-unresolved-model",
params=params,
drop_params=True,
)
== params
)
def test_unmodeled_provider_passes_through_untouched():
params = {"anything": 1, "stream_chunk_size": 4}
assert (
parse_invoke_inference_params(
provider="anthropic",
model="anthropic.claude-sonnet-4-6",
params=params,
drop_params=True,
)
== params
)
def test_guard_raises_on_leaked_control_param():
with pytest.raises(ValueError, match="stream_chunk_size"):
assert_no_control_params({"temperature": 0.5, "stream_chunk_size": 2048})
def test_guard_payload_ignores_non_dict_and_invalid_json():
assert_no_control_params_in_payload("not json")
assert_no_control_params_in_payload("[1, 2, 3]")
assert_no_control_params_in_payload('{"temperature": 0.5}')
with pytest.raises(ValueError, match="stream_chunk_size"):
assert_no_control_params_in_payload('{"stream_chunk_size": 2048}')