diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 1c58a11eebe..8d50641e4c0 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -1655,32 +1655,72 @@ class AWSEventStreamDecoder: self, iterator: Iterator[bytes] ) -> Iterator[Union[GChunk, ModelResponseStream, dict]]: """Given an iterator that yields lines, iterate over it & yield every event encountered""" - from botocore.eventstream import EventStreamBuffer + from botocore.eventstream import ChecksumMismatch, EventStreamBuffer event_stream_buffer = EventStreamBuffer() + accumulated_bytes = b"" # Track raw bytes for error handling for chunk in iterator: - event_stream_buffer.add_data(chunk) - for event in event_stream_buffer: - message = self._parse_message_from_event(event) - if message: - # sse_event = ServerSentEvent(data=message, event="completion") - _data = json.loads(message) - yield self._chunk_parser(chunk_data=_data) + accumulated_bytes += chunk + try: + event_stream_buffer.add_data(chunk) + for event in event_stream_buffer: + message = self._parse_message_from_event(event) + if message: + # sse_event = ServerSentEvent(data=message, event="completion") + _data = json.loads(message) + yield self._chunk_parser(chunk_data=_data) + except ChecksumMismatch: + # Bedrock may return JSON error instead of event stream + # ChecksumMismatch indicates we received non-event-stream data + # See: https://github.com/BerriAI/litellm/issues/20589 + try: + error_response = json.loads(accumulated_bytes.decode("utf-8")) + error_message = error_response.get("message", str(error_response)) + raise BedrockError( + status_code=400, + message=f"Bedrock returned error: {error_message}", + ) + except json.JSONDecodeError: + # Not valid JSON, re-raise original checksum error + raise BedrockError( + status_code=500, + message=f"Bedrock streaming error: received malformed response data", + ) async def aiter_bytes( self, iterator: AsyncIterator[bytes] ) -> AsyncIterator[Union[GChunk, ModelResponseStream, dict]]: """Given an async iterator that yields lines, iterate over it & yield every event encountered""" - from botocore.eventstream import EventStreamBuffer + from botocore.eventstream import ChecksumMismatch, EventStreamBuffer event_stream_buffer = EventStreamBuffer() + accumulated_bytes = b"" # Track raw bytes for error handling async for chunk in iterator: - event_stream_buffer.add_data(chunk) - for event in event_stream_buffer: - message = self._parse_message_from_event(event) - if message: - _data = json.loads(message) - yield self._chunk_parser(chunk_data=_data) + accumulated_bytes += chunk + try: + event_stream_buffer.add_data(chunk) + for event in event_stream_buffer: + message = self._parse_message_from_event(event) + if message: + _data = json.loads(message) + yield self._chunk_parser(chunk_data=_data) + except ChecksumMismatch: + # Bedrock may return JSON error instead of event stream + # ChecksumMismatch indicates we received non-event-stream data + # See: https://github.com/BerriAI/litellm/issues/20589 + try: + error_response = json.loads(accumulated_bytes.decode("utf-8")) + error_message = error_response.get("message", str(error_response)) + raise BedrockError( + status_code=400, + message=f"Bedrock returned error: {error_message}", + ) + except json.JSONDecodeError: + # Not valid JSON, re-raise original checksum error + raise BedrockError( + status_code=500, + message=f"Bedrock streaming error: received malformed response data", + ) def _parse_message_from_event(self, event) -> Optional[str]: response_dict = event.to_response_dict() diff --git a/tests/test_litellm/llms/bedrock/test_bedrock_checksum_mismatch.py b/tests/test_litellm/llms/bedrock/test_bedrock_checksum_mismatch.py new file mode 100644 index 00000000000..f1161015656 --- /dev/null +++ b/tests/test_litellm/llms/bedrock/test_bedrock_checksum_mismatch.py @@ -0,0 +1,82 @@ +""" +Test for Bedrock ChecksumMismatch error handling. + +When Bedrock returns a JSON error response instead of a binary event stream, +the botocore EventStreamBuffer throws ChecksumMismatch. This test verifies +that we properly catch this error and return a meaningful error message. + +Related issue: https://github.com/BerriAI/litellm/issues/20589 +""" + +import json +import pytest +from unittest.mock import AsyncMock, MagicMock + +from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder +from litellm.llms.bedrock.common_utils import BedrockError + + +class TestBedrockChecksumMismatchHandling: + """Test that ChecksumMismatch errors from Bedrock are handled properly.""" + + def test_iter_bytes_handles_json_error_response(self): + """ + Test that iter_bytes properly handles when Bedrock returns a JSON error + instead of a binary event stream. + """ + decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + + # Simulate Bedrock returning a JSON error response + # This is what happens with inference profile ARN format issues + json_error = b'{"message": "Validation error: Invalid inference profile ARN format"}' + + def mock_iterator(): + yield json_error + + # Should raise BedrockError with the actual error message, not ChecksumMismatch + with pytest.raises(BedrockError) as exc_info: + list(decoder.iter_bytes(mock_iterator())) + + assert "Bedrock returned error" in str(exc_info.value.message) + assert "Validation error" in str(exc_info.value.message) + + @pytest.mark.asyncio + async def test_aiter_bytes_handles_json_error_response(self): + """ + Test that aiter_bytes properly handles when Bedrock returns a JSON error + instead of a binary event stream (async version). + """ + decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + + # Simulate Bedrock returning a JSON error response + json_error = b'{"message": "Validation error: Invalid inference profile ARN format"}' + + async def mock_async_iterator(): + yield json_error + + # Should raise BedrockError with the actual error message + with pytest.raises(BedrockError) as exc_info: + chunks = [] + async for chunk in decoder.aiter_bytes(mock_async_iterator()): + chunks.append(chunk) + + assert "Bedrock returned error" in str(exc_info.value.message) + assert "Validation error" in str(exc_info.value.message) + + def test_iter_bytes_handles_malformed_response(self): + """ + Test that iter_bytes properly handles completely malformed responses + that are neither valid event stream nor valid JSON. + """ + decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + + # Random bytes that aren't valid event stream or JSON + malformed_data = b'\x00\x01\x02\x03invalid data' + + def mock_iterator(): + yield malformed_data + + with pytest.raises(BedrockError) as exc_info: + list(decoder.iter_bytes(mock_iterator())) + + assert "malformed response" in str(exc_info.value.message).lower()