diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 7b1064ccef9..04e82e5ebf0 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -14,6 +14,8 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, ) +from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS +from litellm.types.llms.bedrock_invoke import assert_no_control_params_in_payload from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -37,6 +39,7 @@ def make_sync_call( if client is None: client = _get_httpx_client() # Create a new client if none provided + assert_no_control_params_in_payload(data) response = client.post( api_base, headers=headers, @@ -268,7 +271,9 @@ class BedrockConverseLLM(BaseAWSLLM): ): ## SETUP ## stream = optional_params.pop("stream", None) - stream_chunk_size = optional_params.pop("stream_chunk_size", None) + stream_chunk_size = optional_params.get("stream_chunk_size") + for _control_key in LITELLM_CONTROL_PARAM_KEYS: + optional_params.pop(_control_key, None) unencoded_model_id = optional_params.pop("model_id", None) fake_stream = optional_params.pop("fake_stream", False) json_mode = optional_params.get("json_mode", False) diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 75b560b4d6d..753e41b5343 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -45,6 +45,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, ) from litellm.types.llms.bedrock import * +from litellm.types.llms.bedrock_invoke import assert_no_control_params_in_payload from litellm.types.llms.openai import ( ChatCompletionRedactedThinkingBlock, ChatCompletionThinkingBlock, @@ -199,6 +200,7 @@ async def make_call( bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None, stream_chunk_size: Optional[int] = None, ): + assert_no_control_params_in_payload(data) try: if client is None: client = get_async_httpx_client( @@ -296,6 +298,7 @@ def make_sync_call( bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None, stream_chunk_size: Optional[int] = None, ): + assert_no_control_params_in_payload(data) try: if client is None: client = _get_httpx_client( 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 1c95dd0449d..a7b4108a734 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -25,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, ) from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS +from litellm.types.llms.bedrock_invoke import parse_invoke_inference_params from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponse, Usage from litellm.utils import CustomStreamWrapper @@ -170,8 +171,12 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): provider=provider, custom_prompt_dict=custom_prompt_dict, ) - inference_params = self.filter_invoke_request_params( - copy.deepcopy(optional_params) + drop_params = bool(litellm_params.get("drop_params") or litellm.drop_params) + inference_params = parse_invoke_inference_params( + provider=provider, + model=model, + params=self.filter_invoke_request_params(copy.deepcopy(optional_params)), + drop_params=drop_params, ) request_data: dict = {} if provider == "cohere": diff --git a/litellm/types/llms/bedrock_invoke.py b/litellm/types/llms/bedrock_invoke.py new file mode 100644 index 00000000000..60b734777fe --- /dev/null +++ b/litellm/types/llms/bedrock_invoke.py @@ -0,0 +1,160 @@ +"""Typed request bodies for the Bedrock Invoke sub-providers. + +Each sub-provider declares the exact wire keys it accepts in its inference +params. Parsing `optional_params` into one of these models splits the payload +into known fields (kept) and unknown ones captured as `extra_body` passthrough. +The passthrough is forwarded by default and stripped when `drop_params` is set, +so a strict caller never ships keys the provider would reject. +""" + +from __future__ import annotations + +import json +from typing import Dict, List, Optional, Type + +from pydantic import BaseModel, ConfigDict + +from litellm.types.llms.bedrock import LITELLM_CONTROL_PARAM_KEYS + +BEDROCK_INVOKE_PROVIDER = str + + +class BedrockInvokeInferenceParams(BaseModel): + """Base for a sub-provider's inference-param body. + + Declared fields are the provider's wire keys; anything else is captured as + passthrough (`model_extra`) and dropped only when `drop_params` is set. + """ + + model_config = ConfigDict(extra="allow") + + def to_body(self, *, drop_params: bool) -> Dict[str, object]: + known = self.model_dump( + exclude_none=True, exclude=set((self.model_extra or {}).keys()) + ) + if drop_params: + return known + return {**known, **(self.model_extra or {})} + + +class _CohereCommandRBody(BedrockInvokeInferenceParams): + max_tokens: Optional[int] = None + stream: Optional[bool] = None + temperature: Optional[float] = None + p: Optional[float] = None + k: Optional[float] = None + seed: Optional[int] = None + frequency_penalty: Optional[float] = None + presence_penalty: Optional[float] = None + stop_sequences: Optional[List[str]] = None + preamble: Optional[str] = None + prompt_truncation: Optional[str] = None + return_prompt: Optional[bool] = None + raw_prompting: Optional[bool] = None + search_queries_only: Optional[bool] = None + + +class _CohereLegacyBody(BedrockInvokeInferenceParams): + max_tokens: Optional[int] = None + stream: Optional[bool] = None + temperature: Optional[float] = None + p: Optional[float] = None + seed: Optional[int] = None + frequency_penalty: Optional[float] = None + presence_penalty: Optional[float] = None + stop_sequences: Optional[List[str]] = None + num_generations: Optional[int] = None + return_likelihood: Optional[str] = None + + +class _AI21Body(BedrockInvokeInferenceParams): + maxTokens: Optional[int] = None + temperature: Optional[float] = None + topP: Optional[float] = None + stream: Optional[bool] = None + stopSequences: Optional[List[str]] = None + + +class _MistralBody(BedrockInvokeInferenceParams): + max_tokens: Optional[int] = None + temperature: Optional[float] = None + top_p: Optional[float] = None + top_k: Optional[float] = None + stream: Optional[bool] = None + stop: Optional[List[str]] = None + + +class _TitanTextGenerationConfig(BedrockInvokeInferenceParams): + maxTokenCount: Optional[int] = None + temperature: Optional[float] = None + topP: Optional[float] = None + stopSequences: Optional[List[str]] = None + + +class _LlamaBody(BedrockInvokeInferenceParams): + max_gen_len: Optional[int] = None + temperature: Optional[float] = None + top_p: Optional[float] = None + topP: Optional[float] = None + stream: Optional[bool] = None + + +_COHERE_COMMAND_R_PREFIX = "cohere.command-r" + +_INVOKE_BODY_MODELS: Dict[str, Type[BedrockInvokeInferenceParams]] = { + "cohere_command_r": _CohereCommandRBody, + "cohere": _CohereLegacyBody, + "ai21": _AI21Body, + "mistral": _MistralBody, + "amazon": _TitanTextGenerationConfig, + "meta": _LlamaBody, + "llama": _LlamaBody, + "deepseek_r1": _LlamaBody, +} + + +def _resolve_invoke_body_model( + provider: Optional[str], model: str +) -> Optional[Type[BedrockInvokeInferenceParams]]: + if provider == "cohere" and model.startswith(_COHERE_COMMAND_R_PREFIX): + return _CohereCommandRBody + if provider is None: + return None + return _INVOKE_BODY_MODELS.get(provider) + + +def parse_invoke_inference_params( + provider: Optional[str], + model: str, + params: Dict[str, object], + drop_params: bool, +) -> Dict[str, object]: + """Validate inference params against the sub-provider's typed body. + + Providers without a typed body (e.g. delegating ones) pass through + unchanged, preserving today's behavior. + """ + body_model = _resolve_invoke_body_model(provider, model) + if body_model is None: + return params + return body_model.model_validate(params).to_body(drop_params=drop_params) + + +def assert_no_control_params(body: Dict[str, object]) -> None: + """Guard run right before dispatch: control keys must never reach the wire.""" + leaked = LITELLM_CONTROL_PARAM_KEYS & body.keys() + if leaked: + raise ValueError( + "litellm control params leaked into the Bedrock request body: " + f"{sorted(leaked)}" + ) + + +def assert_no_control_params_in_payload(data: str) -> None: + """Dispatch-time guard over the serialized request payload.""" + try: + parsed = json.loads(data) + except (TypeError, ValueError): + return + if isinstance(parsed, dict): + assert_no_control_params(parsed) 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 6666d051aea..96c53591b3c 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 @@ -3,7 +3,9 @@ import sys import pytest -sys.path.insert(0, os.path.abspath("../../../../../..")) # Adds the parent directory to the system path +sys.path.insert( + 0, os.path.abspath("../../../../../..") +) # Adds the parent directory to the system path from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( AmazonAnthropicClaudeConfig, @@ -42,7 +44,9 @@ def test_signed_invoke_body_drops_stream_chunk_size(config, model): headers={}, optional_params={}, request_data=request_body, - api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/{}/invoke".format(model), + api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/{}/invoke".format( + model + ), api_key="test-bearer-token", model=model, stream=True, @@ -51,3 +55,35 @@ def test_signed_invoke_body_drops_stream_chunk_size(config, model): assert signed_body is not None assert "stream_chunk_size" not in signed_body.decode() assert "max_tokens" in signed_body.decode() + + +def test_extra_body_passthrough_by_default(): + """Unknown body keys are forwarded verbatim when drop_params is off, so the + soft-allowlist escape hatch keeps working.""" + cfg = AmazonInvokeConfig() + request_body = cfg.transform_request( + model="mistral.mistral-7b-instruct-v0:2", + messages=[{"role": "user", "content": "hi"}], + optional_params={"temperature": 0.5, "made_up_param": "x"}, + litellm_params={}, + headers={}, + ) + + assert request_body["temperature"] == 0.5 + assert request_body["made_up_param"] == "x" + + +def test_drop_params_strips_extra_body_but_keeps_known_params(): + """drop_params selects the strict body: typed provider keys survive, unknown + passthrough keys are dropped instead of being shipped to the provider.""" + cfg = AmazonInvokeConfig() + request_body = cfg.transform_request( + model="mistral.mistral-7b-instruct-v0:2", + messages=[{"role": "user", "content": "hi"}], + optional_params={"temperature": 0.5, "made_up_param": "x"}, + litellm_params={"drop_params": True}, + headers={}, + ) + + assert request_body["temperature"] == 0.5 + assert "made_up_param" not in request_body diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py index 61987d25d9c..34ce75abe9a 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py @@ -297,6 +297,47 @@ def test_make_sync_call_honors_explicit_stream_chunk_size(): response.iter_bytes.assert_called_once_with(chunk_size=2048) +@pytest.mark.asyncio +async def test_make_call_guards_against_leaked_control_param(): + """The dispatch wrapper validates the payload right before sending, so a + control param that slipped into the body aborts the request instead of + being rejected downstream by Bedrock.""" + client = MagicMock() + client.post = AsyncMock() + + with pytest.raises(ValueError, match="stream_chunk_size"): + await make_call( + client=client, + api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/cohere.command-text-v14/invoke-with-response-stream", + headers={}, + data='{"prompt": "hi", "stream_chunk_size": 2048}', + model="cohere.command-text-v14", + messages=[], + logging_obj=MagicMock(), + ) + + client.post.assert_not_called() + + +def test_make_sync_call_guards_against_leaked_control_param(): + client = MagicMock() + client.post = MagicMock() + + with pytest.raises(ValueError, match="stream_chunk_size"): + make_sync_call( + client=client, + api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/cohere.command-text-v14/invoke-with-response-stream", + headers={}, + data='{"prompt": "hi", "stream_chunk_size": 2048}', + signed_json_body=None, + model="cohere.command-text-v14", + messages=[], + logging_obj=MagicMock(), + ) + + client.post.assert_not_called() + + def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default(): mock_response = MagicMock() mock_response.status_code = 200 diff --git a/tests/test_litellm/types/llms/test_bedrock_invoke.py b/tests/test_litellm/types/llms/test_bedrock_invoke.py new file mode 100644 index 00000000000..3c5bb33c5be --- /dev/null +++ b/tests/test_litellm/types/llms/test_bedrock_invoke.py @@ -0,0 +1,78 @@ +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.types.llms.bedrock_invoke import ( + assert_no_control_params, + assert_no_control_params_in_payload, + parse_invoke_inference_params, +) + + +def test_passthrough_keeps_unknown_keys(): + body = parse_invoke_inference_params( + provider="mistral", + model="mistral.mistral-7b-instruct-v0:2", + params={"temperature": 0.5, "unknown_key": "x"}, + drop_params=False, + ) + assert body == {"temperature": 0.5, "unknown_key": "x"} + + +def test_drop_params_strips_only_unknown_keys(): + body = parse_invoke_inference_params( + provider="mistral", + model="mistral.mistral-7b-instruct-v0:2", + params={"temperature": 0.5, "unknown_key": "x"}, + drop_params=True, + ) + assert body == {"temperature": 0.5} + + +def test_command_r_and_legacy_resolve_to_different_bodies(): + """command-r exposes k/p; the legacy text model does not, so the same key + must be classified differently per model id.""" + command_r = parse_invoke_inference_params( + provider="cohere", + model="cohere.command-r-v1:0", + params={"k": 2}, + drop_params=True, + ) + legacy = parse_invoke_inference_params( + provider="cohere", + model="cohere.command-text-v14", + params={"k": 2}, + drop_params=True, + ) + assert command_r == {"k": 2.0} + assert legacy == {} + + +def test_unmodeled_provider_passes_through_untouched(): + params = {"anything": 1, "stream_chunk_size": 4} + assert ( + parse_invoke_inference_params( + provider="anthropic", + model="anthropic.claude-sonnet-4-6", + params=params, + drop_params=True, + ) + == params + ) + + +def test_guard_raises_on_leaked_control_param(): + with pytest.raises(ValueError, match="stream_chunk_size"): + assert_no_control_params({"temperature": 0.5, "stream_chunk_size": 2048}) + + +def test_guard_payload_ignores_non_dict_and_invalid_json(): + assert_no_control_params_in_payload("not json") + assert_no_control_params_in_payload("[1, 2, 3]") + assert_no_control_params_in_payload('{"temperature": 0.5}') + + with pytest.raises(ValueError, match="stream_chunk_size"): + assert_no_control_params_in_payload('{"stream_chunk_size": 2048}')