From ad116cf38af5e1d2c8f578dd289a83b899373695 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Mon, 4 May 2026 22:02:16 +0000 Subject: [PATCH] Exploring streaming options for doing Bedrock passthrough streaming response bookkeeping. Very much an idea - tangent --- litellm/litellm_core_utils/litellm_logging.py | 18 ++-- .../base_llm/passthrough/transformation.py | 11 +- .../bedrock/passthrough/transformation.py | 23 ++-- litellm/passthrough/main.py | 26 ++--- litellm/passthrough/stream_flush_buffer.py | 101 ++++++++++++++++++ ...test_streaming_interrupt_spend_tracking.py | 8 +- 6 files changed, 149 insertions(+), 38 deletions(-) create mode 100644 litellm/passthrough/stream_flush_buffer.py diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index a815442c2f9..4816bf73d3f 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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, ) diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index 9d4396dce47..805b6e4783c 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -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() diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index 274b0282acc..7bf661ada69 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -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) diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index c4c9aea6f64..7cf317ac4cd 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -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, ) diff --git a/litellm/passthrough/stream_flush_buffer.py b/litellm/passthrough/stream_flush_buffer.py new file mode 100644 index 00000000000..fb83143e58a --- /dev/null +++ b/litellm/passthrough/stream_flush_buffer.py @@ -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") diff --git a/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py b/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py index f3fe3ae5c38..455b6789b4e 100644 --- a/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py +++ b/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py @@ -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]]