fix(bedrock): type invoke request params

Use provider-specific typed request builders for Bedrock invoke bodies so LiteLLM internal params cannot be splatted into provider payloads.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
mateo-berri 2026-06-13 21:00:21 -07:00
parent 2655d1dd5e
commit ed082b1ef2
3 changed files with 375 additions and 42 deletions

View file

@ -2,7 +2,18 @@ import copy
import json
import time
from functools import partial
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast, get_args
from typing import (
TYPE_CHECKING,
Any,
Dict,
FrozenSet,
List,
Optional,
Tuple,
Union,
cast,
get_args,
)
import httpx
@ -24,6 +35,19 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
_get_httpx_client,
)
from litellm.types.llms.bedrock import (
BedrockInvokeAI21InferenceParams,
BedrockInvokeAI21Request,
BedrockInvokeCohereChatRequest,
BedrockInvokeCohereCompletionRequest,
BedrockInvokeCohereInferenceParams,
BedrockInvokeLlamaInferenceParams,
BedrockInvokeLlamaRequest,
BedrockInvokeMistralInferenceParams,
BedrockInvokeMistralRequest,
BedrockInvokeTitanInferenceParams,
BedrockInvokeTitanRequest,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, Usage
from litellm.utils import CustomStreamWrapper
@ -39,6 +63,22 @@ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
BEDROCK_INVOKE_COHERE_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
BedrockInvokeCohereInferenceParams.__annotations__.keys()
)
BEDROCK_INVOKE_AI21_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
BedrockInvokeAI21InferenceParams.__annotations__.keys()
)
BEDROCK_INVOKE_MISTRAL_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
BedrockInvokeMistralInferenceParams.__annotations__.keys()
)
BEDROCK_INVOKE_TITAN_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
BedrockInvokeTitanInferenceParams.__annotations__.keys()
)
BEDROCK_INVOKE_LLAMA_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
BedrockInvokeLlamaInferenceParams.__annotations__.keys()
)
def __init__(self, **kwargs):
BaseConfig.__init__(self, **kwargs)
BaseAWSLLM.__init__(self, **kwargs)
@ -134,11 +174,144 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
fake_stream=fake_stream,
)
def _apply_config_to_params(self, config: dict, inference_params: dict) -> None:
"""Apply config values to inference_params if not already set."""
for k, v in config.items():
if k not in inference_params:
inference_params[k] = v
def _get_filtered_inference_params(
self,
optional_params: dict,
config: dict,
allowed_fields: FrozenSet[str],
) -> Dict[str, object]:
copied_params = copy.deepcopy(optional_params)
inference_params = {
k: v
for k, v in copied_params.items()
if k in allowed_fields and k not in self.aws_authentication_params
}
config_params = {
k: v
for k, v in config.items()
if k in allowed_fields
and k not in self.aws_authentication_params
and k not in inference_params
}
return {**config_params, **inference_params}
def _get_cohere_inference_params(
self, optional_params: dict, config: dict
) -> BedrockInvokeCohereInferenceParams:
return cast(
BedrockInvokeCohereInferenceParams,
self._get_filtered_inference_params(
optional_params=optional_params,
config=config,
allowed_fields=self.BEDROCK_INVOKE_COHERE_ALLOWED_INFERENCE_FIELDS,
),
)
def _get_ai21_inference_params(
self, optional_params: dict, config: dict
) -> BedrockInvokeAI21InferenceParams:
return cast(
BedrockInvokeAI21InferenceParams,
self._get_filtered_inference_params(
optional_params=optional_params,
config=config,
allowed_fields=self.BEDROCK_INVOKE_AI21_ALLOWED_INFERENCE_FIELDS,
),
)
def _get_mistral_inference_params(
self, optional_params: dict, config: dict
) -> BedrockInvokeMistralInferenceParams:
return cast(
BedrockInvokeMistralInferenceParams,
self._get_filtered_inference_params(
optional_params=optional_params,
config=config,
allowed_fields=self.BEDROCK_INVOKE_MISTRAL_ALLOWED_INFERENCE_FIELDS,
),
)
def _get_titan_inference_params(
self, optional_params: dict, config: dict
) -> BedrockInvokeTitanInferenceParams:
return cast(
BedrockInvokeTitanInferenceParams,
self._get_filtered_inference_params(
optional_params=optional_params,
config=config,
allowed_fields=self.BEDROCK_INVOKE_TITAN_ALLOWED_INFERENCE_FIELDS,
),
)
def _get_llama_inference_params(
self, optional_params: dict, config: dict
) -> BedrockInvokeLlamaInferenceParams:
return cast(
BedrockInvokeLlamaInferenceParams,
self._get_filtered_inference_params(
optional_params=optional_params,
config=config,
allowed_fields=self.BEDROCK_INVOKE_LLAMA_ALLOWED_INFERENCE_FIELDS,
),
)
@staticmethod
def _build_cohere_chat_request(
prompt: str,
inference_params: BedrockInvokeCohereInferenceParams,
chat_history: Optional[List[Dict[str, object]]],
) -> BedrockInvokeCohereChatRequest:
request: Dict[str, object] = {"message": prompt, **inference_params}
if chat_history is not None:
request["chat_history"] = chat_history
return cast(BedrockInvokeCohereChatRequest, request)
@staticmethod
def _build_cohere_completion_request(
prompt: str,
inference_params: BedrockInvokeCohereInferenceParams,
) -> BedrockInvokeCohereCompletionRequest:
return cast(
BedrockInvokeCohereCompletionRequest,
{"prompt": prompt, **inference_params},
)
@staticmethod
def _build_ai21_request(
prompt: str,
inference_params: BedrockInvokeAI21InferenceParams,
) -> BedrockInvokeAI21Request:
return cast(BedrockInvokeAI21Request, {"prompt": prompt, **inference_params})
@staticmethod
def _build_mistral_request(
prompt: str,
inference_params: BedrockInvokeMistralInferenceParams,
) -> BedrockInvokeMistralRequest:
return cast(
BedrockInvokeMistralRequest,
{"prompt": prompt, **inference_params},
)
@staticmethod
def _build_titan_request(
prompt: str,
inference_params: BedrockInvokeTitanInferenceParams,
) -> BedrockInvokeTitanRequest:
return BedrockInvokeTitanRequest(
inputText=prompt,
textGenerationConfig=inference_params,
)
@staticmethod
def _build_llama_request(
prompt: str,
inference_params: BedrockInvokeLlamaInferenceParams,
) -> BedrockInvokeLlamaRequest:
return cast(
BedrockInvokeLlamaRequest,
{"prompt": prompt, **inference_params},
)
def transform_request(
self,
@ -162,31 +335,36 @@ 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
}
request_data: dict = {}
if provider == "cohere":
if model.startswith("cohere.command-r"):
## LOAD CONFIG
config = litellm.AmazonCohereChatConfig().get_config()
self._apply_config_to_params(config, inference_params)
_data = {"message": prompt, **inference_params}
if chat_history is not None:
_data["chat_history"] = chat_history
request_data = _data
cohere_inference_params = self._get_cohere_inference_params(
optional_params=optional_params,
config=config,
)
return cast(
dict,
self._build_cohere_chat_request(
prompt=prompt,
inference_params=cohere_inference_params,
chat_history=chat_history,
),
)
else:
## LOAD CONFIG
config = litellm.AmazonCohereConfig.get_config()
self._apply_config_to_params(config, inference_params)
cohere_inference_params = self._get_cohere_inference_params(
optional_params=optional_params,
config=config,
)
if stream is True:
inference_params["stream"] = (
True # cohere requires stream = True in inference params
)
request_data = {"prompt": prompt, **inference_params}
cohere_inference_params["stream"] = True
return cast(
dict,
self._build_cohere_completion_request(
prompt=prompt,
inference_params=cohere_inference_params,
),
)
elif provider == "anthropic":
transformed_request = (
litellm.AmazonAnthropicClaudeConfig().transform_request(
@ -208,28 +386,57 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
headers=headers,
)
elif provider == "ai21":
## LOAD CONFIG
config = litellm.AmazonAI21Config.get_config()
self._apply_config_to_params(config, inference_params)
request_data = {"prompt": prompt, **inference_params}
ai21_inference_params = self._get_ai21_inference_params(
optional_params=optional_params,
config=config,
)
return cast(
dict,
self._build_ai21_request(
prompt=prompt,
inference_params=ai21_inference_params,
),
)
elif provider == "mistral":
## LOAD CONFIG
config = litellm.AmazonMistralConfig.get_config()
self._apply_config_to_params(config, inference_params)
request_data = {"prompt": prompt, **inference_params}
mistral_inference_params = self._get_mistral_inference_params(
optional_params=optional_params,
config=config,
)
return cast(
dict,
self._build_mistral_request(
prompt=prompt,
inference_params=mistral_inference_params,
),
)
elif provider == "amazon": # amazon titan
## LOAD CONFIG
config = litellm.AmazonTitanConfig.get_config()
self._apply_config_to_params(config, inference_params)
request_data = {
"inputText": prompt,
"textGenerationConfig": inference_params,
}
titan_inference_params = self._get_titan_inference_params(
optional_params=optional_params,
config=config,
)
return cast(
dict,
self._build_titan_request(
prompt=prompt,
inference_params=titan_inference_params,
),
)
elif provider == "meta" or provider == "llama" or provider == "deepseek_r1":
## LOAD CONFIG
config = litellm.AmazonLlamaConfig.get_config()
self._apply_config_to_params(config, inference_params)
request_data = {"prompt": prompt, **inference_params}
llama_inference_params = self._get_llama_inference_params(
optional_params=optional_params,
config=config,
)
return cast(
dict,
self._build_llama_request(
prompt=prompt,
inference_params=llama_inference_params,
),
)
elif provider == "twelvelabs":
return litellm.AmazonTwelveLabsPegasusConfig().transform_request(
model=model,
@ -255,8 +462,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
),
)
return request_data
def transform_response( # noqa: PLR0915
self,
model: str,

View file

@ -432,6 +432,85 @@ class BedrockInvokeNovaRequest(TypedDict, total=False):
guardrailConfig: Optional[GuardrailConfigBlock]
class BedrockInvokeCohereInferenceParams(TypedDict, total=False):
max_tokens: int
temperature: float
return_likelihood: str
p: float
k: int
stop_sequences: List[str]
num_generations: int
frequency_penalty: float
presence_penalty: float
truncate: str
stream: bool
tools: List[Dict[str, object]]
tool_results: List[Dict[str, object]]
seed: int
force_single_step: bool
class BedrockInvokeCohereCompletionRequest(
BedrockInvokeCohereInferenceParams, total=False
):
prompt: Required[str]
class BedrockInvokeCohereChatRequest(BedrockInvokeCohereInferenceParams, total=False):
message: Required[str]
chat_history: List[Dict[str, object]]
class BedrockInvokeAI21InferenceParams(TypedDict, total=False):
maxTokens: int
temperature: float
topP: float
stopSequences: List[str]
frequencyPenalty: Dict[str, object]
frequencePenalty: Dict[str, object]
presencePenalty: Dict[str, object]
countPenalty: Dict[str, object]
class BedrockInvokeAI21Request(BedrockInvokeAI21InferenceParams, total=False):
prompt: Required[str]
class BedrockInvokeMistralInferenceParams(TypedDict, total=False):
max_tokens: int
temperature: float
top_p: float
top_k: float
stop: List[str]
class BedrockInvokeMistralRequest(BedrockInvokeMistralInferenceParams, total=False):
prompt: Required[str]
class BedrockInvokeTitanInferenceParams(TypedDict, total=False):
maxTokenCount: int
stopSequences: List[str]
temperature: float
topP: int
class BedrockInvokeTitanRequest(TypedDict):
inputText: str
textGenerationConfig: BedrockInvokeTitanInferenceParams
class BedrockInvokeLlamaInferenceParams(TypedDict, total=False):
max_gen_len: int
temperature: float
top_p: float
topP: float
class BedrockInvokeLlamaRequest(BedrockInvokeLlamaInferenceParams, total=False):
prompt: Required[str]
class GenericStreamingChunk(TypedDict):
text: Required[str]
tool_use: Optional[ChatCompletionToolCallChunk]

View file

@ -39,3 +39,52 @@ def test_transform_request_drops_stream_chunk_size(config, model):
)
assert "stream_chunk_size" not in json.dumps(request_body)
@pytest.mark.parametrize(
"model,valid_param,valid_value,valid_param_path",
[
("cohere.command-text-v14", "max_tokens", 10, ("max_tokens",)),
("ai21.j2-ultra-v1", "maxTokens", 10, ("maxTokens",)),
("mistral.mistral-7b-instruct-v0:2", "max_tokens", 10, ("max_tokens",)),
(
"amazon.titan-text-express-v1",
"maxTokenCount",
10,
("textGenerationConfig", "maxTokenCount"),
),
("meta.llama2-13b-chat-v1", "max_gen_len", 10, ("max_gen_len",)),
],
)
def test_transform_request_drops_internal_params_from_typed_invoke_body(
model, valid_param, valid_value, valid_param_path
):
request_body = AmazonInvokeConfig().transform_request(
model=model,
messages=[{"role": "user", "content": "hi"}],
optional_params={
"skip_mcp_handler": True,
"_skip_mcp_handler": True,
"mcp_handler_context": {"request_id": "test"},
"stream_chunk_size": 2048,
"aws_region_name": "us-east-1",
valid_param: valid_value,
},
litellm_params={},
headers={},
)
serialized_request = json.dumps(request_body)
for internal_param in (
"skip_mcp_handler",
"_skip_mcp_handler",
"mcp_handler_context",
"stream_chunk_size",
"aws_region_name",
):
assert internal_param not in serialized_request
value = request_body
for path_part in valid_param_path:
value = value[path_part]
assert value == valid_value