diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 6bb2da1ad44..b00f7ba563b 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -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, diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index fa8c3a93ef3..2ca2469fb0f 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -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] diff --git a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py index aff89f02ff2..e11701e06b1 100644 --- a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py @@ -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