From 1a5a856e3e2d1889fdf4f6042eafa89ecc2f84d9 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 28 Aug 2026 16:05:45 -0700 Subject: [PATCH] fix(guardrails): defer native /v1/messages stream logging until post_call scans finish --- .../messages/streaming_iterator.py | 39 +++-- litellm/proxy/common_request_processing.py | 22 ++- litellm/proxy/utils.py | 27 +++- .../messages/test_streaming_iterator.py | 88 +++++++++++- .../proxy_logging/test_streaming_hooks.py | 133 ++++++++++++++++++ 5 files changed, 288 insertions(+), 21 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index 8387dd8310d..e0ddd54c24f 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -342,7 +342,7 @@ class BaseAnthropicMessagesStreamingIterator: self.start_time = datetime.now() self.completion_start_time: datetime | None = None - async def _handle_streaming_logging(self, collected_chunks: list[bytes]): + async def _handle_streaming_logging(self, collected_chunks: list[bytes], *, stream_teardown: bool = False): """Handle the logging after all chunks have been collected.""" from litellm.proxy.pass_through_endpoints.streaming_handler import ( PassThroughStreamingHandler, @@ -354,21 +354,32 @@ class BaseAnthropicMessagesStreamingIterator: if self.completion_start_time is not None: self.litellm_logging_obj.completion_start_time = self.completion_start_time self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time + logging_coroutine: Final = PassThroughStreamingHandler._route_streaming_logging_to_handler( + litellm_logging_obj=self.litellm_logging_obj, + passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, + url_route="/v1/messages", + request_body=self.request_body or {}, + endpoint_type=EndpointType.ANTHROPIC, + start_time=self.start_time, + raw_bytes=collected_chunks, + end_time=end_time, + ) + deferred_dispatch_armed: Final = ( + getattr(self.litellm_logging_obj, "_on_deferred_stream_complete", None) is not None + ) + # Post-call guardrails run their end-of-stream scan AFTER this iterator + # is exhausted, so enqueueing now would build the spend log before the + # scan writes guardrail_information. Park the coroutine instead; the + # proxy fires it via _fire_deferred_stream_logging once the guardrail + # chain drains. Teardown (client disconnect) keeps enqueueing + # immediately: the scan never runs there and billing must not be lost. + if deferred_dispatch_armed and not stream_teardown: + self.litellm_logging_obj._deferred_stream_complete_args = (logging_coroutine,) + return # Enqueue on the rooted logging worker rather than asyncio.create_task: # this also runs during generator teardown after a client disconnect, # where an unrooted task could be garbage-collected before it bills. - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler( - litellm_logging_obj=self.litellm_logging_obj, - passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, - url_route="/v1/messages", - request_body=self.request_body or {}, - endpoint_type=EndpointType.ANTHROPIC, - start_time=self.start_time, - raw_bytes=collected_chunks, - end_time=end_time, - ) - ) + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine) def get_async_streaming_response_iterator( self, @@ -433,7 +444,7 @@ class BaseAnthropicMessagesStreamingIterator: # post-loop logging below never runs and the tokens already streamed # (and billed by the provider) would never reach spend tracking. See LIT-5839. if collected_chunks: - await self._handle_streaming_logging(collected_chunks) + await self._handle_streaming_logging(collected_chunks, stream_teardown=True) raise if not saw_terminal_event: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index cbd50da9c3e..936ca1607b5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -4,7 +4,7 @@ import json import logging import math import traceback -from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType @@ -2372,6 +2372,26 @@ class ProxyBaseLLMRequestProcessing: ) logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete + elif ( + _post_call_guardrails_active + and route_type == "anthropic_messages" + and self._is_streaming_response(response) + ): + # Native /v1/messages SSE streams bypass CSW, so the raw + # iterator parks its logging coroutine at stream end (see + # BaseAnthropicMessagesStreamingIterator._handle_streaming_logging) + # and _fire_deferred_stream_logging hands it here after the + # guardrail end-of-stream scans complete. + from litellm.litellm_core_utils.logging_worker import ( + GLOBAL_LOGGING_WORKER, + ) + + async def _on_deferred_native_stream_complete( + logging_coroutine: Coroutine[object, object, object], + ) -> None: + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine) + + logging_obj._on_deferred_stream_complete = _on_deferred_native_stream_complete if route_type == "allm_passthrough_route": # Check if response is an async generator diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d880b529727..516a50610d0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3167,8 +3167,14 @@ class ProxyLogging: # through each of them adds N pass-through trampolines per chunk for # zero behavior change. Skip the chain entirely and stream through. if not caps.iterator_overrides: - async for chunk in response: - yield chunk + try: + async for chunk in response: + yield chunk + except (GeneratorExit, asyncio.CancelledError): + raise + except Exception: + ProxyLogging._fire_deferred_stream_logging(request_data) + raise ProxyLogging._fire_deferred_stream_logging(request_data) return @@ -3221,9 +3227,20 @@ class ProxyLogging: ), ) - # Actually iterate through the chained async generator and yield chunks - async for chunk in current_response: - yield chunk + # Actually iterate through the chained async generator and yield chunks. + # A guardrail block raised after upstream exhaustion (e.g. + # unified_guardrail re-raising HTTPException) must still flush any + # parked deferred logging, or the blocked stream loses its spend log. + # GeneratorExit/CancelledError stay untouched: disconnect cleanup owns + # those. + try: + async for chunk in current_response: + yield chunk + except (GeneratorExit, asyncio.CancelledError): + raise + except Exception: + ProxyLogging._fire_deferred_stream_logging(request_data) + raise # Fire deferred logging AFTER all guardrail end-of-stream blocks # completed. unified_guardrail writes guardrail_information during diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py index cb31280c2d5..b8ce11db8d1 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py @@ -6,6 +6,9 @@ import pytest from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.anthropic.experimental_pass_through.messages import ( + streaming_iterator as streaming_iterator_module, +) from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( INCOMPLETE_STREAM_ERROR_MESSAGE, AnthropicMessagesStreamHiddenParams, @@ -26,7 +29,7 @@ class _RecordingLoggingIterator(BaseAnthropicMessagesStreamingIterator): self.logged_chunks: list = [] self.logging_call_count: int = 0 - async def _handle_streaming_logging(self, collected_chunks): + async def _handle_streaming_logging(self, collected_chunks, *, stream_teardown=False): self.logged_chunks = list(collected_chunks) self.logging_call_count += 1 @@ -543,3 +546,86 @@ def test_anthropic_messages_response_as_sse_events_no_content_blocks(): response = {"id": "msg_4", "content": [], "stop_reason": "end_turn"} decoded = _decode_sse_events(anthropic_messages_response_as_sse_events(response)) assert [event_type for event_type, _ in decoded] == ["message_start", "message_delta", "message_stop"] + + +class _RecordingLoggingWorker: + def __init__(self): + self.enqueued = [] + + def ensure_initialized_and_enqueue(self, async_coroutine): + self.enqueued.append(async_coroutine) + + def close_enqueued(self): + for coroutine in self.enqueued: + coroutine.close() + + +async def _noop_deferred_dispatch(logging_coroutine): + logging_coroutine.close() + + +async def _stream_of(events): + for event in events: + yield event + + +COMPLETE_STREAM_EVENTS = TRUNCATED_TOOL_USE_EVENTS + ({"type": "message_stop"},) + + +@pytest.mark.asyncio +async def test_normal_end_with_deferred_dispatch_armed_parks_logging_coroutine(monkeypatch): + """ + Regression test for LIT-6409: with post_call guardrails active the proxy + arms logging_obj._on_deferred_stream_complete, and the native /v1/messages + iterator must park its logging coroutine instead of enqueueing it at + upstream exhaustion, otherwise the spend log is built before the + guardrail end-of-stream scan writes its post_call entry. + """ + worker = _RecordingLoggingWorker() + monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker) + iterator = _make_iterator("test_deferred_parks_logging_coroutine") + iterator.litellm_logging_obj._on_deferred_stream_complete = _noop_deferred_dispatch + + await _collect(iterator, _stream_of(COMPLETE_STREAM_EVENTS)) + + parked = getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) + assert worker.enqueued == [] + assert parked is not None + assert len(parked) == 1 + assert asyncio.iscoroutine(parked[0]) + parked[0].close() + + +@pytest.mark.asyncio +async def test_client_disconnect_enqueues_immediately_even_when_deferred_dispatch_armed(monkeypatch): + """ + On client disconnect the guardrail end-of-stream scan never runs, so + deferral would strand the spend log; the teardown path must keep + enqueueing immediately (LIT-5839) even when the deferred callback is armed. + """ + worker = _RecordingLoggingWorker() + monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker) + iterator = _make_iterator("test_disconnect_enqueues_when_armed") + iterator.litellm_logging_obj._on_deferred_stream_complete = _noop_deferred_dispatch + + wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS)) + for _ in range(len(TRUNCATED_TOOL_USE_EVENTS)): + await wrapped.__anext__() + await wrapped.aclose() + + assert len(worker.enqueued) == 1 + assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None + worker.close_enqueued() + + +@pytest.mark.asyncio +async def test_normal_end_without_deferred_dispatch_enqueues_immediately(monkeypatch): + worker = _RecordingLoggingWorker() + monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker) + iterator = _make_iterator("test_unarmed_enqueues_at_stream_end") + + await _collect(iterator, _stream_of(COMPLETE_STREAM_EVENTS)) + + assert len(worker.enqueued) == 1 + assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None + worker.close_enqueued() diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index 65d3c3c8079..ec5b994f147 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -10,6 +10,7 @@ Covers ``_wrap_streaming_iterator_with_enrichment``, from __future__ import annotations import asyncio +from datetime import datetime from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock @@ -18,6 +19,10 @@ from fastapi import HTTPException import litellm from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + BaseAnthropicMessagesStreamingIterator, +) from litellm.proxy.utils import ProxyLogging @@ -346,6 +351,134 @@ async def test_async_post_call_streaming_iterator_hook_upstream_error_raises(pro pass +# --------------------------------------------------------------------------- +# deferred native /v1/messages stream logging (LIT-6409) +# --------------------------------------------------------------------------- + + +_NATIVE_MESSAGES_STREAM_EVENTS = ( + {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 3, "output_tokens": 1}}}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}}, + {"type": "message_stop"}, +) + + +def _armed_native_messages_stream(test_name: str, request_data: Dict[str, Any], events: List[Any]): + """The proxy-side setup for a native /v1/messages stream with post_call + guardrails active: a real BaseAnthropicMessagesStreamingIterator whose + logging_obj carries the deferred-dispatch callback the proxy arms in + common_request_processing. The callback records what the guardrail + metadata contained at the moment the deferred logging was dispatched.""" + logging_obj = LiteLLMLoggingObj( + model="bedrock/invoke/anthropic.claude-sonnet-4-20250514-v1:0", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id=test_name, + function_id=test_name, + ) + + async def _dispatch_deferred_logging(logging_coroutine): + events.append( + ( + "logging_dispatched", + "post_call_entry_visible", + bool(request_data.get("metadata", {}).get("standard_logging_guardrail_information")), + ) + ) + logging_coroutine.close() + + logging_obj._on_deferred_stream_complete = _dispatch_deferred_logging + request_data["litellm_logging_obj"] = logging_obj + + iterator = BaseAnthropicMessagesStreamingIterator(litellm_logging_obj=logging_obj, request_body={}) + + async def _upstream(): + for event in _NATIVE_MESSAGES_STREAM_EVENTS: + yield event + + return logging_obj, iterator.async_sse_wrapper(_upstream()) + + +@pytest.mark.asyncio +async def test_native_messages_stream_logging_fires_after_guardrail_end_of_stream_scan( + proxy_logging, make_user_api_key_auth, monkeypatch +): + """ + Regression test for LIT-6409: on native /v1/messages streams the + end-of-stream guardrail scan writes its post_call entry AFTER the + upstream iterator is exhausted, so success logging dispatched at + upstream exhaustion never sees it. The deferred dispatch must fire + only after the guardrail chain fully drains. + """ + events: List[Any] = [] + request_data: Dict[str, Any] = {"metadata": {}} + _, native_stream = _armed_native_messages_stream( + "test_native_stream_deferred_ordering", request_data, events + ) + + class _EndOfStreamScanGuardrail(CustomLogger): + async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data): + async for chunk in response: + yield chunk + request_data.setdefault("metadata", {})["standard_logging_guardrail_information"] = [ + {"guardrail_mode": "post_call", "guardrail_status": "success"} + ] + events.append("scan_appended") + + monkeypatch.setattr(litellm, "callbacks", [_EndOfStreamScanGuardrail()]) + + async for _ in proxy_logging.async_post_call_streaming_iterator_hook( + response=native_stream, + user_api_key_dict=make_user_api_key_auth(), + request_data=request_data, + ): + pass + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert events == ["scan_appended", ("logging_dispatched", "post_call_entry_visible", True)] + + +@pytest.mark.asyncio +async def test_native_messages_stream_logging_fires_when_guardrail_blocks_after_stream_end( + proxy_logging, make_user_api_key_auth, monkeypatch +): + """ + A guardrail block raised after upstream exhaustion (unified_guardrail + re-raises HTTPException for blocked content) must still flush the + parked deferred logging, or the blocked stream loses its spend log. + """ + events: List[Any] = [] + request_data: Dict[str, Any] = {"metadata": {}} + logging_obj, native_stream = _armed_native_messages_stream( + "test_native_stream_deferred_block", request_data, events + ) + + class _BlockingGuardrail(CustomLogger): + async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data): + async for chunk in response: + yield chunk + raise HTTPException(status_code=400, detail={"error": "Violated guardrail policy"}) + + monkeypatch.setattr(litellm, "callbacks", [_BlockingGuardrail()]) + + with pytest.raises(HTTPException): + async for _ in proxy_logging.async_post_call_streaming_iterator_hook( + response=native_stream, + user_api_key_dict=make_user_api_key_auth(), + request_data=request_data, + ): + pass + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert [event[0] for event in events] == ["logging_dispatched"] + assert logging_obj._deferred_stream_complete_args is None + + # --------------------------------------------------------------------------- # _fire_deferred_stream_logging # ---------------------------------------------------------------------------