mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #9515 from BerriAI/litellm_sagemaker_fix_stream
[Bug fix] - Sagemaker endpoint with inference component streaming
This commit is contained in:
commit
0d48652012
4 changed files with 131 additions and 18 deletions
|
|
@ -214,7 +214,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
|
|||
|
||||
# Set up the Sentinel client
|
||||
sentinel = redis.Sentinel(
|
||||
sentinel_nodes,
|
||||
sentinel_nodes,
|
||||
socket_timeout=0.1,
|
||||
password=sentinel_password,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -127,21 +127,25 @@ class AWSEventStreamDecoder:
|
|||
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:
|
||||
verbose_logger.debug("sagemaker parsed chunk bytes %s", message)
|
||||
# remove data: prefix and "\n\n" at the end
|
||||
message = (
|
||||
litellm.CustomStreamWrapper._strip_sse_data_from_chunk(message)
|
||||
or ""
|
||||
)
|
||||
message = message.replace("\n\n", "")
|
||||
try:
|
||||
message = self._parse_message_from_event(event)
|
||||
if message:
|
||||
verbose_logger.debug(
|
||||
"sagemaker parsed chunk bytes %s", message
|
||||
)
|
||||
# remove data: prefix and "\n\n" at the end
|
||||
message = (
|
||||
litellm.CustomStreamWrapper._strip_sse_data_from_chunk(
|
||||
message
|
||||
)
|
||||
or ""
|
||||
)
|
||||
message = message.replace("\n\n", "")
|
||||
|
||||
# Accumulate JSON data
|
||||
accumulated_json += message
|
||||
# Accumulate JSON data
|
||||
accumulated_json += message
|
||||
|
||||
# Try to parse the accumulated JSON
|
||||
try:
|
||||
# Try to parse the accumulated JSON
|
||||
_data = json.loads(accumulated_json)
|
||||
if self.is_messages_api:
|
||||
yield self._chunk_parser_messages_api(chunk_data=_data)
|
||||
|
|
@ -149,9 +153,19 @@ class AWSEventStreamDecoder:
|
|||
yield self._chunk_parser(chunk_data=_data)
|
||||
# Reset accumulated_json after successful parsing
|
||||
accumulated_json = ""
|
||||
except json.JSONDecodeError:
|
||||
# If it's not valid JSON yet, continue to the next event
|
||||
continue
|
||||
except json.JSONDecodeError:
|
||||
# If it's not valid JSON yet, continue to the next event
|
||||
continue
|
||||
except UnicodeDecodeError as e:
|
||||
verbose_logger.warning(
|
||||
f"UnicodeDecodeError: {e}. Attempting to combine with next event."
|
||||
)
|
||||
continue
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Error parsing message: {e}. Attempting to combine with next event."
|
||||
)
|
||||
continue
|
||||
|
||||
# Handle any remaining data after the iterator is exhausted
|
||||
if accumulated_json:
|
||||
|
|
@ -167,6 +181,8 @@ class AWSEventStreamDecoder:
|
|||
f"Warning: Unparseable JSON data remained: {accumulated_json}"
|
||||
)
|
||||
yield None
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Final error parsing accumulated JSON: {e}")
|
||||
|
||||
def _parse_message_from_event(self, event) -> Optional[str]:
|
||||
response_dict = event.to_response_dict()
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ MCP Client Manager
|
|||
|
||||
This class is responsible for managing MCP SSE clients.
|
||||
|
||||
This is a Proxy
|
||||
This is a Proxy
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
|
|
|||
97
tests/litellm/llms/sagemaker/test_sagemaker_common_utils.py
Normal file
97
tests/litellm/llms/sagemaker/test_sagemaker_common_utils.py
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
from litellm.llms.sagemaker.common_utils import AWSEventStreamDecoder
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aiter_bytes_unicode_decode_error():
|
||||
"""
|
||||
Test that AWSEventStreamDecoder.aiter_bytes() does not raise an error when encountering invalid UTF-8 bytes. (UnicodeDecodeError)
|
||||
|
||||
|
||||
Ensures stream processing continues despite the error.
|
||||
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/9165
|
||||
"""
|
||||
# Create an instance of AWSEventStreamDecoder
|
||||
decoder = AWSEventStreamDecoder(model="test-model")
|
||||
|
||||
# Create a mock event that will trigger a UnicodeDecodeError
|
||||
mock_event = MagicMock()
|
||||
mock_event.to_response_dict.return_value = {
|
||||
"status_code": 200,
|
||||
"headers": {},
|
||||
"body": b"\xff\xfe", # Invalid UTF-8 bytes
|
||||
}
|
||||
|
||||
# Create a mock EventStreamBuffer that yields our mock event
|
||||
mock_buffer = MagicMock()
|
||||
mock_buffer.__iter__.return_value = [mock_event]
|
||||
|
||||
# Mock the EventStreamBuffer class
|
||||
with patch("botocore.eventstream.EventStreamBuffer", return_value=mock_buffer):
|
||||
# Create an async generator that yields some test bytes
|
||||
async def mock_iterator():
|
||||
yield b""
|
||||
|
||||
# Process the stream
|
||||
chunks = []
|
||||
async for chunk in decoder.aiter_bytes(mock_iterator()):
|
||||
if chunk is not None:
|
||||
print("chunk=", chunk)
|
||||
chunks.append(chunk)
|
||||
|
||||
# Verify that processing continued despite the error
|
||||
# The chunks list should be empty since we only sent invalid data
|
||||
assert len(chunks) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aiter_bytes_valid_chunk_followed_by_unicode_error():
|
||||
"""
|
||||
Test that valid chunks are processed correctly even when followed by Unicode decode errors.
|
||||
This ensures errors don't corrupt or prevent processing of valid data that came before.
|
||||
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/9165
|
||||
"""
|
||||
decoder = AWSEventStreamDecoder(model="test-model")
|
||||
|
||||
# Create two mock events - first valid, then invalid
|
||||
mock_valid_event = MagicMock()
|
||||
mock_valid_event.to_response_dict.return_value = {
|
||||
"status_code": 200,
|
||||
"headers": {},
|
||||
"body": json.dumps({"token": {"text": "hello"}}).encode(), # Valid data first
|
||||
}
|
||||
|
||||
mock_invalid_event = MagicMock()
|
||||
mock_invalid_event.to_response_dict.return_value = {
|
||||
"status_code": 200,
|
||||
"headers": {},
|
||||
"body": b"\xff\xfe", # Invalid UTF-8 bytes second
|
||||
}
|
||||
|
||||
# Create a mock EventStreamBuffer that yields valid event first, then invalid
|
||||
mock_buffer = MagicMock()
|
||||
mock_buffer.__iter__.return_value = [mock_valid_event, mock_invalid_event]
|
||||
|
||||
with patch("botocore.eventstream.EventStreamBuffer", return_value=mock_buffer):
|
||||
|
||||
async def mock_iterator():
|
||||
yield b"test_bytes"
|
||||
|
||||
chunks = []
|
||||
async for chunk in decoder.aiter_bytes(mock_iterator()):
|
||||
if chunk is not None:
|
||||
chunks.append(chunk)
|
||||
|
||||
# Verify we got our valid chunk despite the subsequent error
|
||||
assert len(chunks) == 1
|
||||
assert chunks[0]["text"] == "hello" # Verify the content of the valid chunk
|
||||
Loading…
Add table
Reference in a new issue