mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Exploring streaming options for doing Bedrock passthrough streaming response bookkeeping.
Very much an idea - tangent
This commit is contained in:
parent
675e49ed94
commit
ad116cf38a
6 changed files with 149 additions and 38 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
101
litellm/passthrough/stream_flush_buffer.py
Normal file
101
litellm/passthrough/stream_flush_buffer.py
Normal 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")
|
||||
|
|
@ -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]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue