fix(passthrough): parse Bedrock stream spend incrementally instead of buffering the whole response (#40724)

* fix(passthrough): parse Bedrock stream spend incrementally instead of buffering the whole response

Bedrock pass-through streaming kept every relayed chunk in memory until EOF and
then decoded, parsed and translated the whole stream again for spend logging.
Large or concurrent streams could exhaust proxy worker memory.

Sync and async passthrough wrappers now hand each chunk to a provider stream
collector as it is relayed. Bedrock decodes event-stream frames incrementally,
folds consecutive text deltas, and keeps only what stream_chunk_builder needs
for usage, tool calls and metadata. Text deltas are no longer retained in the
Bedrock and Anthropic stream decoders either. Providers without a collector
keep the previous raw-bytes behavior. Collector failures are isolated so spend
tracking can never interrupt the customer stream

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(passthrough): assert the spend payload the collector builds instead of mock internals

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(passthrough): type the Bedrock collector helpers by the collector protocol instead of asserting the class

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-11 09:53:37 -07:00 • committed by GitHub
parent 3df127b439
commit 47bba14336
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 507 additions and 174 deletions

View file

@ -14,7 +14,7 @@ from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from types import MappingProxyType, TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
from httpx import Response
from pydantic import BaseModel, JsonValue
@ -211,7 +211,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, LoggedRelayResponse
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
@ -2396,52 +2396,17 @@ class Logging(LiteLLMLoggingBaseClass):
for scope in [key for key in spans_logged if isinstance(key, tuple) and key[-1:] == ("success",)]:
del spans_logged[scope]
def _flush_passthrough_collected_chunks_helper(
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
) -> Optional["LoggedRelayResponse"]:
all_chunks: Final = provider_config._convert_raw_bytes_to_str_lines(raw_bytes)
complete_streaming_response: Final = provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,
litellm_logging_obj=self,
model=self.model,
custom_llm_provider=self.model_call_details.get("custom_llm_provider", ""),
endpoint=self.model_call_details.get("endpoint", ""),
)
return complete_streaming_response
def flush_passthrough_collected_chunks(
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
):
def flush_passthrough_collected_chunks(self, collector: "PassthroughStreamCollector"):
"""
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
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
Log the response a passthrough stream collector assembled once streaming is done (trigger success handler)
"""
complete_streaming_response: Final = self._flush_passthrough_collected_chunks_helper(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
complete_streaming_response: Final = collector.build_logged_response(litellm_logging_obj=self)
if complete_streaming_response is not None:
self.success_handler(result=complete_streaming_response)
async def async_flush_passthrough_collected_chunks(
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
):
complete_streaming_response: Final = self._flush_passthrough_collected_chunks_helper(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
async def async_flush_passthrough_collected_chunks(self, collector: "PassthroughStreamCollector"):
complete_streaming_response: Final = collector.build_logged_response(litellm_logging_obj=self)
if complete_streaming_response is not None:
await self.async_success_handler(result=complete_streaming_response)
@ -6505,7 +6470,7 @@ def _get_traceback_str_for_error(error_str: str) -> str:
from decimal import Decimal
# used for unit testing
from typing import Any, Optional, Union
from typing import Any, Union
def create_dummy_standard_logging_payload() -> StandardLoggingPayload:

View file

@ -724,10 +724,11 @@ class ModelResponseIterator:
content_block: Final = ContentBlockDelta(**chunk)
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = []
self.content_blocks.append(content_block)
if "text" in content_block["delta"]:
text = content_block["delta"]["text"]
elif "partial_json" in content_block["delta"]:
return text, tool_use, thinking_blocks, provider_specific_fields, reasoning_content
self.content_blocks.append(content_block)
if "partial_json" in content_block["delta"]:
# Only emit tool calls if we're in a tool_use or server_tool_use block
# web_search_tool_result blocks also have input_json_delta but should not be treated as tool calls
# See: https://github.com/BerriAI/litellm/issues/17254

View file

@ -4,7 +4,7 @@ import re
from abc import abstractmethod
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, TypeAlias
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
from pydantic import TypeAdapter, ValidationError
@ -80,6 +80,38 @@ def logged_relay_shape(
return parsed
class PassthroughStreamCollector(Protocol):
"""Consumes relayed stream bytes as they arrive and builds the response logged for spend tracking."""
def add(self, chunk: bytes) -> None: ...
def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None: ...
class RawBytesStreamCollector:
def __init__(
self, provider_config: BasePassthroughConfig, model: str, custom_llm_provider: str, endpoint: str
) -> None:
self._provider_config = provider_config
self._model = model
self._custom_llm_provider = custom_llm_provider
self._endpoint = endpoint
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
def add(self, chunk: bytes) -> None:
self._raw_bytes.append(chunk)
def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None:
all_chunks: Final = self._provider_config._convert_raw_bytes_to_str_lines(self._raw_bytes)
return self._provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=self._model,
custom_llm_provider=self._custom_llm_provider,
endpoint=self._endpoint,
)
class BasePassthroughConfig(BaseLLMModelInfo):
@abstractmethod
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
@ -182,6 +214,13 @@ class BasePassthroughConfig(BaseLLMModelInfo):
) -> LoggedRelayResponse | None:
return None
def create_stream_collector(
self, model: str, custom_llm_provider: str, endpoint: str
) -> PassthroughStreamCollector:
return RawBytesStreamCollector(
provider_config=self, model=model, custom_llm_provider=custom_llm_provider, endpoint=endpoint
)
def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]:
"""
Converts a list of raw bytes into a list of string lines, similar to aiter_lines()

View file

@ -490,10 +490,10 @@ class AWSEventStreamDecoder:
reasoning_content: str | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
self.content_blocks.append(delta_obj)
if "text" in delta_obj:
text = delta_obj["text"]
elif "toolUse" in delta_obj:
self.content_blocks.append(delta_obj)
# When json_mode is True and this is the internal json_tool_call,
# convert tool input to text content instead of tool call arguments
if self.json_mode is True and self._current_tool_name == RESPONSE_FORMAT_TOOL_NAME:

View file

@ -1,23 +1,129 @@
import json
from collections.abc import Mapping
from collections.abc import Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Final, Optional, cast
import httpx
from httpx import Response
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, PassthroughStreamCollector
from litellm.types.utils import ModelResponseStream
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, BedrockEventStreamDecoderBase, BedrockModelInfo
if TYPE_CHECKING:
from botocore.eventstream import EventStreamMessage
from httpx import URL
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
from litellm.types.utils import CostResponseTypes
_TEXT_ONLY_DELTA_FIELDS: Final = frozenset({"content", "role"})
def _plain_text_delta(chunk: ModelResponseStream) -> str | None:
"""Return the delta text when the chunk carries nothing else that stream_chunk_builder reads."""
if chunk.get("usage") is not None or chunk.provider_specific_fields or len(chunk.choices) != 1:
return None
choice: Final = chunk.choices[0]
if choice.finish_reason or choice.logprobs is not None:
return None
populated: Final = frozenset(key for key, value in choice.delta.model_dump().items() if value is not None)
if not populated <= _TEXT_ONLY_DELTA_FIELDS:
return None
content: Final = choice.delta.get("content")
return content if isinstance(content, str) else None
class _CoalescedChunks:
"""Retains translated chunks with consecutive text deltas folded into one, so memory tracks the response text,
not the event count."""
def __init__(self) -> None:
self._chunks: list[ModelResponseStream] = [] # mutable-ok: instance accumulator for streaming chunks
self._open_text_parts: list[str] = [] # mutable-ok: text deltas pending fold into self._chunks[-1]
def add(self, chunk: ModelResponseStream) -> None:
text: Final = _plain_text_delta(chunk)
if text is not None and self._open_text_parts:
self._open_text_parts.append(text)
return
self._seal_text_run()
self._chunks.append(chunk)
if text is not None:
self._open_text_parts.append(text)
def _seal_text_run(self) -> None:
if len(self._open_text_parts) > 1:
self._chunks[-1].choices[0].delta.content = "".join(self._open_text_parts)
self._open_text_parts.clear()
def chunks(self) -> Sequence[ModelResponseStream]:
self._seal_text_run()
return self._chunks
def _translate_message(decoder: "AWSEventStreamDecoder", message: str) -> ModelResponseStream | None:
from litellm.litellm_core_utils.streaming_handler import (
convert_generic_chunk_to_model_response_stream,
generic_chunk_has_all_required_fields,
)
from litellm.types.utils import GenericStreamingChunk
translated_chunk: Final = decoder._chunk_parser(chunk_data=json.loads(message))
if isinstance(translated_chunk, ModelResponseStream):
return translated_chunk
if generic_chunk_has_all_required_fields(cast(dict, translated_chunk)):
return convert_generic_chunk_to_model_response_stream(cast(GenericStreamingChunk, translated_chunk))
return None
def _build_logged_response(
chunks: Sequence[ModelResponseStream], litellm_logging_obj: "LiteLLMLoggingObj"
) -> Optional["CostResponseTypes"]:
from litellm.main import stream_chunk_builder
if len(chunks) == 0:
return None
return stream_chunk_builder(chunks=list(chunks), logging_obj=litellm_logging_obj)
class BedrockEventStreamCollector:
"""Decodes and translates Bedrock event-stream frames as they are relayed instead of buffering the stream."""
def __init__(
self,
parse_event: Callable[["EventStreamMessage"], str | None],
decoder: Optional["AWSEventStreamDecoder"],
) -> None:
from botocore.eventstream import EventStreamBuffer
self._parse_event = parse_event
self._decoder = decoder
self._event_stream_buffer: Final[EventStreamBuffer] = EventStreamBuffer()
self._chunks: Final = _CoalescedChunks()
def add(self, chunk: bytes) -> None:
if self._decoder is None:
return
self._event_stream_buffer.add_data(chunk)
for event in self._event_stream_buffer:
self._add_event(self._decoder, event)
def _add_event(self, decoder: "AWSEventStreamDecoder", event: "EventStreamMessage") -> None:
message: Final = self._parse_event(event)
translated: Final = _translate_message(decoder, message) if message is not None else None
if translated is not None:
self._chunks.add(translated)
def build_logged_response(self, litellm_logging_obj: "LiteLLMLoggingObj") -> Optional["CostResponseTypes"]:
return _build_logged_response(self._chunks.chunks(), litellm_logging_obj)
class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamDecoderBase, BasePassthroughConfig):
def get_error_class(
self,
@ -168,87 +274,32 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD
return litellm_model_response
def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]:
from botocore.eventstream import EventStreamBuffer
all_chunks: Final = []
event_stream_buffer: Final = EventStreamBuffer()
for chunk in raw_bytes:
event_stream_buffer.add_data(chunk)
for event in event_stream_buffer:
message = self._parse_message_from_event(event)
if message is not None:
all_chunks.append(message)
return all_chunks
def handle_logging_collected_chunks(
self,
all_chunks: list[str],
litellm_logging_obj: "LiteLLMLoggingObj",
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["CostResponseTypes"]:
"""
1. Convert all_chunks to a ModelResponseStream
2. combine model_response_stream to model_response
3. Return the model_response
"""
from litellm.litellm_core_utils.streaming_handler import (
convert_generic_chunk_to_model_response_stream,
generic_chunk_has_all_required_fields,
def create_stream_collector(
self, model: str, custom_llm_provider: str, endpoint: str
) -> PassthroughStreamCollector:
return BedrockEventStreamCollector(
parse_event=self._parse_message_from_event,
decoder=self._get_event_stream_decoder(model=model, endpoint=endpoint),
)
def _get_event_stream_decoder(self, model: str, endpoint: str) -> Optional["AWSEventStreamDecoder"]:
from litellm.llms.bedrock.chat import get_bedrock_event_stream_decoder
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
from litellm.main import stream_chunk_builder
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
all_translated_chunks: Final = []
if "invoke" in endpoint:
invoke_provider: Final = AmazonInvokeConfig.get_bedrock_invoke_provider(model)
if invoke_provider is None:
raise ValueError(f"Invalid invoke provider: {invoke_provider}, for model: {model}")
obj = get_bedrock_event_stream_decoder(
invoke_provider=invoke_provider,
model=model,
sync_stream=True,
json_mode=False,
)
elif "converse" in endpoint:
obj = get_bedrock_event_stream_decoder(
invoke_provider=None,
model=model,
sync_stream=True,
json_mode=False,
)
else:
return None
for chunk in all_chunks:
message = json.loads(chunk)
translated_chunk = obj._chunk_parser(chunk_data=message)
if isinstance(translated_chunk, dict) and generic_chunk_has_all_required_fields(
cast(dict, translated_chunk)
):
chunk_obj = convert_generic_chunk_to_model_response_stream(
cast(GenericStreamingChunk, translated_chunk)
verbose_logger.warning(
"Bedrock passthrough spend tracking skipped: no invoke provider for model %s", model
)
elif isinstance(translated_chunk, ModelResponseStream):
chunk_obj = translated_chunk
else:
continue
all_translated_chunks.append(chunk_obj)
if len(all_translated_chunks) > 0:
model_response: Final = stream_chunk_builder(
chunks=all_translated_chunks,
logging_obj=litellm_logging_obj,
return None
return get_bedrock_event_stream_decoder(
invoke_provider=invoke_provider, model=model, sync_stream=True, json_mode=False
)
if "converse" in endpoint:
return get_bedrock_event_stream_decoder(
invoke_provider=None, model=model, sync_stream=True, json_mode=False
)
return model_response
return None

View file

@ -17,7 +17,7 @@ from httpx._types import CookieTypes, QueryParamTypes, RequestContent, RequestFi
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, PassthroughStreamCollector
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.passthrough.utils import CommonUtils
@ -36,6 +36,35 @@ def _as_generator(iterable: Iterator[bytes]) -> Generator[bytes, bytes, None]:
yield from iterable
class _SpendCollection:
"""Feeds relayed chunks to the provider's stream collector without letting spend tracking break the relay."""
def __init__(self, provider_config: BasePassthroughConfig, litellm_logging_obj: LiteLLMLoggingObj) -> None:
self.collector: Final[PassthroughStreamCollector] = provider_config.create_stream_collector(
model=litellm_logging_obj.model,
custom_llm_provider=litellm_logging_obj.model_call_details.get("custom_llm_provider", ""),
endpoint=litellm_logging_obj.model_call_details.get("endpoint", ""),
)
self.chunk_count = 0
self._failed = False
def add(self, chunk: bytes) -> None:
self.chunk_count += 1
if self._failed:
return
try:
self.collector.add(chunk)
except Exception as e: # noqa: BLE001 # Safe catch-all: spend tracking must never break the relayed stream
self._failed = True
verbose_logger.exception(
"Passthrough spend-tracking collector failed; spend dropped for this stream: %s", e
)
@property
def should_flush(self) -> bool:
return self.chunk_count > 0 and not self._failed
class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
def __init__(
self,
@ -50,8 +79,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
self._response: httpx.Response
self._iterator: AsyncGenerator[bytes, bytes]
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._spend = _SpendCollection(provider_config, litellm_logging_obj)
self._flush_scheduled = False
self._background_tasks: set[asyncio.Task] = set() # mutable-ok: instance set for background task tracking
self._hidden_params: dict[str, object] = {} # mutable-ok: router attaches response headers here in place
@ -101,16 +129,13 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
return _init().__await__()
def _start_flush(self) -> None:
if self._flush_scheduled or not self._raw_bytes:
if self._flush_scheduled or not self._spend.should_flush:
return
self._flush_scheduled = True
try:
task: Final = asyncio.create_task(
self._litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
)
self._litellm_logging_obj.async_flush_passthrough_collected_chunks(collector=self._spend.collector)
)
self._background_tasks.add(task)
@ -118,8 +143,8 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
task.add_done_callback(self._background_tasks.discard)
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s",
len(self._raw_bytes),
"Failed to schedule passthrough spend-tracking flush; %d collected chunks dropped: %s",
self._spend.chunk_count,
e,
)
@ -134,7 +159,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
await self # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__
try:
chunk: Final = await anext(self._iterator)
self._raw_bytes.append(chunk)
self._spend.add(chunk)
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
self._start_flush()
try:
@ -181,13 +206,12 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
self.headers = response.headers
self.status_code = response.status_code
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._iterator: Generator[bytes, bytes, None] = _as_generator(response.iter_bytes())
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._spend = _SpendCollection(provider_config, litellm_logging_obj)
self._flush_scheduled = False
def _start_flush(self) -> None:
if self._flush_scheduled or not self._raw_bytes:
if self._flush_scheduled or not self._spend.should_flush:
return
self._flush_scheduled = True
@ -195,14 +219,12 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
try:
executor.submit(
self._litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
self._litellm_logging_obj.flush_passthrough_collected_chunks, collector=self._spend.collector
)
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s",
len(self._raw_bytes),
"Failed to schedule passthrough spend-tracking flush; %d collected chunks dropped: %s",
self._spend.chunk_count,
e,
)
@ -212,7 +234,7 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
def __next__(self) -> bytes:
try:
chunk: Final = next(self._iterator)
self._raw_bytes.append(chunk)
self._spend.add(chunk)
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
self._start_flush()
try:

View file

@ -608,9 +608,11 @@ async def test_streaming_responses_relay_flush_reaches_the_success_callbacks_wit
)
stream = "event: response.completed\ndata: " + json.dumps(RESPONSES_COMPLETED_EVENT) + "\n\n"
await logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=[stream.encode()], provider_config=AzureAIPassthroughConfig()
collector = AzureAIPassthroughConfig().create_stream_collector(
model="gpt-5.4-mini", custom_llm_provider="azure_ai", endpoint="gpt/openai/responses"
)
collector.add(stream.encode())
await logging_obj.async_flush_passthrough_collected_chunks(collector=collector)
info = litellm.get_model_info("azure_ai/gpt-5.4-mini")
assert probe.logged_call_type == "allm_passthrough_route"

View file

@ -1,7 +1,19 @@
import base64
import json
import struct
import tracemalloc
from binascii import crc32
from datetime import datetime
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig
from litellm.types.utils import ModelResponse
CONVERSE_MODEL = "anthropic.claude-sonnet-4-5-20250929-v1:0"
CONVERSE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/converse-stream"
INVOKE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/invoke-with-response-stream"
def test_bedrock_passthrough_get_complete_url_default_endpoint():
@ -500,3 +512,186 @@ def test_bedrock_passthrough_model_id_without_arn():
f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model_id}/converse"
)
assert url_str == expected_url
def _event_frame(event_type: str, payload: dict) -> bytes:
def header(name: str, value: str) -> bytes:
name_b, value_b = name.encode(), value.encode()
return struct.pack("!B", len(name_b)) + name_b + struct.pack("!B", 7) + struct.pack("!H", len(value_b)) + value_b
payload_b = json.dumps(payload, separators=(",", ":")).encode()
headers_b = (
header(":event-type", event_type)
+ header(":content-type", "application/json")
+ header(":message-type", "event")
)
prelude = struct.pack("!II", 12 + len(headers_b) + len(payload_b) + 4, len(headers_b))
prelude_crc = crc32(prelude) & 0xFFFFFFFF
message = struct.pack("!I", prelude_crc) + headers_b + payload_b
return prelude + message + struct.pack("!I", crc32(message, prelude_crc) & 0xFFFFFFFF)
def _text_block(index: int, texts: list[str]) -> bytes:
return (
_event_frame("contentBlockStart", {"contentBlockIndex": index, "start": {}})
+ b"".join(
_event_frame("contentBlockDelta", {"contentBlockIndex": index, "delta": {"text": text}}) for text in texts
)
+ _event_frame("contentBlockStop", {"contentBlockIndex": index})
)
def _stream_tail(stop_reason: str, output_tokens: int) -> bytes:
return _event_frame("messageStop", {"stopReason": stop_reason}) + _event_frame(
"metadata",
{
"metrics": {"latencyMs": 1234},
"usage": {"inputTokens": 25, "outputTokens": output_tokens, "totalTokens": 25 + output_tokens},
},
)
def _invoke_chunk(payload: dict) -> bytes:
return _event_frame("chunk", {"bytes": base64.b64encode(json.dumps(payload).encode()).decode()})
def _stream_logging_obj(endpoint: str) -> Logging:
logging_obj = Logging(
model=CONVERSE_MODEL,
messages=[],
stream=True,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="call-1",
function_id="fn-1",
)
logging_obj.model_call_details["custom_llm_provider"] = "bedrock"
logging_obj.model_call_details["endpoint"] = endpoint
return logging_obj
def _converse_stream_logging_obj() -> Logging:
return _stream_logging_obj(CONVERSE_STREAM_ENDPOINT)
def _stream_collector(endpoint: str) -> PassthroughStreamCollector:
return BedrockPassthroughConfig().create_stream_collector(
model=CONVERSE_MODEL, custom_llm_provider="bedrock", endpoint=endpoint
)
def _converse_stream_collector() -> PassthroughStreamCollector:
return _stream_collector(CONVERSE_STREAM_ENDPOINT)
def _feed(collector: PassthroughStreamCollector, stream: bytes, chunk_size: int = 16384) -> None:
for offset in range(0, len(stream), chunk_size):
collector.add(stream[offset : offset + chunk_size])
def test_converse_stream_collector_keeps_usage_without_retaining_the_stream():
texts = [f"tok{i} " for i in range(4000)]
stream = _event_frame("messageStart", {"role": "assistant"}) + _text_block(0, texts) + _stream_tail("end_turn", 4000)
_feed(_converse_stream_collector(), stream)
tracemalloc.start()
try:
base = tracemalloc.get_traced_memory()[0]
collector = _converse_stream_collector()
_feed(collector, stream)
retained = tracemalloc.get_traced_memory()[0] - base
finally:
tracemalloc.stop()
assert retained < len(stream) // 4
response = collector.build_logged_response(_converse_stream_logging_obj())
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "".join(texts)
assert response.choices[0].finish_reason == "stop"
assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 4000)
def test_converse_stream_collector_keeps_tool_calls_between_text_runs():
stream = (
_event_frame("messageStart", {"role": "assistant"})
+ _text_block(0, ["Let me ", "check."])
+ _event_frame(
"contentBlockStart",
{"contentBlockIndex": 1, "start": {"toolUse": {"toolUseId": "tool-1", "name": "get_weather"}}},
)
+ _event_frame("contentBlockDelta", {"contentBlockIndex": 1, "delta": {"toolUse": {"input": '{"city": '}}})
+ _event_frame("contentBlockDelta", {"contentBlockIndex": 1, "delta": {"toolUse": {"input": '"Paris"}'}}})
+ _event_frame("contentBlockStop", {"contentBlockIndex": 1})
+ _text_block(2, ["Done", "."])
+ _stream_tail("tool_use", 12)
)
collector = _converse_stream_collector()
_feed(collector, stream, chunk_size=7)
response = collector.build_logged_response(_converse_stream_logging_obj())
assert isinstance(response, ModelResponse)
message = response.choices[0].message
assert message.content == "Let me check.Done."
assert [(call.function.name, call.function.arguments) for call in message.tool_calls] == [
("get_weather", '{"city": "Paris"}')
]
assert response.choices[0].finish_reason == "tool_calls"
assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 12)
def test_invoke_stream_collector_keeps_usage_without_retaining_the_stream():
texts = [f"tok{i} " for i in range(4000)]
stream = (
_invoke_chunk(
{
"type": "message_start",
"message": {
"id": "msg-1",
"type": "message",
"role": "assistant",
"model": CONVERSE_MODEL,
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 25, "output_tokens": 1},
},
}
)
+ _invoke_chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}})
+ b"".join(
_invoke_chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}})
for text in texts
)
+ _invoke_chunk({"type": "content_block_stop", "index": 0})
+ _invoke_chunk(
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4000}}
)
+ _invoke_chunk({"type": "message_stop"})
)
_feed(_stream_collector(INVOKE_STREAM_ENDPOINT), stream)
tracemalloc.start()
try:
base = tracemalloc.get_traced_memory()[0]
collector = _stream_collector(INVOKE_STREAM_ENDPOINT)
_feed(collector, stream)
retained = tracemalloc.get_traced_memory()[0] - base
finally:
tracemalloc.stop()
assert retained < len(stream) // 4
response = collector.build_logged_response(_stream_logging_obj(INVOKE_STREAM_ENDPOINT))
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "".join(texts)
assert response.choices[0].finish_reason == "stop"
assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 4000)
def test_stream_collector_logs_nothing_for_an_unrecognized_endpoint():
collector = BedrockPassthroughConfig().create_stream_collector(
model=CONVERSE_MODEL, custom_llm_provider="bedrock", endpoint=f"/model/{CONVERSE_MODEL}/rerank"
)
collector.add(_event_frame("messageStart", {"role": "assistant"}))
assert collector.build_logged_response(_converse_stream_logging_obj()) is None

View file

@ -34,6 +34,34 @@ class _ImmediateExecutor:
fn(*args, **kwargs)
class _RecordingCollector:
def __init__(self) -> None:
self.chunks: List[bytes] = []
def add(self, chunk: bytes) -> None:
self.chunks.append(chunk)
def build_logged_response(self, litellm_logging_obj: MagicMock) -> bytes:
return b"".join(self.chunks)
class _FailingCollector(_RecordingCollector):
def add(self, chunk: bytes) -> None:
raise ValueError("bad frame")
def _provider_config(collector: _RecordingCollector) -> MagicMock:
provider_config = MagicMock()
provider_config.create_stream_collector.return_value = collector
return provider_config
def _spend_payload(flush_mock: MagicMock) -> bytes:
flush_mock.assert_called_once()
collector = flush_mock.call_args.kwargs["collector"]
return collector.build_logged_response(litellm_logging_obj=MagicMock())
@pytest.mark.asyncio
async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -48,13 +76,12 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
received_response = AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
async for chunk in received_response:
@ -67,12 +94,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == chunks
assert call_kwargs["provider_config"] is provider_config
assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == b"".join(chunks)
@pytest.mark.asyncio
@ -93,12 +115,11 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
gen = AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
received = [await gen.__anext__()]
@ -108,11 +129,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect():
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == [chunks[0]]
assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == chunks[0]
@pytest.mark.asyncio
@ -178,14 +195,13 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
async def _drain():
async for chunk in AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
):
received.append(chunk)
@ -196,11 +212,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == partial_chunks
assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == b"".join(partial_chunks)
def test_passthroughstreamingresponse_flushes_on_normal_completion():
@ -221,12 +233,11 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion():
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
received_responce = PassthroughStreamingResponse(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
with patch("litellm.utils.executor", _ImmediateExecutor()):
@ -237,7 +248,7 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion():
assert received_responce.headers["content-type"] == "application/octet-stream"
assert received_responce.headers["x-request-id"] == "req-123"
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
assert _spend_payload(mock_logging_obj.flush_passthrough_collected_chunks) == b"".join(chunks)
def test_passthroughstreamingresponse_flushes_on_early_close():
@ -258,19 +269,66 @@ def test_passthroughstreamingresponse_flushes_on_early_close():
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
with patch("litellm.utils.executor", _ImmediateExecutor()):
gen = PassthroughStreamingResponse(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
first = next(gen)
gen.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 _spend_payload(mock_logging_obj.flush_passthrough_collected_chunks) == chunks[0]
@pytest.mark.asyncio
async def test_asyncpassthroughstreamingresponse_relays_the_stream_when_spend_parsing_fails():
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
mock_response = _make_streaming_response(chunks)
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
received = [
chunk
async for chunk in AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=_provider_config(_FailingCollector()),
)
]
await asyncio.sleep(0)
assert received == chunks
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_not_called()
def test_passthroughstreamingresponse_relays_the_stream_when_spend_parsing_fails():
from litellm.passthrough.main import PassthroughStreamingResponse
chunks = [b"a", b"b", b"c"]
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream"})
mock_response.iter_bytes = lambda: iter(chunks)
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
received = list(
PassthroughStreamingResponse(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=_provider_config(_FailingCollector()),
)
)
assert received == chunks
mock_logging_obj.flush_passthrough_collected_chunks.assert_not_called()