mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(bedrock): type invoke request bodies, split extra_body on drop_params, guard control-param leaks at dispatch
Parse each invoke sub-provider's inference params into a typed Pydantic body so known wire keys are kept and unknown keys ride as extra_body passthrough, dropped only when drop_params is set. Add a dispatch-time guard in make_call/make_sync_call (invoke and converse) that aborts the request if a litellm control param such as stream_chunk_size leaked into the body, and point the converse path at the same shared control-key source of truth used by invoke.
This commit is contained in:
parent
86fd1358e5
commit
f4b56ae89a
7 changed files with 333 additions and 5 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
160
litellm/types/llms/bedrock_invoke.py
Normal file
160
litellm/types/llms/bedrock_invoke.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
78
tests/test_litellm/types/llms/test_bedrock_invoke.py
Normal file
78
tests/test_litellm/types/llms/test_bedrock_invoke.py
Normal file
|
|
@ -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}')
|
||||
Loading…
Add table
Reference in a new issue