refactor(bedrock): type invoke request bodies, split extra_body on drop_params, guard control-param leaks at dispatch

Parse each invoke sub-provider's inference params into a typed Pydantic
body so known wire keys are kept and unknown keys ride as extra_body
passthrough, dropped only when drop_params is set. Add a dispatch-time
guard in make_call/make_sync_call (invoke and converse) that aborts the
request if a litellm control param such as stream_chunk_size leaked into
the body, and point the converse path at the same shared control-key
source of truth used by invoke.
This commit is contained in:
mateo-berri 2026-06-21 03:00:52 +00:00
parent 86fd1358e5
commit f4b56ae89a
No known key found for this signature in database
7 changed files with 333 additions and 5 deletions

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)

View file

@ -45,6 +45,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
)
from litellm.types.llms.bedrock import *
from litellm.types.llms.bedrock_invoke import assert_no_control_params_in_payload
from litellm.types.llms.openai import (
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
@ -199,6 +200,7 @@ async def make_call(
bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None,
stream_chunk_size: Optional[int] = None,
):
assert_no_control_params_in_payload(data)
try:
if client is None:
client = get_async_httpx_client(
@ -296,6 +298,7 @@ def make_sync_call(
bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None,
stream_chunk_size: Optional[int] = None,
):
assert_no_control_params_in_payload(data)
try:
if client is None:
client = _get_httpx_client(

View file

@ -25,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import (
_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
@ -170,8 +171,12 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
provider=provider,
custom_prompt_dict=custom_prompt_dict,
)
inference_params = self.filter_invoke_request_params(
copy.deepcopy(optional_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":

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

@ -3,7 +3,9 @@ import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../../../..")) # Adds the parent directory to the system path
sys.path.insert(
0, os.path.abspath("../../../../../..")
) # Adds the parent directory to the system path
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeConfig,
@ -42,7 +44,9 @@ def test_signed_invoke_body_drops_stream_chunk_size(config, model):
headers={},
optional_params={},
request_data=request_body,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/{}/invoke".format(model),
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/{}/invoke".format(
model
),
api_key="test-bearer-token",
model=model,
stream=True,
@ -51,3 +55,35 @@ def test_signed_invoke_body_drops_stream_chunk_size(config, model):
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

@ -297,6 +297,47 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
@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()
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(),
)
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()
def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default():
mock_response = MagicMock()
mock_response.status_code = 200

View file

@ -0,0 +1,78 @@
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_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}')