fix(bedrock): honour stream_chunk_size in Invoke streaming (#42686)

This commit is contained in:
devin-ai-integration[bot] 2026-09-23 11:25:00 -07:00 • committed by GitHub
parent efb93e62f8
commit c2b388ebe6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 247 additions and 7 deletions

View file

@ -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)

View file

@ -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,

View file

@ -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(

View file

@ -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()

View file

@ -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

View file

@ -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,