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 b00f7ba563b..49e0f56fad2 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -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, ) diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 2ca2469fb0f..f3c26c982e8 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -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]