diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py index 09e0ec7a44a..e84989e4c1a 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -24,7 +24,6 @@ from litellm.llms.bedrock.common_utils import ( remove_custom_field_from_tools, ) from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER -from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponse from litellm.utils import _supports_factory @@ -191,12 +190,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): litellm_params: dict, headers: dict, ) -> dict: - filtered_params = { - k: v - for k, v in optional_params.items() - if k not in self.aws_authentication_params - and k not in LITELLM_CONTROL_PARAM_KEYS - } + filtered_params = self.filter_invoke_request_params(optional_params) output_config = filtered_params.get("output_config") if isinstance(output_config, dict): filtered_params["output_config"] = dict(output_config) 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 36be6818ab8..1c95dd0449d 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -141,6 +141,14 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): if k not in inference_params: inference_params[k] = v + def filter_invoke_request_params(self, optional_params: dict) -> dict: + return { + k: v + for k, v in optional_params.items() + if k not in self.aws_authentication_params + and k not in LITELLM_CONTROL_PARAM_KEYS + } + def transform_request( self, model: str, @@ -162,13 +170,9 @@ 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 - and k not in LITELLM_CONTROL_PARAM_KEYS - } + inference_params = self.filter_invoke_request_params( + copy.deepcopy(optional_params) + ) request_data: dict = {} if provider == "cohere": if model.startswith("cohere.command-r"):