From c2b388ebe65b4e06d30130fc1f77cc7218c27b27 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 11:25:00 -0700 Subject: [PATCH] fix(bedrock): honour stream_chunk_size in Invoke streaming (#42686) --- litellm/llms/bedrock/chat/converse_handler.py | 4 +- .../base_invoke_transformation.py | 6 +- litellm/llms/bedrock/common_utils.py | 11 +- .../test_base_invoke_transformation.py | 169 +++++++++++++++++- tests/unit/llms/bedrock/test_common_utils.py | 20 +++ tests/unit/llms/chat/test_converse_handler.py | 44 ++++- 6 files changed, 247 insertions(+), 7 deletions(-) create mode 100644 tests/unit/llms/bedrock/test_common_utils.py diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 3abe7c88fcd..bd358805743 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -18,7 +18,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 +from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text, stream_chunk_size_from from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call @@ -278,7 +278,7 @@ class BedrockConverseLLM(BaseAWSLLM): ): ## SETUP ## stream: Final = optional_params.pop("stream", None) - stream_chunk_size: Final = litellm_params.get("stream_chunk_size") + stream_chunk_size: Final = stream_chunk_size_from(litellm_params) 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/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index dcc5e249d8a..629806b58e2 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -18,7 +18,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 +from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from from litellm.llms.bedrock.request_metadata import ( bedrock_request_metadata_headers, merge_bedrock_invoke_headers, @@ -453,6 +453,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): json_mode: bool | None = None, signed_json_body: bytes | None = None, ) -> CustomStreamWrapper: + chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params) completion_stream, response_headers = await make_call( client=client, api_base=api_base, @@ -464,6 +465,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): fake_stream=True if "ai21" in api_base else False, bedrock_invoke_provider=self.get_bedrock_invoke_provider(model), json_mode=json_mode, + stream_chunk_size=chunk_size, ) streaming_response: Final = CustomStreamWrapper( completion_stream=completion_stream, @@ -491,6 +493,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): 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) completion_stream, response_headers = make_sync_call( client=sync_client, api_base=api_base, @@ -503,6 +506,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): fake_stream=True if "ai21" in api_base else False, bedrock_invoke_provider=self.get_bedrock_invoke_provider(model), json_mode=json_mode, + stream_chunk_size=chunk_size, ) streaming_response: Final = CustomStreamWrapper( completion_stream=completion_stream, diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 2e20aafffcb..c60ba4e802f 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 TypeAdapter, ValidationError +from pydantic import ConfigDict, TypeAdapter, ValidationError import litellm from litellm import verbose_logger @@ -86,6 +86,15 @@ 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/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 96a2fa6ec67..ed172fdfbff 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,10 +1,11 @@ import json -from unittest.mock import MagicMock +from typing import Final +from unittest.mock import AsyncMock, MagicMock import httpx import pytest - +import litellm from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( AmazonAnthropicClaudeConfig, ) @@ -12,6 +13,12 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation 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, +) @pytest.mark.parametrize( @@ -234,3 +241,161 @@ def test_transform_response_hands_json_mode_to_nova(): assert result.choices[0].message.tool_calls is None 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) + 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) + + litellm.completion( + model="bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + aws_access_key_id="fake", + aws_secret_access_key="fake", + aws_region_name="us-east-1", + **kwargs, + ) + return mock_response.iter_bytes, client.post, recorder + + +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) + + 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) + + 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 _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() + client.post = AsyncMock(return_value=mock_response) + + await litellm.acompletion( + model="bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + aws_access_key_id="fake", + aws_secret_access_key="fake", + aws_region_name="us-east-1", + **kwargs, + ) + return aiter_bytes_spy, client.post, recorder + + +@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 + ) + + 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) + + 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 = { + "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}), + } + ] + ) + + router.completion( + model="invoke-chunked", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + ) + + 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) + + with pytest.raises(litellm.BadRequestError): + litellm.completion( + model="bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + aws_access_key_id="fake", + aws_secret_access_key="fake", + aws_region_name="us-east-1", + stream_chunk_size="sixty-four", + ) + + client.post.assert_not_called() diff --git a/tests/unit/llms/bedrock/test_common_utils.py b/tests/unit/llms/bedrock/test_common_utils.py new file mode 100644 index 00000000000..cfcc15f186b --- /dev/null +++ b/tests/unit/llms/bedrock/test_common_utils.py @@ -0,0 +1,20 @@ +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 c342a7e1806..cbb8e3acf78 100644 --- a/tests/unit/llms/chat/test_converse_handler.py +++ b/tests/unit/llms/chat/test_converse_handler.py @@ -18,7 +18,6 @@ from tests._support.stream_chunk_size import ( ) - def test_encode_model_id_with_inference_profile(): """ Test instance profile is properly encoded when used as a model @@ -459,6 +458,49 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes( 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() + + 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, + client=client, + aws_access_key_id="fake", + aws_secret_access_key="fake", + aws_region_name="us-east-1", + 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() + + def _bedrock_error_response(status_code: int, request_id: str) -> httpx.Response: return httpx.Response( status_code=status_code,