From 47bba14336080810bbc3e0e9b1160438cca1c161 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 09:53:37 -0700 Subject: [PATCH] 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 Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/litellm_logging.py | 51 +---- litellm/llms/anthropic/chat/handler.py | 5 +- .../base_llm/passthrough/transformation.py | 41 +++- litellm/llms/bedrock/chat/invoke_handler.py | 2 +- .../bedrock/passthrough/transformation.py | 199 +++++++++++------- litellm/passthrough/main.py | 62 ++++-- ...est_azure_ai_passthrough_transformation.py | 6 +- ...test_bedrock_passthrough_transformation.py | 197 ++++++++++++++++- ...test_streaming_interrupt_spend_tracking.py | 118 ++++++++--- 9 files changed, 507 insertions(+), 174 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 6ee68ab21c5..498d662a906 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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: diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index d1fe4cadf40..82189461403 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -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 diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index 20180c5cfa2..ec938889b88 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -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() diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 5f8a5544d65..5c489ecb360 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -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: diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index fb8bc4f191f..6a120f41cb6 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -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 diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 7076683f294..73d8bab686b 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -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: diff --git a/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py b/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py index c8007acf70f..f00698a6624 100644 --- a/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py +++ b/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py @@ -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" diff --git a/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py b/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py index b005d77ac8b..f2a9af11af7 100644 --- a/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py +++ b/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py @@ -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 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 a88b0ef0c4b..5e13db9439b 100644 --- a/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py +++ b/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py @@ -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()