mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(bedrock): honour stream_chunk_size in Invoke streaming (#42686)
This commit is contained in:
parent
efb93e62f8
commit
c2b388ebe6
6 changed files with 247 additions and 7 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
20
tests/unit/llms/bedrock/test_common_utils.py
Normal file
20
tests/unit/llms/bedrock/test_common_utils.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue