refactor(bedrock): freeze invoke params before request build

Parse Bedrock invoke params into provider-specific frozen dataclasses before constructing wire request bodies so builders cannot accept raw optional_params.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
mateo-berri 2026-06-13 21:10:10 -07:00
parent ed082b1ef2
commit 73a4dcb9d9
2 changed files with 343 additions and 122 deletions

View file

@ -1,12 +1,12 @@
import copy
import json
import time
from dataclasses import dataclass, replace
from functools import partial
from typing import (
TYPE_CHECKING,
Any,
Dict,
FrozenSet,
List,
Optional,
Tuple,
@ -36,14 +36,10 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
)
from litellm.types.llms.bedrock import (
BedrockInvokeAI21InferenceParams,
BedrockInvokeAI21Request,
BedrockInvokeCohereChatRequest,
BedrockInvokeCohereCompletionRequest,
BedrockInvokeCohereInferenceParams,
BedrockInvokeLlamaInferenceParams,
BedrockInvokeLlamaRequest,
BedrockInvokeMistralInferenceParams,
BedrockInvokeMistralRequest,
BedrockInvokeTitanInferenceParams,
BedrockInvokeTitanRequest,
@ -62,23 +58,63 @@ else:
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()
)
@dataclass(frozen=True, slots=True)
class BedrockInvokeCohereParams:
max_tokens: Optional[int] = None
temperature: Optional[float] = None
return_likelihood: Optional[str] = None
p: Optional[float] = None
k: Optional[int] = None
stop_sequences: Optional[List[str]] = None
num_generations: Optional[int] = None
frequency_penalty: Optional[float] = None
presence_penalty: Optional[float] = None
truncate: Optional[str] = None
stream: Optional[bool] = None
tools: Optional[List[Dict[str, object]]] = None
tool_results: Optional[List[Dict[str, object]]] = None
seed: Optional[int] = None
force_single_step: Optional[bool] = None
@dataclass(frozen=True, slots=True)
class BedrockInvokeAI21Params:
maxTokens: Optional[int] = None
temperature: Optional[float] = None
topP: Optional[float] = None
stopSequences: Optional[List[str]] = None
frequencyPenalty: Optional[Dict[str, object]] = None
frequencePenalty: Optional[Dict[str, object]] = None
presencePenalty: Optional[Dict[str, object]] = None
countPenalty: Optional[Dict[str, object]] = None
@dataclass(frozen=True, slots=True)
class BedrockInvokeMistralParams:
max_tokens: Optional[int] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
top_k: Optional[float] = None
stop: Optional[List[str]] = None
@dataclass(frozen=True, slots=True)
class BedrockInvokeTitanParams:
maxTokenCount: Optional[int] = None
stopSequences: Optional[List[str]] = None
temperature: Optional[float] = None
topP: Optional[int] = None
@dataclass(frozen=True, slots=True)
class BedrockInvokeLlamaParams:
max_gen_len: Optional[int] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
topP: Optional[float] = None
class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
def __init__(self, **kwargs):
BaseConfig.__init__(self, **kwargs)
BaseAWSLLM.__init__(self, **kwargs)
@ -174,94 +210,200 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
fake_stream=fake_stream,
)
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}
@staticmethod
def _get_param_value(optional_params: dict, config: dict, key: str) -> object:
if key in optional_params:
return copy.deepcopy(optional_params[key])
return copy.deepcopy(config.get(key))
def _get_cohere_inference_params(
@staticmethod
def _drop_none(data: Dict[str, object]) -> Dict[str, object]:
return {key: value for key, value in data.items() if value is not None}
def _parse_cohere_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,
) -> BedrockInvokeCohereParams:
return BedrockInvokeCohereParams(
max_tokens=cast(
Optional[int],
self._get_param_value(optional_params, config, "max_tokens"),
),
temperature=cast(
Optional[float],
self._get_param_value(optional_params, config, "temperature"),
),
return_likelihood=cast(
Optional[str],
self._get_param_value(optional_params, config, "return_likelihood"),
),
p=cast(
Optional[float],
self._get_param_value(optional_params, config, "p"),
),
k=cast(
Optional[int],
self._get_param_value(optional_params, config, "k"),
),
stop_sequences=cast(
Optional[List[str]],
self._get_param_value(optional_params, config, "stop_sequences"),
),
num_generations=cast(
Optional[int],
self._get_param_value(optional_params, config, "num_generations"),
),
frequency_penalty=cast(
Optional[float],
self._get_param_value(optional_params, config, "frequency_penalty"),
),
presence_penalty=cast(
Optional[float],
self._get_param_value(optional_params, config, "presence_penalty"),
),
truncate=cast(
Optional[str],
self._get_param_value(optional_params, config, "truncate"),
),
stream=cast(
Optional[bool],
self._get_param_value(optional_params, config, "stream"),
),
tools=cast(
Optional[List[Dict[str, object]]],
self._get_param_value(optional_params, config, "tools"),
),
tool_results=cast(
Optional[List[Dict[str, object]]],
self._get_param_value(optional_params, config, "tool_results"),
),
seed=cast(
Optional[int],
self._get_param_value(optional_params, config, "seed"),
),
force_single_step=cast(
Optional[bool],
self._get_param_value(optional_params, config, "force_single_step"),
),
)
def _get_ai21_inference_params(
def _parse_ai21_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,
) -> BedrockInvokeAI21Params:
return BedrockInvokeAI21Params(
maxTokens=cast(
Optional[int],
self._get_param_value(optional_params, config, "maxTokens"),
),
temperature=cast(
Optional[float],
self._get_param_value(optional_params, config, "temperature"),
),
topP=cast(
Optional[float],
self._get_param_value(optional_params, config, "topP"),
),
stopSequences=cast(
Optional[List[str]],
self._get_param_value(optional_params, config, "stopSequences"),
),
frequencyPenalty=cast(
Optional[Dict[str, object]],
self._get_param_value(optional_params, config, "frequencyPenalty"),
),
frequencePenalty=cast(
Optional[Dict[str, object]],
self._get_param_value(optional_params, config, "frequencePenalty"),
),
presencePenalty=cast(
Optional[Dict[str, object]],
self._get_param_value(optional_params, config, "presencePenalty"),
),
countPenalty=cast(
Optional[Dict[str, object]],
self._get_param_value(optional_params, config, "countPenalty"),
),
)
def _get_mistral_inference_params(
def _parse_mistral_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,
) -> BedrockInvokeMistralParams:
return BedrockInvokeMistralParams(
max_tokens=cast(
Optional[int],
self._get_param_value(optional_params, config, "max_tokens"),
),
temperature=cast(
Optional[float],
self._get_param_value(optional_params, config, "temperature"),
),
top_p=cast(
Optional[float],
self._get_param_value(optional_params, config, "top_p"),
),
top_k=cast(
Optional[float],
self._get_param_value(optional_params, config, "top_k"),
),
stop=cast(
Optional[List[str]],
self._get_param_value(optional_params, config, "stop"),
),
)
def _get_titan_inference_params(
def _parse_titan_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,
) -> BedrockInvokeTitanParams:
return BedrockInvokeTitanParams(
maxTokenCount=cast(
Optional[int],
self._get_param_value(optional_params, config, "maxTokenCount"),
),
stopSequences=cast(
Optional[List[str]],
self._get_param_value(optional_params, config, "stopSequences"),
),
temperature=cast(
Optional[float],
self._get_param_value(optional_params, config, "temperature"),
),
topP=cast(
Optional[int],
self._get_param_value(optional_params, config, "topP"),
),
)
def _get_llama_inference_params(
def _parse_llama_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,
) -> BedrockInvokeLlamaParams:
return BedrockInvokeLlamaParams(
max_gen_len=cast(
Optional[int],
self._get_param_value(optional_params, config, "max_gen_len"),
),
temperature=cast(
Optional[float],
self._get_param_value(optional_params, config, "temperature"),
),
top_p=cast(
Optional[float],
self._get_param_value(optional_params, config, "top_p"),
),
topP=cast(
Optional[float],
self._get_param_value(optional_params, config, "topP"),
),
)
@staticmethod
def _build_cohere_chat_request(
prompt: str,
inference_params: BedrockInvokeCohereInferenceParams,
inference_params: BedrockInvokeCohereParams,
chat_history: Optional[List[Dict[str, object]]],
) -> BedrockInvokeCohereChatRequest:
request: Dict[str, object] = {"message": prompt, **inference_params}
request: Dict[str, object] = {
"message": prompt,
**AmazonInvokeConfig._cohere_params_to_wire_dict(inference_params),
}
if chat_history is not None:
request["chat_history"] = chat_history
return cast(BedrockInvokeCohereChatRequest, request)
@ -269,48 +411,123 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
@staticmethod
def _build_cohere_completion_request(
prompt: str,
inference_params: BedrockInvokeCohereInferenceParams,
inference_params: BedrockInvokeCohereParams,
) -> BedrockInvokeCohereCompletionRequest:
return cast(
BedrockInvokeCohereCompletionRequest,
{"prompt": prompt, **inference_params},
{
"prompt": prompt,
**AmazonInvokeConfig._cohere_params_to_wire_dict(inference_params),
},
)
@staticmethod
def _build_ai21_request(
prompt: str,
inference_params: BedrockInvokeAI21InferenceParams,
inference_params: BedrockInvokeAI21Params,
) -> BedrockInvokeAI21Request:
return cast(BedrockInvokeAI21Request, {"prompt": prompt, **inference_params})
return cast(
BedrockInvokeAI21Request,
{
"prompt": prompt,
**AmazonInvokeConfig._drop_none(
{
"maxTokens": inference_params.maxTokens,
"temperature": inference_params.temperature,
"topP": inference_params.topP,
"stopSequences": inference_params.stopSequences,
"frequencyPenalty": inference_params.frequencyPenalty,
"frequencePenalty": inference_params.frequencePenalty,
"presencePenalty": inference_params.presencePenalty,
"countPenalty": inference_params.countPenalty,
}
),
},
)
@staticmethod
def _build_mistral_request(
prompt: str,
inference_params: BedrockInvokeMistralInferenceParams,
inference_params: BedrockInvokeMistralParams,
) -> BedrockInvokeMistralRequest:
return cast(
BedrockInvokeMistralRequest,
{"prompt": prompt, **inference_params},
{
"prompt": prompt,
**AmazonInvokeConfig._drop_none(
{
"max_tokens": inference_params.max_tokens,
"temperature": inference_params.temperature,
"top_p": inference_params.top_p,
"top_k": inference_params.top_k,
"stop": inference_params.stop,
}
),
},
)
@staticmethod
def _build_titan_request(
prompt: str,
inference_params: BedrockInvokeTitanInferenceParams,
inference_params: BedrockInvokeTitanParams,
) -> BedrockInvokeTitanRequest:
return BedrockInvokeTitanRequest(
inputText=prompt,
textGenerationConfig=inference_params,
textGenerationConfig=cast(
BedrockInvokeTitanInferenceParams,
AmazonInvokeConfig._drop_none(
{
"maxTokenCount": inference_params.maxTokenCount,
"stopSequences": inference_params.stopSequences,
"temperature": inference_params.temperature,
"topP": inference_params.topP,
}
),
),
)
@staticmethod
def _build_llama_request(
prompt: str,
inference_params: BedrockInvokeLlamaInferenceParams,
inference_params: BedrockInvokeLlamaParams,
) -> BedrockInvokeLlamaRequest:
return cast(
BedrockInvokeLlamaRequest,
{"prompt": prompt, **inference_params},
{
"prompt": prompt,
**AmazonInvokeConfig._drop_none(
{
"max_gen_len": inference_params.max_gen_len,
"temperature": inference_params.temperature,
"top_p": inference_params.top_p,
"topP": inference_params.topP,
}
),
},
)
@staticmethod
def _cohere_params_to_wire_dict(
inference_params: BedrockInvokeCohereParams,
) -> Dict[str, object]:
return AmazonInvokeConfig._drop_none(
{
"max_tokens": inference_params.max_tokens,
"temperature": inference_params.temperature,
"return_likelihood": inference_params.return_likelihood,
"p": inference_params.p,
"k": inference_params.k,
"stop_sequences": inference_params.stop_sequences,
"num_generations": inference_params.num_generations,
"frequency_penalty": inference_params.frequency_penalty,
"presence_penalty": inference_params.presence_penalty,
"truncate": inference_params.truncate,
"stream": inference_params.stream,
"tools": inference_params.tools,
"tool_results": inference_params.tool_results,
"seed": inference_params.seed,
"force_single_step": inference_params.force_single_step,
}
)
def transform_request(
@ -338,7 +555,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
if provider == "cohere":
if model.startswith("cohere.command-r"):
config = litellm.AmazonCohereChatConfig().get_config()
cohere_inference_params = self._get_cohere_inference_params(
cohere_inference_params = self._parse_cohere_params(
optional_params=optional_params,
config=config,
)
@ -352,12 +569,15 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
)
else:
config = litellm.AmazonCohereConfig.get_config()
cohere_inference_params = self._get_cohere_inference_params(
cohere_inference_params = self._parse_cohere_params(
optional_params=optional_params,
config=config,
)
if stream is True:
cohere_inference_params["stream"] = True
cohere_inference_params = replace(
cohere_inference_params,
stream=True,
)
return cast(
dict,
self._build_cohere_completion_request(
@ -387,7 +607,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
)
elif provider == "ai21":
config = litellm.AmazonAI21Config.get_config()
ai21_inference_params = self._get_ai21_inference_params(
ai21_inference_params = self._parse_ai21_params(
optional_params=optional_params,
config=config,
)
@ -400,7 +620,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
)
elif provider == "mistral":
config = litellm.AmazonMistralConfig.get_config()
mistral_inference_params = self._get_mistral_inference_params(
mistral_inference_params = self._parse_mistral_params(
optional_params=optional_params,
config=config,
)
@ -413,7 +633,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
)
elif provider == "amazon": # amazon titan
config = litellm.AmazonTitanConfig.get_config()
titan_inference_params = self._get_titan_inference_params(
titan_inference_params = self._parse_titan_params(
optional_params=optional_params,
config=config,
)
@ -426,7 +646,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
)
elif provider == "meta" or provider == "llama" or provider == "deepseek_r1":
config = litellm.AmazonLlamaConfig.get_config()
llama_inference_params = self._get_llama_inference_params(
llama_inference_params = self._parse_llama_params(
optional_params=optional_params,
config=config,
)

View file

@ -432,7 +432,8 @@ class BedrockInvokeNovaRequest(TypedDict, total=False):
guardrailConfig: Optional[GuardrailConfigBlock]
class BedrockInvokeCohereInferenceParams(TypedDict, total=False):
class BedrockInvokeCohereCompletionRequest(TypedDict, total=False):
prompt: Required[str]
max_tokens: int
temperature: float
return_likelihood: str
@ -450,18 +451,28 @@ class BedrockInvokeCohereInferenceParams(TypedDict, total=False):
force_single_step: bool
class BedrockInvokeCohereCompletionRequest(
BedrockInvokeCohereInferenceParams, total=False
):
prompt: Required[str]
class BedrockInvokeCohereChatRequest(BedrockInvokeCohereInferenceParams, total=False):
class BedrockInvokeCohereChatRequest(TypedDict, total=False):
message: Required[str]
chat_history: List[Dict[str, object]]
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 BedrockInvokeAI21InferenceParams(TypedDict, total=False):
class BedrockInvokeAI21Request(TypedDict, total=False):
prompt: Required[str]
maxTokens: int
temperature: float
topP: float
@ -472,11 +483,8 @@ class BedrockInvokeAI21InferenceParams(TypedDict, total=False):
countPenalty: Dict[str, object]
class BedrockInvokeAI21Request(BedrockInvokeAI21InferenceParams, total=False):
class BedrockInvokeMistralRequest(TypedDict, total=False):
prompt: Required[str]
class BedrockInvokeMistralInferenceParams(TypedDict, total=False):
max_tokens: int
temperature: float
top_p: float
@ -484,10 +492,6 @@ class BedrockInvokeMistralInferenceParams(TypedDict, total=False):
stop: List[str]
class BedrockInvokeMistralRequest(BedrockInvokeMistralInferenceParams, total=False):
prompt: Required[str]
class BedrockInvokeTitanInferenceParams(TypedDict, total=False):
maxTokenCount: int
stopSequences: List[str]
@ -500,17 +504,14 @@ class BedrockInvokeTitanRequest(TypedDict):
textGenerationConfig: BedrockInvokeTitanInferenceParams
class BedrockInvokeLlamaInferenceParams(TypedDict, total=False):
class BedrockInvokeLlamaRequest(TypedDict, total=False):
prompt: Required[str]
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]