fix(guardrails): match deferred stream dispatch shape per stream owner and defer passthrough logging until guardrail eos

This commit is contained in:
mateo-berri 2026-08-29 02:23:12 -07:00
parent 64eec53fd8
commit f60ccf6234
6 changed files with 381 additions and 79 deletions

View file

@ -2341,53 +2341,14 @@ class ProxyBaseLLMRequestProcessing:
if requested_model_from_client:
self.data["_litellm_client_requested_model"] = requested_model_from_client
# Streaming: attach a closure that fires after all guardrail
# end-of-stream blocks complete. CSW.__anext__ stores the
# assembled response on logging_obj; the outer consumer
# (ProxyLogging._fire_deferred_stream_logging) fires the
# closure after the full streaming pipeline finishes.
# The closure runs non-apply_guardrail hooks on the
# assembled response, then fires success logging.
# Only for CustomStreamWrapper — raw async generators from
# passthrough routes bypass CSW and would orphan the closure.
from litellm.litellm_core_utils.streaming_handler import (
CustomStreamWrapper,
)
if _post_call_guardrails_active and isinstance(response, CustomStreamWrapper):
# Intentionally a live reference (not a copy) — mirrors
# ProxyLogging.post_call_success_hook which also mutates
# data["guardrail_to_apply"] during iteration.
_captured_data: Final = self.data
_captured_user_api_key_dict: Final = user_api_key_dict
_captured_logging_obj: Final = logging_obj
async def _on_deferred_stream_complete(assembled_response: object, cache_hit: object) -> None:
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=_captured_data,
captured_user_api_key_dict=_captured_user_api_key_dict,
captured_logging_obj=_captured_logging_obj,
assembled_response=assembled_response,
cache_hit=cache_hit,
)
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
elif (
_post_call_guardrails_active
and route_type in ("anthropic_messages", "aresponses")
and self._is_streaming_response(response)
):
from litellm.litellm_core_utils.logging_worker import (
GLOBAL_LOGGING_WORKER,
if _post_call_guardrails_active:
self._arm_deferred_stream_dispatch(
response=response,
route_type=route_type,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
)
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
if self._is_streaming_response(response):
@ -3057,6 +3018,87 @@ class ProxyBaseLLMRequestProcessing:
except Exception as e:
verbose_proxy_logger.exception("Error firing deferred logging: %s", e)
def _arm_deferred_stream_dispatch(
self,
response: object,
route_type: str,
user_api_key_dict: "UserAPIKeyAuth",
logging_obj: LiteLLMLoggingObj,
) -> None:
"""
Streaming with post-call guardrails active: attach a closure that
ProxyLogging._fire_deferred_stream_logging fires after all guardrail
end-of-stream blocks complete, so the spend log sees
guardrail_information.
Three closure shapes, matching who owns logging for the stream:
- CustomStreamWrapper (chat completions) stores
(assembled_response, cache_hit); the closure also runs
non-apply_guardrail post-call hooks via
_run_deferred_stream_guardrails.
- Bridged /v1/responses (LiteLLMCompletionStreamingIterator) shares
its inner CustomStreamWrapper's logging_obj, so it stores the same
(assembled_response, cache_hit) shape; the closure only dispatches
success logging, matching the route's pre-existing hook surface.
- Native anthropic_messages/aresponses iterators store a single
ready-made logging coroutine to enqueue.
Raw async generators from passthrough routes bypass all three and
would orphan the closure, so they are not armed here.
"""
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
if isinstance(response, CustomStreamWrapper):
# Intentionally a live reference (not a copy) — mirrors
# ProxyLogging.post_call_success_hook which also mutates
# data["guardrail_to_apply"] during iteration.
_captured_data: Final = self.data
_captured_user_api_key_dict: Final = user_api_key_dict
_captured_logging_obj: Final = logging_obj
async def _on_deferred_stream_complete(assembled_response: object, cache_hit: object) -> None:
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=_captured_data,
captured_user_api_key_dict=_captured_user_api_key_dict,
captured_logging_obj=_captured_logging_obj,
assembled_response=assembled_response,
cache_hit=cache_hit,
)
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
return
if route_type not in ("anthropic_messages", "aresponses") or not self._is_streaming_response(response):
return
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
if isinstance(response, LiteLLMCompletionStreamingIterator):
_captured_bridge_logging_obj: Final = logging_obj
async def _on_deferred_bridged_stream_complete(assembled_response: object, cache_hit: object) -> None:
await _as_success_dispatcher(_captured_bridge_logging_obj).dispatch_success_handlers(
assembled_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
prefer_async_handlers=True,
)
logging_obj._on_deferred_stream_complete = _on_deferred_bridged_stream_complete
return
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
@staticmethod
async def _run_deferred_stream_guardrails(
captured_data: dict,

View file

@ -65,6 +65,19 @@ class PassThroughStreamingHandler:
route_streaming_logging or PassThroughStreamingHandler._route_streaming_logging_to_handler
)
raw_bytes: Final[list[bytes]] = []
def _build_logging_coroutine() -> Coroutine[None, None, None]:
return resolved_route_streaming_logging(
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body or {},
endpoint_type=endpoint_type,
start_time=start_time,
raw_bytes=raw_bytes,
end_time=datetime.now(),
)
logging_scheduled = False
model_name: Final = PassThroughStreamingHandler._extract_model_for_cost_injection(
request_body=request_body,
@ -114,6 +127,21 @@ class PassThroughStreamingHandler:
)
if pending:
yield pending
# Stream completed cleanly. When the proxy armed deferred
# dispatch (post-call guardrails active), park the logging
# coroutine on logging_obj instead of enqueueing now, so
# ProxyLogging._fire_deferred_stream_logging fires it after
# guardrail end-of-stream blocks populate guardrail_information.
# Disconnect/exception paths skip this and fall through to the
# immediate enqueue in ``finally`` to keep partial billing
# (LIT-2642).
if (
getattr(litellm_logging_obj, "_on_deferred_stream_complete", None) is not None
and raw_bytes
and response.status_code < 400
):
logging_scheduled = True
litellm_logging_obj._deferred_stream_complete_args = (_build_logging_coroutine(),)
except Exception as e:
verbose_proxy_logger.error("Error in chunk_processor: %s", e)
raise
@ -128,18 +156,7 @@ class PassThroughStreamingHandler:
if not logging_scheduled and raw_bytes and response.status_code < 400:
logging_scheduled = True
try:
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=resolved_route_streaming_logging(
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body or {},
endpoint_type=endpoint_type,
start_time=start_time,
raw_bytes=raw_bytes,
end_time=datetime.now(),
)
)
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=_build_logging_coroutine())
except Exception as e:
verbose_proxy_logger.error("Error scheduling chunk_processor logging: %s", e)

View file

@ -1297,3 +1297,126 @@ class TestResponsesIteratorDeferredLogging:
assert len(created) == 1
await created[0]
assert recorded["dispatched"] is True
class TestArmDeferredStreamDispatch:
"""Regression for PR #38722: the closure shape armed on logging_obj must
match the args the stream's logging owner stores. Bridged /v1/responses
(LiteLLMCompletionStreamingIterator) shares its inner CustomStreamWrapper's
logging_obj, which stores (assembled_response, cache_hit); arming the
single-coroutine native closure there made _fire_deferred_stream_logging
raise TypeError inside the streaming hook, leaking an in-stream 500 error
frame on every streamed /v1/responses request."""
def _processor(self):
return ProxyBaseLLMRequestProcessing(data={"model": "gpt-test"})
def _dispatch_recording_logging_obj(self):
recorded = {}
async def dispatch_success_handlers(
result=None, start_time=None, end_time=None, cache_hit=None, prefer_async_handlers=False
):
recorded["result"] = result
recorded["cache_hit"] = cache_hit
recorded["prefer_async_handlers"] = prefer_async_handlers
logging_obj = MagicMock()
logging_obj.dispatch_success_handlers = dispatch_success_handlers
logging_obj._on_deferred_stream_complete = None
logging_obj._deferred_stream_complete_args = None
return logging_obj, recorded
@pytest.mark.asyncio
async def test_bridged_responses_iterator_gets_csw_arg_shape(self):
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
logging_obj, recorded = self._dispatch_recording_logging_obj()
bridged = object.__new__(LiteLLMCompletionStreamingIterator)
self._processor()._arm_deferred_stream_dispatch(
response=bridged,
route_type="aresponses",
user_api_key_dict=MagicMock(),
logging_obj=logging_obj,
)
assembled = object()
logging_obj._deferred_stream_complete_args = (assembled, False)
ProxyLogging._fire_deferred_stream_logging({"litellm_logging_obj": logging_obj})
await asyncio.sleep(0)
assert recorded["result"] is assembled
assert recorded["cache_hit"] is False
assert recorded["prefer_async_handlers"] is True
@pytest.mark.asyncio
async def test_native_stream_closure_enqueues_single_coroutine(self):
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
logging_obj, _ = self._dispatch_recording_logging_obj()
async def _agen():
yield b"x"
self._processor()._arm_deferred_stream_dispatch(
response=_agen(),
route_type="anthropic_messages",
user_api_key_dict=MagicMock(),
logging_obj=logging_obj,
)
closure = logging_obj._on_deferred_stream_complete
assert closure is not None
async def _logging_coroutine():
return None
coro = _logging_coroutine()
with patch.object( # test-quality-ok: GLOBAL_LOGGING_WORKER is a process-global singleton with no injection seam
GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue"
) as mock_enqueue:
await closure(coro)
mock_enqueue.assert_called_once_with(async_coroutine=coro)
coro.close()
@pytest.mark.asyncio
async def test_csw_closure_routes_through_deferred_stream_guardrails(self, monkeypatch):
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
logging_obj, recorded = self._dispatch_recording_logging_obj()
csw = object.__new__(CustomStreamWrapper)
processor = self._processor()
monkeypatch.setattr( # test-quality-ok: empty the process-global callback registry so no ambient guardrail runs
litellm, "callbacks", []
)
processor._arm_deferred_stream_dispatch(
response=csw,
route_type="acompletion",
user_api_key_dict=MagicMock(),
logging_obj=logging_obj,
)
assembled = object()
await logging_obj._on_deferred_stream_complete(assembled, False)
await asyncio.sleep(0)
assert recorded["result"] is assembled
assert recorded["cache_hit"] is False
assert recorded["prefer_async_handlers"] is True
def test_non_native_route_generator_not_armed(self):
logging_obj, _ = self._dispatch_recording_logging_obj()
async def _agen():
yield b"x"
self._processor()._arm_deferred_stream_dispatch(
response=_agen(),
route_type="acompletion",
user_api_key_dict=MagicMock(),
logging_obj=logging_obj,
)
assert logging_obj._on_deferred_stream_complete is None

View file

@ -28,12 +28,21 @@ def _make_streaming_response(chunks):
return mock
def _unarmed_logging_obj():
"""Real Logging objects only carry _on_deferred_stream_complete when the
proxy arms deferred dispatch; a bare MagicMock's auto-attribute is truthy
and would spuriously trigger the deferral branch."""
obj = MagicMock()
obj._on_deferred_stream_complete = None
return obj
@pytest.mark.asyncio
async def test_chunk_processor_logs_on_normal_completion():
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
response = _make_streaming_response(chunks)
mock_logging_obj = MagicMock()
mock_logging_obj = _unarmed_logging_obj()
mock_passthrough_handler = MagicMock()
with patch.object(
@ -66,7 +75,7 @@ async def test_chunk_processor_logs_on_client_disconnect():
chunks = [b"event-1", b"event-2", b"event-3"]
response = _make_streaming_response(chunks)
mock_logging_obj = MagicMock()
mock_logging_obj = _unarmed_logging_obj()
mock_passthrough_handler = MagicMock()
with patch.object(
@ -104,7 +113,7 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er
response = _make_streaming_response(chunks)
response.status_code = 403
mock_logging_obj = MagicMock()
mock_logging_obj = _unarmed_logging_obj()
mock_passthrough_handler = MagicMock()
with patch.object(
@ -134,7 +143,7 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er
async def test_chunk_processor_does_not_schedule_logging_when_no_chunks():
response = _make_streaming_response([])
mock_logging_obj = MagicMock()
mock_logging_obj = _unarmed_logging_obj()
mock_passthrough_handler = MagicMock()
with patch.object(
@ -189,7 +198,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker():
async for chunk in PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=MagicMock(),
litellm_logging_obj=_unarmed_logging_obj(),
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
@ -230,7 +239,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne
gen = PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=MagicMock(),
litellm_logging_obj=_unarmed_logging_obj(),
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
@ -246,7 +255,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne
def _logging_obj_with_write_once_cst():
"""Build a MagicMock that mirrors the real Logging behavior: _update_completion_start_time
latches self.completion_start_time so the write-once guard actually latches."""
obj = MagicMock()
obj = _unarmed_logging_obj()
obj.completion_start_time = None
def _update(*, completion_start_time):
@ -301,7 +310,7 @@ async def test_chunk_processor_does_not_reset_completion_start_time_on_later_chu
response = _make_streaming_response(chunks)
real_first = datetime(2020, 1, 1, 0, 0, 0)
mock_logging_obj = MagicMock()
mock_logging_obj = _unarmed_logging_obj()
# Simulate first-chunk stamp having already landed (e.g. under contention or a
# prior wrapper that already set it): later chunks must be no-ops.
mock_logging_obj.completion_start_time = real_first
@ -387,7 +396,7 @@ async def _collect_openai_passthrough_chunks(chunks, endpoint_type):
async for chunk in PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "gpt-4o-mini", "stream": True},
litellm_logging_obj=MagicMock(),
litellm_logging_obj=_unarmed_logging_obj(),
endpoint_type=endpoint_type,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
@ -517,3 +526,109 @@ def test_convert_raw_bytes_survives_truncated_multibyte_sequence():
lines = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(raw_bytes)
assert any('"type": "message_delta"' in line for line in lines)
@pytest.mark.asyncio
async def test_chunk_processor_defers_logging_until_fire_when_armed():
"""Regression for PR #38722: native /v1/messages streams route through
chunk_processor, which enqueued the spend log the moment the stream ended,
racing the guardrail end-of-stream scan and logging
guardrail_information as null. With deferred dispatch armed, the completed
stream must park the logging coroutine on logging_obj and only enqueue it
when ProxyLogging._fire_deferred_stream_logging fires after the scan."""
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import ProxyLogging
chunks = [b"event-1", b"event-2"]
response = _make_streaming_response(chunks)
logging_obj = _unarmed_logging_obj()
logging_obj._deferred_stream_complete_args = None
enqueued = []
def _capture(async_coroutine):
enqueued.append(async_coroutine)
async_coroutine.close()
with patch.object( # test-quality-ok: GLOBAL_LOGGING_WORKER is a process-global singleton with no injection seam
GLOBAL_LOGGING_WORKER,
"ensure_initialized_and_enqueue",
side_effect=_capture,
) as mock_enqueue:
gen = PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=logging_obj,
endpoint_type=EndpointType.ANTHROPIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
url_route="/v1/messages",
route_streaming_logging=AsyncMock(),
)
ProxyBaseLLMRequestProcessing(data={})._arm_deferred_stream_dispatch(
response=gen,
route_type="anthropic_messages",
user_api_key_dict=MagicMock(),
logging_obj=logging_obj,
)
received = []
async for chunk in gen:
received.append(chunk)
await asyncio.sleep(0)
assert received == chunks
mock_enqueue.assert_not_called()
parked = logging_obj._deferred_stream_complete_args
assert isinstance(parked, tuple) and len(parked) == 1
assert asyncio.iscoroutine(parked[0])
ProxyLogging._fire_deferred_stream_logging({"litellm_logging_obj": logging_obj})
await asyncio.sleep(0)
mock_enqueue.assert_called_once()
@pytest.mark.asyncio
async def test_chunk_processor_enqueues_immediately_on_disconnect_even_when_armed():
"""Client disconnects never reach _fire_deferred_stream_logging, so parking
the coroutine there would lose the partial-usage spend log (LIT-2642); the
disconnect path must keep enqueueing immediately."""
chunks = [b"event-1", b"event-2", b"event-3"]
response = _make_streaming_response(chunks)
logging_obj = _unarmed_logging_obj()
async def _armed_closure(logging_coroutine):
raise AssertionError("deferred closure must not fire on disconnect")
logging_obj._on_deferred_stream_complete = _armed_closure
logging_obj._deferred_stream_complete_args = None
enqueued = []
def _capture(async_coroutine):
enqueued.append(async_coroutine)
async_coroutine.close()
with patch.object( # test-quality-ok: GLOBAL_LOGGING_WORKER is a process-global singleton with no injection seam
GLOBAL_LOGGING_WORKER,
"ensure_initialized_and_enqueue",
side_effect=_capture,
) as mock_enqueue:
gen = PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=logging_obj,
endpoint_type=EndpointType.ANTHROPIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
url_route="/v1/messages",
route_streaming_logging=AsyncMock(),
)
await gen.__anext__()
await gen.aclose()
mock_enqueue.assert_called_once()
assert logging_obj._deferred_stream_complete_args is None

View file

@ -1656,7 +1656,7 @@ class TestCommonRequestProcessingHelpers:
assert payload["error"]["message"] == "MCP request blocked: no rewritable argument field present"
assert payload["error"]["provider_specific_fields"]["error"]["code"] == "panw_prisma_airs_blocked"
async def testserialize_http_exception_detail_helper(self):
async def test_serialize_http_exception_detail_helper(self):
"""Direct unit coverage for the L1 helper across all branches."""
from litellm.proxy.common_request_processing import (
serialize_http_exception_detail,

View file

@ -346,14 +346,14 @@ async def test_post_call_stream_guardrail_keeps_own_iterator_on_chat_completions
@pytest.mark.asyncio
async def test_unified_guardrail_iterator_accepts_explicit_guardrail(monkeypatch):
async def test_unified_guardrail_iterator_accepts_explicit_guardrail():
"""
The dispatch passes each guardrail explicitly instead of through a shared
request_data key, so chaining two unified-routed guardrails cannot drop
all but the last one.
all but the last one. The block fires after the deltas were already
flushed to the client, so it surfaces as a trailing in-stream error frame
rather than a raised HTTPException.
"""
from fastapi import HTTPException
from litellm.proxy.utils import unified_guardrail
guardrail = _content_filter_guardrail("BLOCK")
@ -367,14 +367,19 @@ async def test_unified_guardrail_iterator_accepts_explicit_guardrail(monkeypatch
for chunk in _anthropic_stream_chunks(["the", " zebra runs"]):
yield chunk
with pytest.raises(HTTPException):
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/messages"),
response=fake_stream(),
request_data=request_data,
guardrail_to_apply=guardrail,
):
pass
delivered = []
async for item in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/messages"),
response=fake_stream(),
request_data=request_data,
guardrail_to_apply=guardrail,
):
delivered.append(item)
raw = b"".join(c for c in delivered if isinstance(c, bytes)).decode()
assert "event: error" in raw
assert "guardrail_error" in raw
assert raw.index("guardrail_error") > raw.index(" zebra runs")
@pytest.mark.asyncio