From e4947231058394bcfaac14f6e3f0700d4be5644c Mon Sep 17 00:00:00 2001 From: shrey-berri Date: Sat, 26 Sep 2026 16:01:20 -0700 Subject: [PATCH] fix(params): validate stream_chunk_size once, before any provider call (#43222) * fix(params): validate stream_chunk_size once and carry it as typed control options Checks stream_chunk_size at the top of completion() and acompletion(), accepts digit strings, returns a 400 naming the param unless drop_params is set, and stores the checked value under _litellm_control. Bedrock Converse and Invoke read it from litellm_params; the Bedrock-only checker and the dead Invoke pops are gone. Owned-kwarg filtering now runs through one helper everywhere. Refs LIT-8317 * test(bedrock): drop tests for the removed stream_chunk_size_from helper Refs LIT-8317 * fix(params): check stream_chunk_size before the MCP gateway branch Refs LIT-8317 * fix(params): return assert_never in the exhaustive control-options match Refs LIT-8317 * fix(params): address council review of the control options change Read all_litellm_params live so names registered after import stay LiteLLM-owned, make litellm_params a required keyword on the stream wrapper hooks, give digit strings and ints the same 18-digit range, share the default-chunking test table, test the Responses bridge through litellm.responses, and revert formatting-only churn in existing tests. Refs LIT-8317 * fix(params): address the second council review of control options Keep the Responses bridge on its original all_litellm_params forwarding, narrow _int_from_decimal_string inline so it type-checks, bound nested huge ints in the error message, store _litellm_control only when a value is set, simplify the parser to its single field, drop the one-caller wrapper, and tighten the tests. Refs LIT-8317 * fix(params): keep the 18-digit length check on stream_chunk_size strings A 19-character string with leading zeros such as 0000000000000000001 would otherwise pass as 1, although the rule and the error message say at most 18 digits. Refs LIT-8317 * test(params): tidy control options tests after council sign-off Move the Responses bridge test into the existing bridge test file, drop the rebind test that pinned an implementation detail, assert through stored_control_options instead of the storage key, and cover drop_params="true" through Bedrock streaming. Refs LIT-8317 * test(params): wrap a chunking test row that went past 120 characters Refs LIT-8317 --- litellm/caching/caching.py | 5 +- litellm/constants.py | 1 + litellm/images/main.py | 11 +- .../litellm_core_utils/get_litellm_params.py | 53 +++- litellm/llms/base_llm/chat/transformation.py | 4 + .../bedrock/chat/agentcore/transformation.py | 6 +- litellm/llms/bedrock/chat/converse_handler.py | 5 +- .../anthropic_claude3_transformation.py | 1 - .../base_invoke_transformation.py | 13 +- litellm/llms/bedrock/common_utils.py | 11 +- litellm/llms/bytez/chat/transformation.py | 5 + litellm/llms/custom_httpx/llm_http_handler.py | 2 + litellm/llms/langgraph/chat/transformation.py | 5 + litellm/llms/oci/chat/transformation.py | 6 +- litellm/llms/sagemaker/chat/transformation.py | 5 + .../vertex_ai/agent_engine/transformation.py | 5 + litellm/main.py | 37 ++- litellm/types/litellm_params.py | 28 ++- litellm/utils.py | 24 +- tests/_support/stream_chunk_size.py | 34 +-- .../providers/test_internal_params_wire.py | 9 +- tests/unit/caching/test_caching.py | 14 ++ .../test_get_litellm_params.py | 84 ++++++- .../test_base_invoke_transformation.py | 169 +++++++------ tests/unit/llms/bedrock/test_common_utils.py | 20 -- tests/unit/llms/chat/test_converse_handler.py | 135 ++++------ .../oci/chat/test_oci_chat_transformation.py | 2 + .../unit/llms/oci/test_oci_coverage_boost.py | 2 + .../test_sagemaker_chat_transformation.py | 3 + .../test_sagemaker_nova_transformation.py | 2 + .../test_responses_api_bridge_flag.py | 18 ++ tests/unit/test_filter_out_litellm_params.py | 20 ++ tests/unit/test_main.py | 233 +++++++++++++++++- tests/unit/types/test_litellm_params.py | 10 +- 34 files changed, 695 insertions(+), 287 deletions(-) delete mode 100644 tests/unit/llms/bedrock/test_common_utils.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index a31cad4af29..d766d1a58bc 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -25,7 +25,7 @@ from litellm._logging import verbose_logger from litellm.constants import CACHED_STREAMING_CHUNK_DELAY from litellm.litellm_core_utils.model_param_helper import ModelParamHelper from litellm.types.caching import * -from litellm.types.utils import EmbeddingResponse, all_litellm_params +from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg from .azure_blob_cache import AzureBlobCache from .base_cache import BaseCache @@ -377,7 +377,6 @@ class Cache: return preset_cache_key combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params() - litellm_param_kwargs: Final = all_litellm_params is_semantic_cache: Final = self._is_semantic_cache() scope_excluded_params: Final = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset() for param in kwargs: @@ -387,7 +386,7 @@ class Cache: param_value: str | None = self._get_param_value(param, kwargs) if param_value is not None: cache_key += f"{param}: {param_value}" - elif param not in litellm_param_kwargs: # check if user passed in optional param - e.g. top_k + elif not is_litellm_owned_kwarg(param): if litellm.enable_caching_on_provider_specific_optional_params is True: # feature flagged for now if kwargs[param] is None: continue # ignore None params diff --git a/litellm/constants.py b/litellm/constants.py index dac15c01fbf..a292b654778 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1611,6 +1611,7 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = { # Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.) PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-" INTERNAL_KWARG_PREFIX: Final = "_litellm_" +CONTROL_OPTIONS_KEY: Final = f"{INTERNAL_KWARG_PREFIX}control" AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech" AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech" diff --git a/litellm/images/main.py b/litellm/images/main.py index 5ca8a726a69..7dc68dafecc 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -25,7 +25,7 @@ from litellm.llms.base_llm import BaseImageEditConfig, BaseImageGenerationConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.custom_llm import CustomLLM -from litellm.utils import exception_type, get_litellm_params +from litellm.utils import exception_type, filter_out_litellm_params, get_litellm_params #################### Initialize provider clients #################### llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler() @@ -52,7 +52,6 @@ from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( LITELLM_IMAGE_VARIATION_PROVIDERS, LlmProviders, - is_litellm_owned_kwarg, ) from litellm.utils import ( ImageResponse, @@ -249,9 +248,7 @@ def image_generation( "size", "style", ] - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k) - } + non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params) image_generation_config: BaseImageGenerationConfig | None = None if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values(): @@ -755,9 +752,7 @@ def image_edit( "style", "async_call", ] - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k) - } + non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params) litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) model_info: Final = kwargs.get("model_info", None) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index f7aaef3a51f..f28259a1b7f 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -1,9 +1,15 @@ +import reprlib from collections.abc import Mapping, MutableMapping +from dataclasses import dataclass, fields from types import MappingProxyType from typing import Final +from pydantic import TypeAdapter, ValidationError + +from litellm.constants import CONTROL_OPTIONS_KEY from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.llms.openai.data_residency import infer_openai_data_residency +from litellm.types.litellm_params import MAX_CONTROL_INT_DIGITS, ControlOptions from litellm.types.router import CustomPricingLiteLLMParams AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset( @@ -70,6 +76,51 @@ OPTIONAL_KWARGS_KEYS: Final = ( # Backward-compatible alias for existing imports/tests. _OPTIONAL_KWARGS_KEYS: Final = OPTIONAL_KWARGS_KEYS +_CONTROL_OPTIONS: Final = TypeAdapter(ControlOptions) +_CONTROL_OPTION_NAMES: Final = tuple(field.name for field in fields(ControlOptions)) +_MAX_SHOWN_INT_BITS: Final = 64 +_EXPECTED: Final = f"expected a positive integer of at most {MAX_CONTROL_INT_DIGITS} digits" + + +class _BoundedRepr(reprlib.Repr): + def repr_int(self, x: int, level: int) -> str: + if x.bit_length() > _MAX_SHOWN_INT_BITS: + return f"" + return super().repr_int(x, level) + + +_BOUNDED_REPR: Final = _BoundedRepr() + + +@dataclass(frozen=True, slots=True) +class InvalidControlOption: + param: str + message: str + + +def parse_control_options(kwargs: Mapping[str, object]) -> ControlOptions | InvalidControlOption: + given: Final = { # mutable-ok: TypeAdapter.validate_python takes a dict + name: kwargs[name] for name in _CONTROL_OPTION_NAMES if name in kwargs + } + try: + return _CONTROL_OPTIONS.validate_python(given) + except ValidationError as e: + param: Final = str(e.errors(include_url=False)[0]["loc"][0]) + return InvalidControlOption( + param=param, message=f"Invalid {param}={_BOUNDED_REPR.repr(given[param])}: {_EXPECTED}" + ) + + +def stored_control_options(litellm_params: Mapping[str, object]) -> ControlOptions: + control: Final = litellm_params.get(CONTROL_OPTIONS_KEY) + return control if isinstance(control, ControlOptions) else ControlOptions() + + +def with_control_options(litellm_params: Mapping[str, object], control: ControlOptions) -> dict[str, object]: + if control == ControlOptions(): + return dict(litellm_params) # mutable-ok: completion() hands litellm_params to provider code typed as dict + return {**litellm_params, CONTROL_OPTIONS_KEY: control} # mutable-ok: same dict contract as above + def _get_base_model_from_litellm_call_metadata( metadata: dict | None, @@ -130,7 +181,6 @@ def get_litellm_params( api_version: str | None = None, max_retries: int | None = None, litellm_request_debug: bool | None = None, - stream_chunk_size: int | None = None, **kwargs, ) -> dict: _litellm_metadata_dict: Final = litellm_metadata if isinstance(litellm_metadata, dict) else None @@ -193,7 +243,6 @@ def get_litellm_params( "max_retries": max_retries, "use_litellm_proxy": use_litellm_proxy, "litellm_request_debug": litellm_request_debug, - "stream_chunk_size": stream_chunk_size, } # Sparse extraction: only add kwargs keys that are actually present diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 948a9bc6852..f1b41a2302d 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -393,6 +393,8 @@ class BaseConfig(ABC): client: AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": raise NotImplementedError @@ -408,6 +410,8 @@ class BaseConfig(ABC): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": raise NotImplementedError diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py index 29133bcfaf9..2ad23e84e9f 100644 --- a/litellm/llms/bedrock/chat/agentcore/transformation.py +++ b/litellm/llms/bedrock/chat/agentcore/transformation.py @@ -5,7 +5,7 @@ https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agentcore_InvokeAgen """ import json -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Mapping from typing import TYPE_CHECKING, Any, Final, Optional, Union from urllib.parse import quote @@ -643,6 +643,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """ Simplified sync streaming - returns a generator that yields ModelResponse chunks. @@ -862,6 +864,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): client: Optional["AsyncHTTPHandler"] = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """ Simplified async streaming - returns an async generator that yields ModelResponse chunks. diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 65a34f72167..28df1af4bc0 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -7,6 +7,7 @@ import litellm from litellm.anthropic_beta_headers_manager import ( update_headers_with_filtered_beta, ) +from litellm.litellm_core_utils.get_litellm_params import stored_control_options from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -18,7 +19,7 @@ from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, pop_aws_auth_params, run_aws_signing -from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text, stream_chunk_size_from +from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call @@ -280,7 +281,7 @@ class BedrockConverseLLM(BaseAWSLLM): ): ## SETUP ## stream: Final = optional_params.pop("stream", None) - stream_chunk_size: Final = stream_chunk_size_from(litellm_params) if stream is True else None + stream_chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size if stream is True else None unencoded_model_id: Final = optional_params.pop("model_id", None) fake_stream = optional_params.pop("fake_stream", False) json_mode: Final = optional_params.get("json_mode", False) 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 a8b94fb5703..c5abb5e9a1c 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -225,7 +225,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): anthropic_request.pop("model", None) anthropic_request.pop("stream", None) - anthropic_request.pop("stream_chunk_size", None) apply_bedrock_invoke_structured_output( model=model, request_body=anthropic_request, 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 629806b58e2..9baf8110b4e 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -1,6 +1,7 @@ import copy import json import time +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast, get_args import httpx @@ -9,6 +10,7 @@ from pydantic import TypeAdapter, ValidationError import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import map_finish_reason +from litellm.litellm_core_utils.get_litellm_params import stored_control_options from litellm.litellm_core_utils.logging_utils import track_llm_api_timing from litellm.litellm_core_utils.prompt_templates.factory import ( cohere_message_pt, @@ -18,7 +20,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.bedrock.chat.invoke_handler import make_call, make_sync_call -from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from +from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.bedrock.request_metadata import ( bedrock_request_metadata_headers, merge_bedrock_invoke_headers, @@ -180,7 +182,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): ) -> dict: ## SETUP ## stream: Final = optional_params.pop("stream", None) - optional_params.pop("stream_chunk_size", None) custom_prompt_dict: Final[dict] = litellm_params.pop("custom_prompt_dict", None) or {} hf_model_name: Final = litellm_params.get("hf_model_name", None) @@ -452,8 +453,10 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): client: AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: - chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params) + chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size completion_stream, response_headers = await make_call( client=client, api_base=api_base, @@ -489,11 +492,13 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: sync_client: Final = ( _get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client ) - chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params) + chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size completion_stream, response_headers = make_sync_call( client=sync_client, api_base=api_base, diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 5f044897b2c..ccc4309fc5d 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -18,7 +18,7 @@ if TYPE_CHECKING: from litellm.types.llms.bedrock import BedrockCreateBatchRequest import httpx -from pydantic import ConfigDict, TypeAdapter, ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm import verbose_logger @@ -86,15 +86,6 @@ class BedrockError(BaseLLMException): _BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (*AWS_AUTH_PARAM_KEYS, "aws_region_name") -_STREAM_CHUNK_SIZE_VALIDATOR: Final[TypeAdapter[int | None]] = TypeAdapter(int | None, config=ConfigDict(strict=True)) - - -def stream_chunk_size_from(litellm_params: Mapping[str, object]) -> int | None: - raw: Final = litellm_params.get("stream_chunk_size") - try: - return _STREAM_CHUNK_SIZE_VALIDATOR.validate_python(raw) - except ValidationError as e: - raise BedrockError(status_code=400, message=f"Invalid stream_chunk_size={raw!r}. Expected int. Error: {e}") def merge_bedrock_aws_request_params( diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index e622761dd7f..5846ba560a8 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -1,6 +1,7 @@ import json import time import traceback +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx @@ -258,6 +259,8 @@ class BytezChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -300,6 +303,8 @@ class BytezChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={}) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 0cb1416db3f..8aa38ff3341 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -790,6 +790,7 @@ class BaseLLMHTTPHandler: messages=messages, client=client, json_mode=json_mode, + litellm_params=litellm_params, ) completion_stream, headers = self.make_sync_call( provider_config=provider_config, @@ -953,6 +954,7 @@ class BaseLLMHTTPHandler: client=client, json_mode=json_mode, signed_json_body=signed_json_body, + litellm_params=litellm_params, ) completion_stream, _response_headers = await self.make_async_call_stream_helper( diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py index c9388ee472f..293672f1ca9 100644 --- a/litellm/llms/langgraph/chat/transformation.py +++ b/litellm/llms/langgraph/chat/transformation.py @@ -9,6 +9,7 @@ Non-streaming endpoint: POST /runs/wait """ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast import httpx @@ -285,6 +286,8 @@ class LangGraphConfig(BaseConfig): client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for synchronous streaming. @@ -344,6 +347,8 @@ class LangGraphConfig(BaseConfig): client: Optional["AsyncHTTPHandler"] = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for asynchronous streaming. diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index ecff823a18d..24f3ddd5162 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -10,7 +10,7 @@ implement the LiteLLM BaseConfig interface. Heavy-lifting lives in: """ import json -from collections.abc import AsyncIterator, Callable, Iterator +from collections.abc import AsyncIterator, Callable, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final import httpx @@ -642,6 +642,8 @@ class OCIChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "OCIStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -681,6 +683,8 @@ class OCIChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "OCIStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.OCI, params={}) diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 04995f32d97..f99a3f9e1bc 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -7,6 +7,7 @@ LiteLLM Docs: https://docs.litellm.ai/docs/providers/aws_sagemaker#sagemaker-mes Huggingface Docs: https://huggingface.co/docs/text-generation-inference/en/messages_api """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast import httpx @@ -149,6 +150,8 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -191,6 +194,8 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: if client is None or isinstance(client, HTTPHandler): try: diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index b37bf473731..cf889c481a3 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -10,6 +10,7 @@ API Reference: """ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast import httpx @@ -365,6 +366,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for synchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( @@ -423,6 +426,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): client: Optional["AsyncHTTPHandler"] = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for asynchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( diff --git a/litellm/main.py b/litellm/main.py index 8c2afe4429a..6c85adf3ae8 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -38,7 +38,7 @@ import dotenv import httpx import openai from pydantic import BaseModel -from typing_extensions import overload +from typing_extensions import assert_never, overload import litellm @@ -48,6 +48,7 @@ from litellm import client # Other utils are imported directly to avoid circular imports from litellm.utils import ( exception_type, + filter_out_litellm_params, get_litellm_params, get_optional_params, peek_reasoning_summary_aliases, @@ -83,6 +84,9 @@ from litellm.litellm_core_utils.get_litellm_params import ( AWS_CREDENTIAL_KWARGS_KEYS, OPTIONAL_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY, + InvalidControlOption, + parse_control_options, + with_control_options, ) from litellm.litellm_core_utils.get_provider_specific_headers import ( ProviderSpecificHeaderUtils, @@ -127,7 +131,7 @@ from litellm.types.completion import ( _CompletionDispatchContext, _CompletionDispatchResult, ) -from litellm.types.litellm_params import RetryStrategy +from litellm.types.litellm_params import ControlOptions, RetryStrategy from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -178,7 +182,7 @@ from litellm.utils import ( from ._logging import verbose_logger from .caching.caching import disable_cache, enable_cache, update_cache -from .litellm_core_utils.core_helpers import safe_deep_copy +from .litellm_core_utils.core_helpers import normalize_drop_params, safe_deep_copy from .litellm_core_utils.fallback_utils import ( async_completion_with_fallbacks, completion_with_fallbacks, @@ -284,7 +288,6 @@ from .types.utils import ( LlmProviders, PromptTokensDetails, ProviderSpecificHeader, - is_litellm_owned_kwarg, ) ####### ENVIRONMENT VARIABLES ################### @@ -335,6 +338,21 @@ ovhcloud_transformation: Final = OVHCloudChatConfig() lemonade_transformation: Final = LemonadeChatConfig() MOCK_RESPONSE_TYPE = str | Exception | dict | ModelResponse | ModelResponseStream + + +def _resolve_control_options(kwargs: Mapping[str, object], model: str) -> ControlOptions: + control: Final = parse_control_options(kwargs) + match control: + case ControlOptions(): + return control + case InvalidControlOption(param=param, message=message): + if litellm.drop_params is True or normalize_drop_params(kwargs.get("drop_params")) is True: + return ControlOptions() + raise litellm.BadRequestError(message=message, model=model, llm_provider=None, body={"param": param}) + case _: + return assert_never(control) + + ####### COMPLETION ENDPOINTS ################ @@ -501,6 +519,7 @@ async def acompletion( loop: Final = asyncio.get_event_loop() custom_llm_provider = kwargs.get("custom_llm_provider", None) + _ = _resolve_control_options(kwargs, model) ## PROMPT MANAGEMENT HOOKS ## ######################################################### @@ -5230,6 +5249,7 @@ def completion( # Responses API config (get_provider_responses_api_config -> None). skip_responses_api_bridge: Final = kwargs.pop("_skip_responses_api_bridge", False) + control_options: Final = _resolve_control_options(kwargs, model) skip_mcp_handler: Final = kwargs.pop("_skip_mcp_handler", False) if not skip_mcp_handler and tools: from litellm.responses.mcp.chat_completions_handler import acompletion_with_mcp @@ -5370,7 +5390,6 @@ def completion( ) ######## end of unpacking kwargs ########### non_default_params: Final = get_non_default_completion_params(kwargs=kwargs) - litellm_params: dict[str, object] = {} # used to prevent unbound var errors ## PROMPT MANAGEMENT HOOKS ## from litellm.integrations.anthropic_cache_control_hook import ( @@ -5622,7 +5641,7 @@ def completion( messages = function_call_prompt(messages=messages, functions=functions_unsupported_model) # For logging - save the values of the litellm-specific params passed in - litellm_params = get_litellm_params( + requested_litellm_params: Final = get_litellm_params( acompletion=acompletion, api_key=api_key, force_timeout=force_timeout, @@ -5670,7 +5689,6 @@ def completion( max_retries=max_retries, timeout=timeout, litellm_request_debug=kwargs.get("litellm_request_debug", False), - stream_chunk_size=kwargs.get("stream_chunk_size"), tpm=kwargs.get("tpm"), rpm=kwargs.get("rpm"), use_xai_oauth=kwargs.get("use_xai_oauth", False), @@ -5683,6 +5701,7 @@ def completion( if key in kwargs }, ) + litellm_params: Final = with_control_options(requested_litellm_params, control_options) if litellm_params.get("provider_affinity_header") is not None: try: headers = add_provider_affinity_header( @@ -6352,9 +6371,7 @@ def embedding( "encoding_format", ] default_params: Final = [*openai_params, "aembedding", "extra_headers"] - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in default_params and not is_litellm_owned_kwarg(k) - } + non_default_params: Final = filter_out_litellm_params(kwargs, excluding=default_params) model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( model=model, diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index f5ba9ebd3da..439858ea2b5 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -5,7 +5,10 @@ from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequenc from dataclasses import dataclass, field, fields, is_dataclass from itertools import chain from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Literal, TypeAlias +from typing import TYPE_CHECKING, Annotated, Final, Literal, TypeAlias + +from pydantic import BeforeValidator, Field +from pydantic.dataclasses import dataclass as pydantic_dataclass if TYPE_CHECKING: import httpx @@ -234,11 +237,31 @@ class ResponseOptions: merge_reasoning_content_in_choices: bool | None = None enable_json_schema_validation: bool | None = None complete_response: bool | None = None - stream_chunk_size: int | None = None keepalive_seconds: float | None = None allow_client_keepalive_override: bool | None = None +MAX_CONTROL_INT_DIGITS: Final = 18 + + +def _int_from_decimal_string(value: object) -> object: + if isinstance(value, str) and value.isascii() and value.isdecimal() and len(value) <= MAX_CONTROL_INT_DIGITS: + return int(value) + return value + + +@pydantic_dataclass(frozen=True, slots=True, kw_only=True) +class ControlOptions: + stream_chunk_size: ( + Annotated[ + int, + BeforeValidator(_int_from_decimal_string), + Field(strict=True, gt=0, lt=10**MAX_CONTROL_INT_DIGITS), + ] + | None + ) = None + + @dataclass(frozen=True, slots=True, kw_only=True) class MockOptions: mock_response: "MockResponse | None" = None @@ -258,6 +281,7 @@ class LiteLLMOptions: guardrails: GuardrailOptions prompt: PromptOptions response: ResponseOptions + control: ControlOptions mock: MockOptions diff --git a/litellm/utils.py b/litellm/utils.py index 092fe936cf9..7ce412e818c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -288,7 +288,7 @@ except (ImportError, AttributeError, TypeError): # Convert to str (if necessary) claude_json_str = json.dumps(json_data) import importlib.metadata -from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence +from collections.abc import AsyncIterator, Callable, Collection, Iterable, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, runtime_checkable from typing_extensions import assert_never @@ -4161,8 +4161,10 @@ def _remove_unsupported_params(non_default_params: dict, supported_openai_params return non_default_params -def filter_out_litellm_params(kwargs: Mapping[str, object]) -> dict: - return {key: value for key, value in kwargs.items() if not is_litellm_owned_kwarg(key)} +def filter_out_litellm_params( + kwargs: Mapping[str, object], excluding: Collection[str] = frozenset() +) -> dict[str, object]: + return {key: value for key, value in kwargs.items() if key not in excluding and not is_litellm_owned_kwarg(key)} def _provider_supports_vertex_params(custom_llm_provider: str) -> bool: @@ -10132,13 +10134,8 @@ def get_standard_openai_params(params: Mapping[str, object]) -> dict: return {k: v for k, v in params.items() if k in litellm.OPENAI_CHAT_COMPLETION_PARAMS and v is not None} -def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict: - openai_params: Final = litellm.OPENAI_CHAT_COMPLETION_PARAMS - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k) - } - - return non_default_params +def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict[str, object]: + return filter_out_litellm_params(kwargs, excluding=litellm.OPENAI_CHAT_COMPLETION_PARAMS) def peek_reasoning_summary_aliases(optional_params: dict) -> object | None: @@ -10184,13 +10181,10 @@ def strip_reasoning_summary_aliases_from_optional_params( return op, rs_val -def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict: +def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict[str, object]: from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in OPENAI_TRANSCRIPTION_PARAMS and not is_litellm_owned_kwarg(k) - } - return non_default_params + return filter_out_litellm_params(kwargs, excluding=OPENAI_TRANSCRIPTION_PARAMS) def add_openai_metadata( diff --git a/tests/_support/stream_chunk_size.py b/tests/_support/stream_chunk_size.py index 051f552e282..6e6256637f0 100644 --- a/tests/_support/stream_chunk_size.py +++ b/tests/_support/stream_chunk_size.py @@ -1,26 +1,28 @@ from collections.abc import Mapping +from types import MappingProxyType from typing import Final -import litellm import pytest -from litellm.integrations.custom_logger import CustomLogger +from litellm.constants import CONTROL_OPTIONS_KEY +from litellm.types.litellm_params import ControlOptions -class LitellmParamsRecorder(CustomLogger): - def __init__(self) -> None: - super().__init__() - self.seen: tuple[Mapping[str, object], ...] = () +DEFAULT_CHUNKING_REQUESTS: Final = ( + pytest.param(MappingProxyType({}), id="unset"), + pytest.param(MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": True}), id="dropped"), + pytest.param( + MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": "true"}), id="dropped_by_string_flag" + ), + pytest.param(MappingProxyType({CONTROL_OPTIONS_KEY: ControlOptions(stream_chunk_size=1)}), id="forged_options"), + pytest.param(MappingProxyType({CONTROL_OPTIONS_KEY: {"stream_chunk_size": 1}}), id="forged_mapping"), +) - def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None: - params: Final = kwargs["litellm_params"] - assert isinstance(params, Mapping) - self.seen = (*self.seen, params) - - -def record_litellm_params(monkeypatch: pytest.MonkeyPatch) -> LitellmParamsRecorder: - recorder: Final = LitellmParamsRecorder() - monkeypatch.setattr(litellm, "input_callback", [recorder]) - return recorder +ROUTER_CHUNK_SIZE_CASES: Final = ( + pytest.param(MappingProxyType({"stream_chunk_size": 64}), 64, id="int"), + pytest.param(MappingProxyType({"stream_chunk_size": "64"}), 64, id="digit_string"), + pytest.param(MappingProxyType({}), None, id="unset"), + pytest.param(MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": True}), None, id="dropped"), +) def keys_at_every_depth(value: object) -> frozenset[str]: diff --git a/tests/integration/providers/test_internal_params_wire.py b/tests/integration/providers/test_internal_params_wire.py index 17b0fc9d815..f9f5d1e5478 100644 --- a/tests/integration/providers/test_internal_params_wire.py +++ b/tests/integration/providers/test_internal_params_wire.py @@ -8,11 +8,12 @@ from collections.abc import Callable, Mapping from pathlib import Path from typing import Final -import litellm import pytest from integration._support.upstream import INTERNAL_FIELDS from integration._support.wire import Reply, Request, wire_server -from tests._support.stream_chunk_size import keys_at_every_depth, record_litellm_params + +import litellm +from tests._support.stream_chunk_size import keys_at_every_depth TEXT: Final = "wire control" OPENAI_RESPONSE: Final = { @@ -277,13 +278,11 @@ def provider_wire_environment(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) - @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("stream", [False, True]) async def test_internal_params_never_reach_provider_body( - monkeypatch: pytest.MonkeyPatch, provider_wire_environment: None, provider: str, asynchronous: bool, stream: bool, ) -> None: - recorder: Final = record_litellm_params(monkeypatch) with wire_server(_peer(provider)) as wire: parameters: Final = { **_request_parameters(provider, wire.url), @@ -308,8 +307,6 @@ async def test_internal_params_never_reach_provider_body( assert result.choices[0].message.content == TEXT requests: Final = wire.drain() assert len(requests) == 1 - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 body: Final = json.loads(requests[0].body) keys: Final = keys_at_every_depth(body) assert "stream_chunk_size" not in keys diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 2e4122530d8..0e0f2b7eac6 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -1,10 +1,12 @@ import asyncio import logging import re +from typing import Final from unittest.mock import MagicMock import pytest +import litellm import litellm.caching.redis_cache as redis_cache_module from litellm.caching.caching import Cache from litellm.caching.caching_handler import _PENDING_CACHE_WRITES @@ -389,3 +391,15 @@ async def test_embedding_cache_serves_base64_string_embeddings_on_repeat(monkeyp assert embedder.provider_calls == 1, "a string embedding written to the cache must be served on repeat" assert [item["embedding"] for item in second.data] == [item["embedding"] for item in first.data] == ["AACAPwAAAEA="] + + +def test_provider_specific_cache_key_ignores_litellm_owned_kwargs(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "enable_caching_on_provider_specific_optional_params", True) + cache: Final = Cache(type=LiteLLMCacheType.LOCAL) + request: Final = {"model": "gpt-4.1-mini", "messages": [{"role": "user", "content": "hi"}], "top_k": 5} + + base_key: Final = cache.get_cache_key(**request) + + assert cache.get_cache_key(**request, _litellm_control={"stream_chunk_size": 64}) == base_key + assert cache.get_cache_key(**request, litellm_trace_id="trace-1") == base_key + assert cache.get_cache_key(**{**request, "top_k": 6}) != base_key diff --git a/tests/unit/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py index 39bc2688ae0..9b5771092ac 100644 --- a/tests/unit/litellm_core_utils/test_get_litellm_params.py +++ b/tests/unit/litellm_core_utils/test_get_litellm_params.py @@ -7,12 +7,18 @@ Ensures backward compatibility after sparse kwargs extraction optimization. from typing import Final import pytest +from pydantic import ValidationError +from litellm.constants import CONTROL_OPTIONS_KEY from litellm.litellm_core_utils.get_litellm_params import ( _OPTIONAL_KWARGS_KEYS, + InvalidControlOption, _get_base_model_from_litellm_call_metadata, get_litellm_params, + parse_control_options, + stored_control_options, ) +from litellm.types.litellm_params import ControlOptions NAMED_PRICE_PARAMS: Final = frozenset( {"input_cost_per_token", "output_cost_per_token", "input_cost_per_second", "output_cost_per_second"} @@ -90,9 +96,8 @@ class TestGetLitellmParamsKwargsExtraction: assert "s3_endpoint_url" not in result_without_s3_kwargs assert "s3_region_name" not in result_without_s3_kwargs - def test_stream_chunk_size_is_carried_as_a_litellm_param(self) -> None: - assert get_litellm_params(stream_chunk_size=64)["stream_chunk_size"] == 64 - assert get_litellm_params()["stream_chunk_size"] is None + def test_a_caller_supplied_control_options_key_is_not_carried(self) -> None: + assert CONTROL_OPTIONS_KEY not in get_litellm_params(**{CONTROL_OPTIONS_KEY: {"stream_chunk_size": 64}}) def test_s3_credential_kwargs_are_forwarded_for_s3_signing(self): result = get_litellm_params(s3_access_key_id="s3-key", s3_secret_access_key="s3-secret") @@ -122,6 +127,79 @@ class TestGetLitellmParamsKwargsExtraction: assert result[key] == f"val_{key}" +@pytest.mark.parametrize( + "kwargs,expected", + [ + ({"stream_chunk_size": 64, "temperature": 0.2}, ControlOptions(stream_chunk_size=64)), + ({"stream_chunk_size": "64"}, ControlOptions(stream_chunk_size=64)), + ({"stream_chunk_size": None}, ControlOptions()), + ({"temperature": 0.2}, ControlOptions()), + ], +) +def test_control_options_are_read_from_the_request_kwargs(kwargs: dict[str, object], expected: ControlOptions) -> None: + assert parse_control_options(kwargs) == expected + + +@pytest.mark.parametrize( + "raw,shown", + [ + ("sixty-four", "'sixty-four'"), + (" 64", "' 64'"), + ("-1", "'-1'"), + ("\uff16\uff14", "'\uff16\uff14'"), + ("x" * 500, "'xxxxxxxxxxxx...xxxxxxxxxxxxx'"), + pytest.param(-(10**5000), "", id="huge_negative_int"), + pytest.param(-(2**64 - 1), "-18446744073709551615", id="64_bit_negative_int"), + pytest.param(-(2**64), "", id="65_bit_negative_int"), + pytest.param([-(10**5000)], "[]", id="nested_huge_int"), + pytest.param(10**18, "1000000000000000000", id="19_digit_int"), + pytest.param("1" + "0" * 18, "'1000000000000000000'", id="19_digit_string"), + pytest.param("9" * 5000, "'999999999999...9999999999999'", id="5000_digit_string"), + pytest.param("0" * 18 + "1", "'0000000000000000001'", id="19_digit_string_with_leading_zeros"), + (64.0, "64.0"), + (True, "True"), + (0, "0"), + ("0", "'0'"), + (-1, "-1"), + ], +) +def test_control_options_reject_a_stream_chunk_size_that_is_not_a_positive_int(raw: object, shown: str) -> None: + assert parse_control_options({"stream_chunk_size": raw}) == InvalidControlOption( + param="stream_chunk_size", + message=f"Invalid stream_chunk_size={shown}: expected a positive integer of at most 18 digits", + ) + + +@pytest.mark.parametrize("raw", [10**18 - 1, "9" * 18], ids=["int", "digit_string"]) +def test_control_options_accept_the_largest_18_digit_value(raw: object) -> None: + assert parse_control_options({"stream_chunk_size": raw}) == ControlOptions(stream_chunk_size=10**18 - 1) + + +def test_control_options_accept_an_18_digit_string_with_leading_zeros() -> None: + assert parse_control_options({"stream_chunk_size": "0" * 17 + "1"}) == ControlOptions(stream_chunk_size=1) + + +@pytest.mark.parametrize("raw", [0, -1, "sixty-four", 64.0, True]) +def test_control_options_enforce_their_rule_at_construction(raw: object) -> None: + with pytest.raises(ValidationError): + ControlOptions(stream_chunk_size=raw) # pyright: ignore[reportArgumentType] # the invalid type is the input + + +@pytest.mark.parametrize( + "litellm_params,expected", + [ + ({CONTROL_OPTIONS_KEY: ControlOptions(stream_chunk_size=64)}, ControlOptions(stream_chunk_size=64)), + ({}, ControlOptions()), + ({CONTROL_OPTIONS_KEY: {"stream_chunk_size": 64}}, ControlOptions()), + ({"stream_chunk_size": 64}, ControlOptions()), + ], +) +def test_stored_control_options_reads_only_the_validated_options( + litellm_params: dict[str, object], expected: ControlOptions +) -> None: + assert stored_control_options(litellm_params) == expected + + class TestGetLitellmParamsBaseModel: """Verify base_model resolution precedence.""" diff --git a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py index ed172fdfbff..d1748e1b38d 100644 --- a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py +++ b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py @@ -1,4 +1,6 @@ import json +from collections.abc import Mapping +from types import MappingProxyType from typing import Final from unittest.mock import AsyncMock, MagicMock @@ -6,44 +8,40 @@ import httpx import pytest import litellm -from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( - AmazonAnthropicClaudeConfig, -) from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, ) from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler -from tests._support.stream_chunk_size import ( - LitellmParamsRecorder, - keys_at_every_depth, - record_litellm_params, -) +from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth @pytest.mark.parametrize( - "config,model", + "model", [ - (AmazonInvokeConfig, "anthropic.claude-3-sonnet-20240229-v1:0"), - (AmazonInvokeConfig, "amazon.titan-text-express-v1"), - (AmazonInvokeConfig, "mistral.mistral-7b-instruct-v0:2"), - (AmazonAnthropicClaudeConfig, "anthropic.claude-sonnet-4-6"), + "anthropic.claude-sonnet-4-6", + "amazon.titan-text-express-v1", + "mistral.mistral-7b-instruct-v0:2", ], ) -def test_transform_request_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( - model=model, +def test_completion_keeps_stream_chunk_size_out_of_invoke_bodies(model: str) -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) + + litellm.completion( + model=f"bedrock/invoke/{model}", messages=[{"role": "user", "content": "hi"}], - optional_params={"stream": True, "stream_chunk_size": 2048, "max_tokens": 10}, - litellm_params={}, - headers={}, + stream=True, + max_tokens=10, + client=client, + aws_access_key_id="fake", + aws_secret_access_key="fake", + aws_region_name="us-east-1", + stream_chunk_size=2048, ) - assert "stream_chunk_size" not in json.dumps(request_body) + request: Final = send.call_args.args[0] + assert "stream_chunk_size" not in keys_at_every_depth(json.loads(request.content)), request.content def test_validate_environment_maps_guardrail_config_to_invoke_headers(): @@ -243,10 +241,7 @@ def test_transform_response_hands_json_mode_to_nova(): assert json.loads(result.choices[0].message.content) == {"city": "Paris", "temperature": 21} -def _stream_invoke_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, **kwargs -) -> tuple[MagicMock, MagicMock, LitellmParamsRecorder]: - recorder: Final = record_litellm_params(monkeypatch) +def _stream_invoke_completion_with_spied_client(**kwargs: object) -> tuple[MagicMock, MagicMock]: mock_response = MagicMock() mock_response.status_code = 200 mock_response.iter_bytes = MagicMock(return_value=iter([])) @@ -263,39 +258,33 @@ def _stream_invoke_completion_with_spied_client( aws_region_name="us-east-1", **kwargs, ) - return mock_response.iter_bytes, client.post, recorder + return mock_response.iter_bytes, client.post -def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_invoke_body( - monkeypatch: pytest.MonkeyPatch, -): - iter_bytes_spy, post_spy, recorder = _stream_invoke_completion_with_spied_client(monkeypatch, stream_chunk_size=64) +def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_invoke_body() -> None: + iter_bytes_spy, post_spy = _stream_invoke_completion_with_spied_client(stream_chunk_size=64) iter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 -def test_completion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch): - iter_bytes_spy, _, recorder = _stream_invoke_completion_with_spied_client(monkeypatch) +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +def test_completion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], +) -> None: + iter_bytes_spy, _ = _stream_invoke_completion_with_spied_client(**request_kwargs) iter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -async def _astream_invoke_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, **kwargs -) -> tuple[MagicMock, AsyncMock, LitellmParamsRecorder]: +async def _astream_invoke_completion_with_spied_client(**kwargs: object) -> tuple[MagicMock, AsyncMock]: async def _no_bytes(): return yield b"" mock_response = MagicMock() mock_response.status_code = 200 - recorder: Final = record_litellm_params(monkeypatch) mock_response.aiter_bytes = MagicMock(return_value=_no_bytes()) aiter_bytes_spy = mock_response.aiter_bytes client = AsyncHTTPHandler() @@ -311,57 +300,49 @@ async def _astream_invoke_completion_with_spied_client( aws_region_name="us-east-1", **kwargs, ) - return aiter_bytes_spy, client.post, recorder + return aiter_bytes_spy, client.post @pytest.mark.asyncio -async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_invoke_body( - monkeypatch: pytest.MonkeyPatch, -): - aiter_bytes_spy, post_spy, recorder = await _astream_invoke_completion_with_spied_client( - monkeypatch, stream_chunk_size=64 - ) +async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_invoke_body() -> None: + aiter_bytes_spy, post_spy = await _astream_invoke_completion_with_spied_client(stream_chunk_size=64) aiter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 @pytest.mark.asyncio -async def test_acompletion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch): - aiter_bytes_spy, _, recorder = await _astream_invoke_completion_with_spied_client(monkeypatch) +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +async def test_acompletion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], +) -> None: + aiter_bytes_spy, _ = await _astream_invoke_completion_with_spied_client(**request_kwargs) aiter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -@pytest.mark.parametrize("stream_chunk_size,expected_chunk_size", [(64, 64), (None, None)]) -def test_router_deployment_stream_chunk_size_reaches_iter_bytes( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size, expected_chunk_size -): - recorder: Final = record_litellm_params(monkeypatch) - 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) - deployment_params = { +INVOKE_DEPLOYMENT: Final = MappingProxyType( + { "model": "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", "aws_access_key_id": "fake", "aws_secret_access_key": "fake", "aws_region_name": "us-east-1", } - router = litellm.Router( - model_list=[ - { - "model_name": "invoke-chunked", - "litellm_params": deployment_params - | ({} if stream_chunk_size is None else {"stream_chunk_size": stream_chunk_size}), - } - ] +) + + +@pytest.mark.parametrize("deployment_extras,expected_chunk_size", ROUTER_CHUNK_SIZE_CASES) +def test_router_deployment_stream_chunk_size_reaches_iter_bytes( + deployment_extras: Mapping[str, object], expected_chunk_size: int | None +) -> None: + mock_response: Final = MagicMock() + mock_response.status_code = 200 + mock_response.iter_bytes = MagicMock(return_value=iter([])) + client: Final = HTTPHandler() + client.post = MagicMock(return_value=mock_response) + router: Final = litellm.Router( + model_list=[{"model_name": "invoke-chunked", "litellm_params": {**INVOKE_DEPLOYMENT, **deployment_extras}}] ) router.completion( @@ -374,17 +355,11 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes( mock_response.iter_bytes.assert_called_once_with(chunk_size=expected_chunk_size) data: Final = client.post.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == stream_chunk_size -def test_stream_wrapper_rejects_non_int_stream_chunk_size(monkeypatch: pytest.MonkeyPatch): - record_litellm_params(monkeypatch) - 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) +def test_invoke_stream_rejects_non_int_stream_chunk_size_before_calling_bedrock() -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) with pytest.raises(litellm.BadRequestError): litellm.completion( @@ -398,4 +373,28 @@ def test_stream_wrapper_rejects_non_int_stream_chunk_size(monkeypatch: pytest.Mo stream_chunk_size="sixty-four", ) - client.post.assert_not_called() + send.assert_not_called() + + +def test_router_deployment_with_a_non_numeric_stream_chunk_size_gets_a_400_before_calling_bedrock() -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "invoke-chunked", + "litellm_params": {**INVOKE_DEPLOYMENT, "stream_chunk_size": "sixty-four"}, + } + ] + ) + + with pytest.raises(litellm.BadRequestError) as exc_info: + router.completion( + model="invoke-chunked", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + ) + + assert exc_info.value.status_code == 400 + send.assert_not_called() diff --git a/tests/unit/llms/bedrock/test_common_utils.py b/tests/unit/llms/bedrock/test_common_utils.py deleted file mode 100644 index cfcc15f186b..00000000000 --- a/tests/unit/llms/bedrock/test_common_utils.py +++ /dev/null @@ -1,20 +0,0 @@ -import pytest - -from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from - - -def test_stream_chunk_size_from_absent_is_none(): - assert stream_chunk_size_from({}) is None - - -def test_stream_chunk_size_from_int_is_returned(): - assert stream_chunk_size_from({"stream_chunk_size": 64}) == 64 - - -@pytest.mark.parametrize("bad_value", ["64", 6.4, True]) -def test_stream_chunk_size_from_rejects_non_int_with_400(bad_value): - with pytest.raises(BedrockError) as excinfo: - stream_chunk_size_from({"stream_chunk_size": bad_value}) - - assert excinfo.value.status_code == 400 - assert repr(bad_value) in excinfo.value.message diff --git a/tests/unit/llms/chat/test_converse_handler.py b/tests/unit/llms/chat/test_converse_handler.py index cbb8e3acf78..57bb9ab771f 100644 --- a/tests/unit/llms/chat/test_converse_handler.py +++ b/tests/unit/llms/chat/test_converse_handler.py @@ -1,5 +1,6 @@ import json -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Mapping +from types import MappingProxyType from typing import Final from unittest.mock import AsyncMock, MagicMock @@ -11,11 +12,7 @@ from litellm.llms.bedrock.chat import BedrockConverseLLM from litellm.llms.bedrock.chat.converse_handler import make_sync_call from litellm.llms.bedrock.common_utils import _get_all_bedrock_regions from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler -from tests._support.stream_chunk_size import ( - LitellmParamsRecorder, - keys_at_every_depth, - record_litellm_params, -) +from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth def test_encode_model_id_with_inference_profile(): @@ -319,10 +316,7 @@ def test_completion_plumbs_stream_chunk_size_through_converse() -> None: iter_bytes_spy.assert_called_once_with(chunk_size=2048) -def _stream_converse_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None = None -) -> tuple[MagicMock, MagicMock, LitellmParamsRecorder]: - recorder: Final = record_litellm_params(monkeypatch) +def _stream_converse_completion_with_spied_client(**request: object) -> tuple[MagicMock, MagicMock]: mock_response: Final = MagicMock() mock_response.status_code = 200 mock_response.iter_bytes = MagicMock(return_value=iter([])) @@ -337,43 +331,35 @@ def _stream_converse_completion_with_spied_client( aws_access_key_id="fake", aws_secret_access_key="fake", aws_region_name="us-east-1", - stream_chunk_size=stream_chunk_size, + **request, ) - return mock_response.iter_bytes, client.post, recorder + return mock_response.iter_bytes, client.post -def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_converse_body( - monkeypatch: pytest.MonkeyPatch, -) -> None: - iter_bytes_spy, post_spy, recorder = _stream_converse_completion_with_spied_client( - monkeypatch, stream_chunk_size=64 - ) +def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_converse_body() -> None: + iter_bytes_spy, post_spy = _stream_converse_completion_with_spied_client(stream_chunk_size=64) iter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 -def test_completion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch) -> None: - iter_bytes_spy, _, recorder = _stream_converse_completion_with_spied_client(monkeypatch) +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +def test_completion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], +) -> None: + iter_bytes_spy, _ = _stream_converse_completion_with_spied_client(**request_kwargs) iter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -async def _astream_converse_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None = None -) -> tuple[MagicMock, AsyncMock, LitellmParamsRecorder]: +async def _astream_converse_completion_with_spied_client(**request: object) -> tuple[MagicMock, AsyncMock]: async def _no_bytes(chunk_size: int | None = None) -> AsyncIterator[bytes]: return yield b"" mock_response: Final = MagicMock() mock_response.status_code = 200 - recorder: Final = record_litellm_params(monkeypatch) mock_response.aiter_bytes = MagicMock(return_value=_no_bytes()) aiter_bytes_spy: Final = mock_response.aiter_bytes client: Final = AsyncHTTPHandler() @@ -387,61 +373,51 @@ async def _astream_converse_completion_with_spied_client( aws_access_key_id="fake", aws_secret_access_key="fake", aws_region_name="us-east-1", - stream_chunk_size=stream_chunk_size, + **request, ) - return aiter_bytes_spy, client.post, recorder + return aiter_bytes_spy, client.post @pytest.mark.asyncio -async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_converse_body( - monkeypatch: pytest.MonkeyPatch, -) -> None: - aiter_bytes_spy, post_spy, recorder = await _astream_converse_completion_with_spied_client( - monkeypatch, stream_chunk_size=64 - ) +async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_converse_body() -> None: + aiter_bytes_spy, post_spy = await _astream_converse_completion_with_spied_client(stream_chunk_size=64) aiter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 @pytest.mark.asyncio -async def test_acompletion_without_stream_chunk_size_uses_default_chunking( - monkeypatch: pytest.MonkeyPatch, +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +async def test_acompletion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], ) -> None: - aiter_bytes_spy, _, recorder = await _astream_converse_completion_with_spied_client(monkeypatch) + aiter_bytes_spy, _ = await _astream_converse_completion_with_spied_client(**request_kwargs) aiter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -@pytest.mark.parametrize("stream_chunk_size,expected_chunk_size", [(64, 64), (None, None)]) -def test_router_deployment_stream_chunk_size_reaches_iter_bytes( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None, expected_chunk_size: int | None -) -> None: - recorder: Final = record_litellm_params(monkeypatch) - mock_response: Final = MagicMock() - mock_response.status_code = 200 - mock_response.iter_bytes = MagicMock(return_value=iter([])) - client: Final = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - deployment_params: Final = { +CONVERSE_DEPLOYMENT: Final = MappingProxyType( + { "model": "bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0", "aws_access_key_id": "fake", "aws_secret_access_key": "fake", "aws_region_name": "us-east-1", } +) + + +@pytest.mark.parametrize("deployment_extras,expected_chunk_size", ROUTER_CHUNK_SIZE_CASES) +def test_router_deployment_stream_chunk_size_reaches_iter_bytes( + deployment_extras: Mapping[str, object], expected_chunk_size: int | None +) -> None: + mock_response: Final = MagicMock() + mock_response.status_code = 200 + mock_response.iter_bytes = MagicMock(return_value=iter([])) + client: Final = HTTPHandler() + client.post = MagicMock(return_value=mock_response) router: Final = litellm.Router( - model_list=[ - { - "model_name": "converse-chunked", - "litellm_params": deployment_params - | ({} if stream_chunk_size is None else {"stream_chunk_size": stream_chunk_size}), - } - ] + model_list=[{"model_name": "converse-chunked", "litellm_params": {**CONVERSE_DEPLOYMENT, **deployment_extras}}] ) router.completion( @@ -454,20 +430,18 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes( mock_response.iter_bytes.assert_called_once_with(chunk_size=expected_chunk_size) data: Final = client.post.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == stream_chunk_size -def test_converse_stream_rejects_non_int_stream_chunk_size_before_calling_bedrock(monkeypatch: pytest.MonkeyPatch): - record_litellm_params(monkeypatch) - client = HTTPHandler() - client.post = MagicMock() +@pytest.mark.parametrize("stream", [True, False], ids=["stream", "non_stream"]) +def test_converse_rejects_non_int_stream_chunk_size_before_calling_bedrock(stream: bool) -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) with pytest.raises(litellm.BadRequestError): litellm.completion( model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0", messages=[{"role": "user", "content": "hi"}], - stream=True, + stream=stream, client=client, aws_access_key_id="fake", aws_secret_access_key="fake", @@ -475,30 +449,7 @@ def test_converse_stream_rejects_non_int_stream_chunk_size_before_calling_bedroc stream_chunk_size="sixty-four", ) - client.post.assert_not_called() - - -def test_converse_non_stream_ignores_invalid_stream_chunk_size(): - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.json = MagicMock(return_value=_converse_response_body()) - mock_response.text = json.dumps(_converse_response_body()) - mock_response.headers = httpx.Headers() - client = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - - response = litellm.completion( - model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0", - messages=[{"role": "user", "content": "hi"}], - client=client, - aws_access_key_id="fake", - aws_secret_access_key="fake", - aws_region_name="us-east-1", - stream_chunk_size="64", - ) - - assert response.choices[0].message.content == "hi" - client.post.assert_called_once() + send.assert_not_called() def _bedrock_error_response(status_code: int, request_id: str) -> httpx.Response: diff --git a/tests/unit/llms/oci/chat/test_oci_chat_transformation.py b/tests/unit/llms/oci/chat/test_oci_chat_transformation.py index 708187b8ae1..462b1d6ea72 100644 --- a/tests/unit/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/unit/llms/oci/chat/test_oci_chat_transformation.py @@ -1247,6 +1247,7 @@ class TestOCIStreamingSignedBody: mock_logging = MagicMock() config.get_sync_custom_stream_wrapper( + litellm_params={}, api_base="https://example.com", headers={}, data={"key": "value"}, @@ -1286,6 +1287,7 @@ class TestOCIStreamingSignedBody: payload = {"key": "value"} config.get_sync_custom_stream_wrapper( + litellm_params={}, api_base="https://example.com", headers={}, data=payload, diff --git a/tests/unit/llms/oci/test_oci_coverage_boost.py b/tests/unit/llms/oci/test_oci_coverage_boost.py index 7c91ece70b5..8f7588c5de7 100644 --- a/tests/unit/llms/oci/test_oci_coverage_boost.py +++ b/tests/unit/llms/oci/test_oci_coverage_boost.py @@ -1111,6 +1111,7 @@ def test_get_sync_custom_stream_wrapper_returns_wrapper(): mock_client.post.return_value = mock_response wrapper = config.get_sync_custom_stream_wrapper( + litellm_params={}, model=_GENERIC_MODEL, custom_llm_provider="oci", logging_obj=MagicMock(), @@ -1143,6 +1144,7 @@ async def test_get_async_custom_stream_wrapper_returns_wrapper(): mock_client.post = AsyncMock(return_value=mock_response) wrapper = await config.get_async_custom_stream_wrapper( + litellm_params={}, model=_GENERIC_MODEL, custom_llm_provider="oci", logging_obj=MagicMock(), diff --git a/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py index 697f5a7ff59..cb20b3390bb 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py @@ -125,6 +125,7 @@ def test_sync_first_event_emitted_after_a_single_frame(): response = httpx.Response(200, stream=stream) wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper( + litellm_params={}, model="phi-4", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), @@ -147,6 +148,7 @@ def test_sync_events_emitted_incrementally_without_bursting(): response = httpx.Response(200, stream=stream) wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper( + litellm_params={}, model="phi-4", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), @@ -171,6 +173,7 @@ async def test_async_first_event_emitted_after_a_single_frame(): response = httpx.Response(200, stream=stream) wrapper = await SagemakerChatConfig().get_async_custom_stream_wrapper( + litellm_params={}, model="phi-4", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), diff --git a/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py index 5cc414819e3..d878bc70a09 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py @@ -309,6 +309,7 @@ class TestSagemakerChatBackwardsCompatibility: ) as mock_csw: mock_csw.return_value = MagicMock() self.config.get_sync_custom_stream_wrapper( + litellm_params={}, model="my-hf-endpoint", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), @@ -348,6 +349,7 @@ class TestSagemakerChatBackwardsCompatibility: mock_csw.return_value = MagicMock() asyncio.run( self.config.get_async_custom_stream_wrapper( + litellm_params={}, model="my-hf-endpoint", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), diff --git a/tests/unit/responses/test_responses_api_bridge_flag.py b/tests/unit/responses/test_responses_api_bridge_flag.py index 642495fab86..fb1361c1f49 100644 --- a/tests/unit/responses/test_responses_api_bridge_flag.py +++ b/tests/unit/responses/test_responses_api_bridge_flag.py @@ -12,6 +12,7 @@ from typing import Final from unittest.mock import MagicMock, patch import httpx +import openai import pytest import respx @@ -592,3 +593,20 @@ class TestUseResponsesApiBridgeFlag: mock_native_handler.assert_called_once() assert result is not None + + def test_bridge_still_rejects_an_invalid_stream_chunk_size(self) -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = openai.OpenAI(api_key="fake-key", http_client=httpx.Client(transport=httpx.MockTransport(send))) + + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.responses( + model="openai/gpt-4.1-mini", + input="hi", + use_chat_completions_api=True, + stream_chunk_size="sixty-four", + client=client, + num_retries=0, + ) + + assert exc_info.value.param == "stream_chunk_size" + send.assert_not_called() diff --git a/tests/unit/test_filter_out_litellm_params.py b/tests/unit/test_filter_out_litellm_params.py index 72f8f5f1478..342e251d611 100644 --- a/tests/unit/test_filter_out_litellm_params.py +++ b/tests/unit/test_filter_out_litellm_params.py @@ -2,6 +2,10 @@ Test filter_out_litellm_params helper function. """ +from typing import Final + + +import litellm from litellm.utils import filter_out_litellm_params @@ -34,3 +38,19 @@ def test_filter_out_litellm_params(): assert "litellm_trace_id" not in filtered assert "proxy_server_request" not in filtered assert "secret_fields" not in filtered + + +def test_filter_out_litellm_params_also_drops_the_excluded_names(): + kwargs = {"temperature": 0.2, "top_k": 5, "litellm_trace_id": "trace-1", "_litellm_control": object()} + + assert filter_out_litellm_params(kwargs, excluding=("temperature",)) == {"top_k": 5} + + +def test_filter_out_litellm_params_sees_a_name_appended_to_the_public_list_after_import(): + litellm.all_litellm_params.append("registered_later") + try: + filtered: Final = filter_out_litellm_params({"registered_later": 1, "top_k": 2}) + finally: + litellm.all_litellm_params.remove("registered_later") + + assert filtered == {"top_k": 2} diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index 7bef35d8559..e0e1fcfe105 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -22,10 +22,16 @@ from unittest.mock import MagicMock, patch import litellm from litellm import main as litellm_main +from litellm.constants import CONTROL_OPTIONS_KEY from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs +from litellm.litellm_core_utils.get_litellm_params import stored_control_options from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage +from litellm.types.litellm_params import ControlOptions +from litellm.types.llms.openai import AllMessageValues +from litellm.types.prompts.init_prompts import PromptSpec +from litellm.types.utils import Delta, ModelResponseStream, StandardCallbackDynamicParams, StreamingChoices, Usage @pytest.fixture(autouse=True) @@ -4273,3 +4279,228 @@ def test_completion_rejects_untranslatable_tool_choice_with_a_400(tool_choice): ) assert exc_info.value.status_code == 400 assert f"tool_choice={tool_choice}" in str(exc_info.value) + + +@pytest.mark.parametrize("raw", ["sixty-four", 0, -1]) +def test_completion_rejects_an_invalid_stream_chunk_size_with_a_400_naming_the_param(raw: object) -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size=raw, + mock_response="unused", + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.param == "stream_chunk_size" + assert f"Invalid stream_chunk_size={raw!r}: expected a positive integer of at most 18 digits" in str(exc_info.value) + + +class _PromptHookRecorder(CustomPromptManagement): + def __init__(self, on_prompt: MagicMock) -> None: + super().__init__() + self.on_prompt: Final = on_prompt + + def get_chat_completion_prompt( + self, + model: str, + messages: list[AllMessageValues], + non_default_params: dict, + prompt_id: str | None, + prompt_variables: dict | None, + dynamic_callback_params: StandardCallbackDynamicParams, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: + self.on_prompt("sync") + return model, messages, non_default_params + + async def async_get_chat_completion_prompt( + self, + model: str, + messages: list[AllMessageValues], + non_default_params: dict, + prompt_id: str | None, + prompt_variables: dict | None, + dynamic_callback_params: StandardCallbackDynamicParams, + litellm_logging_obj: LiteLLMLogging, + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: + self.on_prompt("async") + return model, messages, non_default_params + + +async def _call_completion(is_async: bool, **kwargs: object) -> None: + if is_async: + await litellm.acompletion(**kwargs) + else: + litellm.completion(**kwargs) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async,hook", [(False, "sync"), (True, "async")], ids=["completion", "acompletion"]) +async def test_the_prompt_hook_runs_when_stream_chunk_size_is_valid( + monkeypatch: pytest.MonkeyPatch, is_async: bool, hook: str +) -> None: + on_prompt: Final = MagicMock() + monkeypatch.setattr(litellm, "callbacks", [_PromptHookRecorder(on_prompt)]) + + await _call_completion( + is_async, + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + prompt_id="greeting", + stream_chunk_size=64, + mock_response="hi", + ) + + on_prompt.assert_any_call(hook) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", [False, True], ids=["completion", "acompletion"]) +async def test_an_invalid_stream_chunk_size_is_rejected_before_any_prompt_hook_runs( + monkeypatch: pytest.MonkeyPatch, is_async: bool +) -> None: + on_prompt: Final = MagicMock() + monkeypatch.setattr(litellm, "callbacks", [_PromptHookRecorder(on_prompt)]) + + with pytest.raises(litellm.BadRequestError): + await _call_completion( + is_async, + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + prompt_id="greeting", + stream_chunk_size="sixty-four", + mock_response="hi", + ) + + on_prompt.assert_not_called() + + +def _completion_logging_obj(call_id: str) -> LiteLLMLogging: + return LiteLLMLogging( + model="gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="completion", + start_time=datetime(2026, 1, 1), + litellm_call_id=call_id, + function_id=f"{call_id}-function", + ) + + +def test_completion_carries_the_control_options_into_the_logged_litellm_params() -> None: + logging_obj: Final = _completion_logging_obj("control-params") + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size=64, + mock_response="hi", + litellm_logging_obj=logging_obj, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions(stream_chunk_size=64) + + +def test_completion_ignores_a_caller_supplied_control_options_key() -> None: + logging_obj: Final = _completion_logging_obj("control-params-injection") + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + mock_response="hi", + litellm_logging_obj=logging_obj, + **{CONTROL_OPTIONS_KEY: {"stream_chunk_size": 1}}, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions() + + +@pytest.mark.parametrize("drop_params", [True, "true"]) +def test_drop_params_drops_an_invalid_stream_chunk_size_instead_of_rejecting_it(drop_params: object) -> None: + logging_obj: Final = _completion_logging_obj(f"drop-params-{drop_params}") + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size="sixty-four", + drop_params=drop_params, + mock_response="hi", + litellm_logging_obj=logging_obj, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions() + + +def test_drop_params_keeps_a_dropped_stream_chunk_size_out_of_the_provider_request( + respx_mock: respx.MockRouter, +) -> None: + api_base: Final = "http://localhost:12346/v1" + mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/chat/completions.*").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-drop", + "object": "chat.completion", + "created": 1712697600, + "model": "gpt-4.1-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + ) + ) + + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + api_base=api_base, + api_key="fake_openai_api_key", + stream_chunk_size="sixty-four", + drop_params=True, + ) + + assert mock_route.called + sent: Final = json.loads(respx_mock.calls[0].request.content) + assert "stream_chunk_size" not in sent, sent + assert sent["model"] == "gpt-4.1-mini" + + +@pytest.mark.asyncio +async def test_global_drop_params_drops_an_invalid_stream_chunk_size_on_acompletion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "drop_params", True) + logging_obj: Final = _completion_logging_obj("global-drop-params") + await litellm.acompletion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size=0, + mock_response="hi", + litellm_logging_obj=logging_obj, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions() + + +def test_completion_rejects_an_invalid_stream_chunk_size_before_the_mcp_gateway() -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + tools=[{"type": "mcp", "server_label": "gateway", "server_url": "litellm_proxy"}], + stream_chunk_size="sixty-four", + ) + assert exc_info.value.param == "stream_chunk_size" + + +def test_drop_params_false_still_rejects_an_invalid_stream_chunk_size() -> None: + with pytest.raises(litellm.BadRequestError): + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size="sixty-four", + drop_params=False, + mock_response="hi", + ) diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index 33467a78aa8..a2d944fcf39 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -504,7 +504,8 @@ LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2}, litellm_params.GuardrailOptions: {"guardrails": ("default",)}, litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}}, - litellm_params.ResponseOptions: {"stream_chunk_size": 64}, + litellm_params.ResponseOptions: {"keepalive_seconds": 1.5}, + litellm_params.ControlOptions: {"stream_chunk_size": 64}, litellm_params.MockOptions: {"mock_timeout": True}, litellm_params.CallState: { "completion_call_id": "call", @@ -533,7 +534,8 @@ LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"}, litellm_params.GuardrailOptions: {"guardrails": (1,)}, litellm_params.PromptOptions: {"prompt_id": 1}, - litellm_params.ResponseOptions: {"stream_chunk_size": "64"}, + litellm_params.ResponseOptions: {"keepalive_seconds": "1.5"}, + litellm_params.ControlOptions: {"stream_chunk_size": "sixty-four"}, litellm_params.MockOptions: {"mock_timeout": "true"}, litellm_params.CallState: {"completion_call_id": 1}, litellm_params.AgenticLoopState: {"depth": "1"}, @@ -583,10 +585,8 @@ def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, samp @pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id) def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None: - instance: Final = _leaf_instance(leaf, sample) - with pytest.raises(ValidationError): - _strict_leaf_validation(leaf, instance) + _strict_leaf_validation(leaf, _leaf_instance(leaf, sample)) @pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id)