Exploring streaming options for doing Bedrock passthrough streaming response bookkeeping.

Very much an idea - tangent
This commit is contained in:
harish-berri 2026-05-04 22:02:16 +00:00
parent 675e49ed94
commit ad116cf38a
6 changed files with 149 additions and 38 deletions

View file

@ -170,6 +170,7 @@ from .specialty_caches.dynamic_logging_cache import DynamicLoggingCache
if TYPE_CHECKING:
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.passthrough.stream_flush_buffer import StreamBuffer
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
@ -1954,12 +1955,13 @@ class Logging(LiteLLMLoggingBaseClass):
def _flush_passthrough_collected_chunks_helper(
self,
raw_bytes: List[bytes],
stream_flush_buffer: "StreamBuffer",
provider_config: "BasePassthroughConfig",
) -> Optional["CostResponseTypes"]:
all_chunks = provider_config._convert_raw_bytes_to_str_lines(raw_bytes)
# Consume-once iterator (e.g. Bedrock EventStreamBuffer); flush path invokes once per buffer.
chunk_iter = stream_flush_buffer.iter_chunks_for_logging(provider_config)
complete_streaming_response = provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,
chunks=chunk_iter,
litellm_logging_obj=self,
model=self.model,
custom_llm_provider=self.model_call_details.get("custom_llm_provider", ""),
@ -1969,20 +1971,20 @@ class Logging(LiteLLMLoggingBaseClass):
def flush_passthrough_collected_chunks(
self,
raw_bytes: List[bytes],
stream_flush_buffer: "StreamBuffer",
provider_config: "BasePassthroughConfig",
):
"""
Flush collected chunks from the logging object
This is used to log the collected chunks once streaming is done on passthrough endpoints
1. Decode the raw bytes to string lines
1. Decode buffered stream data to string payloads (provider-specific)
2. Get the complete streaming response from the provider config
3. Log the complete streaming response (trigger success handler)
This is used for passthrough endpoints
"""
complete_streaming_response = self._flush_passthrough_collected_chunks_helper(
raw_bytes=raw_bytes,
stream_flush_buffer=stream_flush_buffer,
provider_config=provider_config,
)
@ -1992,11 +1994,11 @@ class Logging(LiteLLMLoggingBaseClass):
async def async_flush_passthrough_collected_chunks(
self,
raw_bytes: List[bytes],
stream_flush_buffer: "StreamBuffer",
provider_config: "BasePassthroughConfig",
):
complete_streaming_response = self._flush_passthrough_collected_chunks_helper(
raw_bytes=raw_bytes,
stream_flush_buffer=stream_flush_buffer,
provider_config=provider_config,
)

View file

@ -1,5 +1,5 @@
from abc import abstractmethod
from typing import TYPE_CHECKING, List, Optional, Tuple, Union
from typing import TYPE_CHECKING, Iterable, List, Optional, Tuple, Union
from ..base_utils import BaseLLMModelInfo
@ -7,6 +7,7 @@ if TYPE_CHECKING:
from httpx import URL, Headers, Response
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.passthrough.stream_flush_buffer import StreamBuffer
from litellm.types.utils import CostResponseTypes
from ..chat.transformation import BaseLLMException
@ -112,7 +113,7 @@ class BasePassthroughConfig(BaseLLMModelInfo):
def handle_logging_collected_chunks(
self,
all_chunks: List[str],
chunks: Iterable[str],
litellm_logging_obj: "LiteLLMLoggingObj",
model: str,
custom_llm_provider: str,
@ -137,3 +138,9 @@ class BasePassthroughConfig(BaseLLMModelInfo):
lines = [line.strip() for line in combined_str.split("\n") if line.strip()]
return lines
def build_stream_flush_buffer(self) -> "StreamBuffer":
"""Collect streaming response bytes until passthrough flush (default: raw chunks)."""
from litellm.passthrough.stream_flush_buffer import RawBytesStreamBuffer
return RawBytesStreamBuffer()

View file

@ -1,5 +1,5 @@
import json
from typing import TYPE_CHECKING, List, Optional, Tuple, cast
from typing import TYPE_CHECKING, Iterable, List, Optional, Tuple, cast
from httpx import Response
@ -9,18 +9,23 @@ from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConf
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockEventStreamDecoderBase, BedrockModelInfo
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import CostResponseTypes
if TYPE_CHECKING:
from httpx import URL
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.passthrough.stream_flush_buffer import StreamBuffer
from litellm.types.utils import CostResponseTypes
class BedrockPassthroughConfig(
BaseAWSLLM, BedrockModelInfo, BedrockEventStreamDecoderBase, BasePassthroughConfig
):
def build_stream_flush_buffer(self) -> "StreamBuffer":
"""Bedrock streaming uses AWS binary event-stream framing on the wire."""
from litellm.passthrough.stream_flush_buffer import BedrockEventStreamBuffer
return BedrockEventStreamBuffer()
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
return "stream" in endpoint
@ -177,14 +182,14 @@ class BedrockPassthroughConfig(
def handle_logging_collected_chunks(
self,
all_chunks: List[str],
chunks: Iterable[str],
litellm_logging_obj: "LiteLLMLoggingObj",
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["CostResponseTypes"]:
"""
1. Convert all_chunks to a ModelResponseStream
1. Convert streamed chunk payloads to a ModelResponseStream
2. combine model_response_stream to model_response
3. Return the model_response
"""
@ -223,7 +228,7 @@ class BedrockPassthroughConfig(
else:
return None
for chunk in all_chunks:
for chunk in chunks:
message = json.loads(chunk)
translated_chunk = obj._chunk_parser(chunk_data=message)

View file

@ -11,7 +11,6 @@ from typing import (
AsyncGenerator,
Coroutine,
Generator,
List,
Optional,
Union,
cast,
@ -25,6 +24,7 @@ from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.passthrough.stream_flush_buffer import StreamBuffer
from litellm.passthrough.utils import CommonUtils
from litellm.utils import client
@ -391,26 +391,24 @@ def _sync_streaming(
):
from litellm.utils import executor
raw_bytes: List[bytes] = []
stream_flush_buffer: StreamBuffer = provider_config.build_stream_flush_buffer()
flush_scheduled = False
try:
for chunk in response.iter_bytes(): # type: ignore
raw_bytes.append(chunk)
stream_flush_buffer.feed(chunk)
yield chunk
finally:
if not flush_scheduled and raw_bytes:
if not flush_scheduled and stream_flush_buffer.should_flush():
flush_scheduled = True
try:
executor.submit(
litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=raw_bytes,
stream_flush_buffer=stream_flush_buffer,
provider_config=provider_config,
)
except Exception as e:
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush "
"in _sync_streaming; %d buffered chunks dropped: %s",
len(raw_bytes),
"Failed to schedule passthrough spend-tracking flush in _sync_streaming: %s",
e,
)
@ -431,11 +429,11 @@ async def _async_streaming(
pass
raise
raw_bytes: List[bytes] = []
stream_flush_buffer: StreamBuffer = provider_config.build_stream_flush_buffer()
flush_scheduled = False
try:
async for chunk in iter_response.aiter_bytes(): # type: ignore
raw_bytes.append(chunk)
stream_flush_buffer.feed(chunk)
yield chunk
except Exception:
try:
@ -447,19 +445,17 @@ async def _async_streaming(
# GeneratorExit (raised on client disconnect) is not caught by
# `except Exception`; the finally block ensures partial usage
# still gets flushed for spend tracking. See LIT-2642.
if not flush_scheduled and raw_bytes:
if not flush_scheduled and stream_flush_buffer.should_flush():
flush_scheduled = True
try:
asyncio.create_task(
litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=raw_bytes,
stream_flush_buffer=stream_flush_buffer,
provider_config=provider_config,
)
)
except Exception as e:
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush "
"in _async_streaming; %d buffered chunks dropped: %s",
len(raw_bytes),
"Failed to schedule passthrough spend-tracking flush in _async_streaming: %s",
e,
)

View file

@ -0,0 +1,101 @@
"""
Streaming buffers for passthrough spend/logging flush.
``RawBytesStreamBuffer`` retains socket chunks for non–AWS-binary streams.
``BedrockEventStreamBuffer`` buffers raw bytes in botocore ``EventStreamBuffer``;
``iter_chunks_for_logging`` parses and UTF-8-decodes on first consumption only (see class doc).
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Iterator, List
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
class StreamBuffer(ABC):
"""Collect streamed bytes until passthrough flush."""
@abstractmethod
def feed(self, chunk: bytes) -> None:
raise NotImplementedError
@abstractmethod
def should_flush(self) -> bool:
raise NotImplementedError
@abstractmethod
def iter_chunks_for_logging(
self, provider_config: BasePassthroughConfig
) -> Iterator[str]:
"""Yield inner payload strings for ``handle_logging_collected_chunks``.
Passthrough flush calls this once per ``StreamBuffer`` lifetime. Implementations may
return a consume-once iterator (draining internal parsers); callers must not iterate
again without feeding new data unless the concrete buffer documents repeatable reads.
"""
def all_chunks_for_logging(self, provider_config: BasePassthroughConfig) -> List[str]:
"""Materialize logging chunks (calls ``iter_chunks_for_logging`` once)."""
return list(self.iter_chunks_for_logging(provider_config))
class RawBytesStreamBuffer(StreamBuffer):
"""Retain chunks as ``List[bytes]`` (SSE / plain streaming)."""
__slots__ = ("_chunks",)
def __init__(self) -> None:
self._chunks: List[bytes] = []
def feed(self, chunk: bytes) -> None:
self._chunks.append(chunk)
def should_flush(self) -> bool:
return len(self._chunks) > 0
def iter_chunks_for_logging(
self, provider_config: BasePassthroughConfig
) -> Iterator[str]:
"""Yield logging lines; repeatable (``_chunks`` are not drained by iteration)."""
yield from provider_config._convert_raw_bytes_to_str_lines(self._chunks)
@property
def raw_chunks(self) -> List[bytes]:
"""Read-only view for tests."""
return self._chunks
class BedrockEventStreamBuffer(StreamBuffer):
"""Bedrock ``application/vnd.amazon.eventstream``.
Iterating ``EventStreamBuffer`` consumes framed messages from botocore's internal buffer.
``iter_chunks_for_logging`` is therefore consume-once: a second pass yields nothing useful
unless more bytes were fed after the first iteration.
"""
__slots__ = ("_event_stream_buffer", "_received_any")
def __init__(self) -> None:
from botocore.eventstream import EventStreamBuffer
self._event_stream_buffer = EventStreamBuffer()
self._received_any = False
def feed(self, chunk: bytes) -> None:
self._received_any = True
self._event_stream_buffer.add_data(chunk)
def should_flush(self) -> bool:
return self._received_any
def iter_chunks_for_logging(
self, provider_config: BasePassthroughConfig
) -> Iterator[str]:
# Consume-once: botocore removes each parsed message from _event_stream_buffer._data.
for event in self._event_stream_buffer:
payload = event.payload
if payload:
yield payload.decode("utf-8")

View file

@ -63,7 +63,7 @@ async def test_async_streaming_flushes_on_normal_completion():
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == chunks
assert call_kwargs["stream_flush_buffer"].raw_chunks == chunks
assert call_kwargs["provider_config"] is provider_config
@ -101,7 +101,7 @@ async def test_async_streaming_flushes_on_client_disconnect():
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == [chunks[0]]
assert call_kwargs["stream_flush_buffer"].raw_chunks == [chunks[0]]
@pytest.mark.asyncio
@ -180,7 +180,7 @@ async def test_async_streaming_flushes_on_upstream_exception_with_partial_data()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == partial_chunks
assert call_kwargs["stream_flush_buffer"].raw_chunks == partial_chunks
def test_sync_streaming_flushes_on_normal_completion():
@ -241,4 +241,4 @@ def test_sync_streaming_flushes_on_early_close():
assert first == chunks[0]
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = mock_logging_obj.flush_passthrough_collected_chunks.call_args.kwargs
assert call_kwargs["raw_bytes"] == [chunks[0]]
assert call_kwargs["stream_flush_buffer"].raw_chunks == [chunks[0]]