mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge dcf2bc7eef into 6437b812be
This commit is contained in:
commit
d22ab81b5f
12 changed files with 409 additions and 1232 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -101,6 +101,7 @@ STABILIZATION_TODO.md
|
|||
**/playwright-report
|
||||
**/*.storageState.json
|
||||
**/coverage
|
||||
.coverage
|
||||
test-config
|
||||
|
||||
# ---------- Terraform ----------
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -154,34 +154,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
|
||||
return _anthropic_request
|
||||
|
||||
async def async_transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
_anthropic_request = self._build_bedrock_anthropic_request_base(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
await self._async_convert_document_url_sources_to_base64(_anthropic_request)
|
||||
beta_list = self._compute_bedrock_invoke_beta_headers(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
)
|
||||
if beta_list:
|
||||
_anthropic_request["anthropic_beta"] = beta_list
|
||||
|
||||
return _anthropic_request
|
||||
|
||||
def _build_bedrock_anthropic_request_base(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -190,11 +162,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
|
||||
}
|
||||
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)
|
||||
|
|
@ -215,7 +183,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
|
||||
anthropic_request.pop("model", None)
|
||||
anthropic_request.pop("stream", None)
|
||||
anthropic_request.pop("stream_chunk_size", None)
|
||||
output_format = anthropic_request.pop("output_format", None)
|
||||
output_config_format = pop_bedrock_invoke_output_config_format(
|
||||
anthropic_request
|
||||
|
|
|
|||
|
|
@ -24,6 +24,8 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
HTTPHandler,
|
||||
_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
|
||||
|
|
@ -140,6 +142,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,
|
||||
|
|
@ -150,7 +160,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
) -> dict:
|
||||
## SETUP ##
|
||||
stream = optional_params.pop("stream", None)
|
||||
optional_params.pop("stream_chunk_size", None)
|
||||
custom_prompt_dict: dict = litellm_params.pop("custom_prompt_dict", None) or {}
|
||||
hf_model_name = litellm_params.get("hf_model_name", None)
|
||||
|
||||
|
|
@ -162,12 +171,13 @@ 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
|
||||
}
|
||||
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":
|
||||
if model.startswith("cohere.command-r"):
|
||||
|
|
|
|||
|
|
@ -1,11 +1,24 @@
|
|||
import json
|
||||
from typing import Any, Dict, List, Literal, Optional, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TYPE_CHECKING, Required, TypedDict, override
|
||||
|
||||
from .openai import ChatCompletionToolCallChunk
|
||||
|
||||
|
||||
class LiteLLMControlParams(BaseModel):
|
||||
"""LiteLLM-internal control parameters that must never be serialized into a
|
||||
provider request body. They live in optional_params for convenience but
|
||||
govern client-side behavior (e.g. how the HTTP response stream is
|
||||
re-chunked), so Bedrock rejects them as unknown fields."""
|
||||
|
||||
stream_chunk_size: int | None = None
|
||||
|
||||
|
||||
LITELLM_CONTROL_PARAM_KEYS = frozenset(LiteLLMControlParams.model_fields)
|
||||
|
||||
|
||||
class CachePointBlock(TypedDict, total=False):
|
||||
type: Literal["default"]
|
||||
ttl: str
|
||||
|
|
|
|||
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)
|
||||
|
|
@ -3559,89 +3559,6 @@ def test_bedrock_openai_model_id_extraction():
|
|||
print(f"✓ Model ID extracted and encoded: {model_id}")
|
||||
|
||||
|
||||
def test_bedrock_openai_convert_messages_to_prompt():
|
||||
"""
|
||||
Test that convert_messages_to_prompt returns empty string for OpenAI models.
|
||||
"""
|
||||
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
|
||||
|
||||
bedrock_llm = BedrockLLM()
|
||||
messages = [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
|
||||
prompt, chat_history = bedrock_llm.convert_messages_to_prompt(
|
||||
model="test-model", messages=messages, provider="openai", custom_prompt_dict={}
|
||||
)
|
||||
|
||||
# OpenAI models use messages directly, no prompt conversion
|
||||
assert prompt == ""
|
||||
assert chat_history is None
|
||||
print("✓ convert_messages_to_prompt returns empty for OpenAI")
|
||||
|
||||
|
||||
def test_bedrock_openai_response_parsing():
|
||||
"""
|
||||
Test that OpenAI responses are correctly parsed.
|
||||
"""
|
||||
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
|
||||
from litellm import ModelResponse
|
||||
from unittest.mock import Mock
|
||||
import json
|
||||
|
||||
bedrock_llm = BedrockLLM()
|
||||
|
||||
# Mock OpenAI-style response
|
||||
openai_response = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": "The capital of France is Paris.",
|
||||
"role": "assistant",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
|
||||
}
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.json.return_value = openai_response
|
||||
mock_response.text = json.dumps(openai_response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ModelResponse()
|
||||
mock_logging = Mock()
|
||||
|
||||
result = bedrock_llm.process_response(
|
||||
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
|
||||
response=mock_response,
|
||||
model_response=model_response,
|
||||
stream=False,
|
||||
logging_obj=mock_logging,
|
||||
optional_params={},
|
||||
api_key="",
|
||||
data={},
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
print_verbose=lambda x: None,
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
# Verify response content
|
||||
assert result.choices[0].message.content == "The capital of France is Paris."
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
|
||||
# Verify usage
|
||||
assert result.usage.prompt_tokens == 10
|
||||
assert result.usage.completion_tokens == 8
|
||||
assert result.usage.total_tokens == 18
|
||||
|
||||
print("✓ OpenAI response parsing works correctly")
|
||||
|
||||
|
||||
def test_bedrock_openai_request_transformation():
|
||||
"""
|
||||
Test that the request is correctly transformed for OpenAI models.
|
||||
|
|
@ -3831,46 +3748,6 @@ def test_bedrock_openai_multiple_message_types():
|
|||
print("✓ Multiple message types handled correctly")
|
||||
|
||||
|
||||
def test_bedrock_openai_error_handling():
|
||||
"""
|
||||
Test that errors from OpenAI models are properly handled.
|
||||
"""
|
||||
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
|
||||
from litellm import ModelResponse
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from unittest.mock import Mock
|
||||
import json
|
||||
|
||||
bedrock_llm = BedrockLLM()
|
||||
|
||||
# Mock error response
|
||||
mock_response = Mock()
|
||||
mock_response.json.side_effect = Exception("Invalid JSON")
|
||||
mock_response.text = "Invalid response"
|
||||
mock_response.status_code = 422
|
||||
|
||||
model_response = ModelResponse()
|
||||
mock_logging = Mock()
|
||||
|
||||
with pytest.raises(BedrockError) as exc_info:
|
||||
bedrock_llm.process_response(
|
||||
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
|
||||
response=mock_response,
|
||||
model_response=model_response,
|
||||
stream=False,
|
||||
logging_obj=mock_logging,
|
||||
optional_params={},
|
||||
api_key="",
|
||||
data={},
|
||||
messages=[],
|
||||
print_verbose=lambda x: None,
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 422
|
||||
print("✓ Error handling works correctly")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Nova Grounding (web_search_options) Unit Tests (Mocked)
|
||||
# ============================================================================
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
|
@ -25,12 +24,15 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
|
|||
(AmazonAnthropicClaudeConfig, "anthropic.claude-sonnet-4-6"),
|
||||
],
|
||||
)
|
||||
def test_transform_request_drops_stream_chunk_size(config, model):
|
||||
def test_signed_invoke_body_drops_stream_chunk_size(config, model):
|
||||
"""stream_chunk_size is a LiteLLM-internal knob for re-chunking the HTTP
|
||||
response stream. Leaking it into the provider request body makes Bedrock
|
||||
reject the whole request: ValidationException 'stream_chunk_size: Extra
|
||||
inputs are not permitted'."""
|
||||
request_body = config().transform_request(
|
||||
inputs are not permitted'. The two invoke transform entry points build the
|
||||
body differently, so this asserts on the actual signed wire bytes that both
|
||||
funnel through, regardless of which transform produced them."""
|
||||
cfg = config()
|
||||
request_body = cfg.transform_request(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"stream": True, "stream_chunk_size": 2048, "max_tokens": 10},
|
||||
|
|
@ -38,4 +40,50 @@ def test_transform_request_drops_stream_chunk_size(config, model):
|
|||
headers={},
|
||||
)
|
||||
|
||||
assert "stream_chunk_size" not in json.dumps(request_body)
|
||||
_, signed_body = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={},
|
||||
request_data=request_body,
|
||||
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/{}/invoke".format(
|
||||
model
|
||||
),
|
||||
api_key="test-bearer-token",
|
||||
model=model,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -8,14 +8,11 @@ sys.path.insert(
|
|||
0, os.path.abspath("../../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.chat.invoke_handler import (
|
||||
AWSEventStreamDecoder,
|
||||
BedrockLLM,
|
||||
make_call,
|
||||
make_sync_call,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
|
||||
def test_transform_thinking_blocks_with_redacted_content():
|
||||
|
|
@ -297,32 +294,42 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
|
|||
response.iter_bytes.assert_called_once_with(chunk_size=2048)
|
||||
|
||||
|
||||
def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default():
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.iter_bytes = MagicMock(return_value=iter([]))
|
||||
client = HTTPHandler()
|
||||
client.post = MagicMock(return_value=mock_response)
|
||||
@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()
|
||||
|
||||
BedrockLLM().completion(
|
||||
model="cohere.command-text-v14",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_base=None,
|
||||
custom_prompt_dict={},
|
||||
model_response=litellm.ModelResponse(),
|
||||
print_verbose=lambda *args, **kwargs: None,
|
||||
encoding=litellm.encoding,
|
||||
logging_obj=MagicMock(),
|
||||
optional_params={
|
||||
"stream": True,
|
||||
"aws_access_key_id": "fake",
|
||||
"aws_secret_access_key": "fake",
|
||||
"aws_region_name": "us-east-1",
|
||||
},
|
||||
acompletion=False,
|
||||
timeout=None,
|
||||
litellm_params={},
|
||||
client=client,
|
||||
)
|
||||
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(),
|
||||
)
|
||||
|
||||
mock_response.iter_bytes.assert_called_once_with(chunk_size=None)
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -202,6 +202,27 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
|
|||
response.iter_bytes.assert_called_once_with(chunk_size=2048)
|
||||
|
||||
|
||||
def test_make_sync_call_guards_against_leaked_control_param():
|
||||
"""The converse 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 = 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/anthropic.claude-sonnet-4-6/converse-stream",
|
||||
headers={},
|
||||
data='{"messages": [], "stream_chunk_size": 2048}',
|
||||
model="anthropic.claude-sonnet-4-6",
|
||||
messages=[],
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
client.post.assert_not_called()
|
||||
|
||||
|
||||
def test_completion_plumbs_stream_chunk_size_through_converse():
|
||||
iter_bytes_spy = _stream_completion_with_spied_iter_bytes(
|
||||
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0"
|
||||
|
|
|
|||
91
tests/test_litellm/types/llms/test_bedrock_invoke.py
Normal file
91
tests/test_litellm/types/llms/test_bedrock_invoke.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
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_none_provider_passes_through_untouched():
|
||||
params = {"anything": 1}
|
||||
assert (
|
||||
parse_invoke_inference_params(
|
||||
provider=None,
|
||||
model="some-unresolved-model",
|
||||
params=params,
|
||||
drop_params=True,
|
||||
)
|
||||
== params
|
||||
)
|
||||
|
||||
|
||||
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