fix(proxy): attribute only the raising layer in stream and pipeline blocks

The streaming wrapper caught every exception crossing its boundary and named its own
callback, so a block by an inner guardrail or a provider stream failure also named every
outer guardrail. The wrapper now runs the hook over an upstream boundary that remembers
the exception it raised, and skips attribution when the same exception passes through

Pipeline blocks converted from SensitiveDataRouteException or ModifyResponseException into
a generic guardrail_pipeline_error now still record the blocking step's guardrail

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-17 18:58:21 +00:00
parent 9fa85c5da5
commit 73dea4c567
3 changed files with 223 additions and 64 deletions

View file

@ -12,13 +12,36 @@ import sys
import threading
import time
import traceback
from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence
from collections.abc import (
AsyncGenerator,
AsyncIterable,
AsyncIterator,
Awaitable,
Callable,
Coroutine,
Mapping,
Sequence,
)
from dataclasses import dataclass, field
from datetime import date, datetime, timedelta, timezone
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, Protocol, TypeVar, Union, cast, overload
from typing import (
TYPE_CHECKING,
Any,
ClassVar,
Final,
Generic,
Literal,
Optional,
Protocol,
TypeVar,
Union,
cast,
overload,
)
from typing_extensions import ReadOnly, TypedDict
@ -444,6 +467,33 @@ def _record_raising_guardrail(request_data: Mapping[str, object], callback: obje
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=guardrail_name)
class _UpstreamStreamBoundary(Generic[_T]):
"""Remembers the exception the upstream iterator raised, so the wrapper around a
streaming hook can tell a pass-through failure from one the hook raised itself."""
__slots__ = ("_upstream", "failure")
def __init__(self, upstream: AsyncIterable[_T]) -> None:
self._upstream: Final = upstream.__aiter__()
self.failure: BaseException | None = None
def __aiter__(self) -> "_UpstreamStreamBoundary[_T]":
return self
async def __anext__(self) -> _T:
try:
return await self._upstream.__anext__()
except StopAsyncIteration:
raise
except Exception as e:
self.failure = e
raise
class _StreamIteratorHook(Protocol[_T]):
def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ...
def _is_client_error_exception(exc: Exception) -> bool:
if isinstance(exc, HTTPException):
return exc.status_code < 500
@ -2037,14 +2087,18 @@ class ProxyLogging:
_merge_pipeline_metadata_writes(data, result.modified_data)
if result.terminal_action == "block":
blocking_step: Final = result.step_results[-1] if result.step_results else None
callback: Final = (
PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name)
if blocking_step is not None
else None
)
if callback is not None:
_record_raising_guardrail(data, callback)
original_exception: Final = result.original_exception
if original_exception is not None and not _exception_changes_request_flow(original_exception):
blocking_step: Final = result.step_results[-1] if result.step_results else None
if blocking_step is not None:
callback: Final = PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name)
if callback is not None:
_enrich_http_exception_with_guardrail_context(original_exception, callback)
_record_raising_guardrail(data, callback)
if callback is not None:
_enrich_http_exception_with_guardrail_context(original_exception, callback)
raise original_exception
step_results_serializable: Final = [
@ -2490,23 +2544,26 @@ class ProxyLogging:
@staticmethod
async def _wrap_streaming_iterator_with_enrichment(
callback: object,
gen: AsyncGenerator[_T, None],
response: AsyncIterable[_T],
hook: _StreamIteratorHook[_T],
request_data: Mapping[str, object],
) -> AsyncGenerator[_T, None]:
"""
Yield from `gen`; if iteration raises an HTTPException with dict detail,
enrich the detail with the originating callback's `guardrail_name` and
`guardrail_mode` before re-raising. Used to wrap each layer of the
async_post_call_streaming_iterator_hook chain so the enrichment is
attributed to the callback that produced the chunk pipeline at that
point in the chain.
Run `hook` over `response` and yield its chunks. If the hook itself raises,
enrich an HTTPException's dict detail with the callback's `guardrail_name`
and `guardrail_mode` and record the callback in `applied_guardrails` before
re-raising. Failures raised by `response` (the provider stream or an inner
layer of the async_post_call_streaming_iterator_hook chain) pass through
untouched, so only the layer that actually raised is attributed.
"""
upstream: Final = _UpstreamStreamBoundary(response)
try:
async for chunk in gen:
async for chunk in hook(response=upstream):
yield chunk
except Exception as e:
_enrich_http_exception_with_guardrail_context(e, callback)
_record_raising_guardrail(request_data, callback)
if e is not upstream.failure:
_enrich_http_exception_with_guardrail_context(e, callback)
_record_raising_guardrail(request_data, callback)
raise
# Cache for callback-capability detection. Keyed on a signature of
@ -3663,29 +3720,27 @@ class ProxyLogging:
)
else kind
)
if effective_kind == "override":
current_response = self._wrap_streaming_iterator_with_enrichment(
resolved_callback,
resolved_callback.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=current_response,
request_data=request_data,
),
hook: _StreamIteratorHook[object] = (
partial(
resolved_callback.async_post_call_streaming_iterator_hook,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
else:
# kind == "apply_guardrail": route through unified_guardrail
current_response = self._wrap_streaming_iterator_with_enrichment(
resolved_callback,
unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
request_data=request_data,
response=current_response,
guardrail_to_apply=resolved_callback,
buffer_until_moderated_default=(kind == "override"),
),
if effective_kind == "override"
else partial(
unified_guardrail.async_post_call_streaming_iterator_hook,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
guardrail_to_apply=resolved_callback,
buffer_until_moderated_default=(kind == "override"),
)
)
current_response = self._wrap_streaming_iterator_with_enrichment(
resolved_callback,
current_response,
hook,
request_data=request_data,
)
pipeline_translation: Final = (
resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None

View file

@ -551,14 +551,23 @@ def test_handle_pipeline_result_block_does_not_reraise_sensitive_data_route():
session_id="sess-1",
guardrail_name="pii-router",
)
cb = _make_guardrail()
cb.guardrail_name = "pii-router"
result = MagicMock()
result.terminal_action = "block"
result.step_results = [MagicMock(guardrail_name="pii-router")]
result.original_exception = original
with pytest.raises(HTTPException) as info:
ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p")
data: dict[str, object] = {"model": "m"}
saved = litellm.callbacks
litellm.callbacks = [cb]
try:
with pytest.raises(HTTPException) as info:
ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p")
finally:
litellm.callbacks = saved
assert info.value.status_code == 400
assert info.value.detail["error"]["type"] == "guardrail_pipeline_error"
assert data["metadata"] == {"applied_guardrails": ["pii-router"]}
def test_handle_pipeline_result_block_does_not_reraise_modify_response():
@ -571,14 +580,23 @@ def test_handle_pipeline_result_block_does_not_reraise_modify_response():
request_data={"model": "m"},
guardrail_name="masker",
)
cb = _make_guardrail()
cb.guardrail_name = "masker"
result = MagicMock()
result.terminal_action = "block"
result.step_results = [MagicMock(guardrail_name="masker")]
result.original_exception = original
with pytest.raises(HTTPException) as info:
ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p")
data: dict[str, object] = {"model": "m"}
saved = litellm.callbacks
litellm.callbacks = [cb]
try:
with pytest.raises(HTTPException) as info:
ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p")
finally:
litellm.callbacks = saved
assert info.value.status_code == 400
assert info.value.detail["error"]["type"] == "guardrail_pipeline_error"
assert data["metadata"] == {"applied_guardrails": ["masker"]}
def test_handle_pipeline_result_modify_response_raises_modify_exception():

View file

@ -170,6 +170,15 @@ def test_init_response_taking_too_long_task_no_slack_instance_no_error_raises(pr
# ---------------------------------------------------------------------------
async def _passthrough_hook(*, response: AsyncIterator[object]) -> AsyncGenerator[object, None]:
async for chunk in response:
yield chunk
async def _one_chunk() -> AsyncGenerator[object, None]:
yield "chunk"
@pytest.mark.asyncio
async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(proxy_logging):
async def gen():
@ -177,7 +186,9 @@ async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(pro
yield ch
cb = MagicMock(guardrail_name="g", event_hook="pre_call")
wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=gen(), request_data={})
wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(
callback=cb, response=gen(), hook=_passthrough_hook, request_data={}
)
out = [ch async for ch in wrapped]
snapshot = {
"chunks": out,
@ -197,18 +208,43 @@ async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(pro
async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_raises(proxy_logging):
detail = {"error": "blocked"}
async def boom_gen():
async def boom_hook(*, response: AsyncIterator[object]) -> AsyncGenerator[object, None]:
if False:
yield # pragma: no cover
raise HTTPException(status_code=400, detail=detail)
cb = MagicMock(guardrail_name="presidio", event_hook="post_call")
wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=boom_gen(), request_data={})
request_data: dict[str, object] = {}
wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(
callback=cb, response=_one_chunk(), hook=boom_hook, request_data=request_data
)
with pytest.raises(HTTPException):
async for _ in wrapped:
pass
assert detail["guardrail_name"] == "presidio"
assert detail["guardrail_mode"] == "post_call"
assert request_data["metadata"]["applied_guardrails"] == ["presidio"]
@pytest.mark.asyncio
async def test_wrap_streaming_iterator_leaves_upstream_http_exception_unattributed(proxy_logging):
detail = {"error": "upstream rejected the stream"}
async def failing_upstream() -> AsyncGenerator[object, None]:
if False:
yield # pragma: no cover
raise HTTPException(status_code=502, detail=detail)
cb = MagicMock(guardrail_name="presidio", event_hook="post_call")
request_data: dict[str, object] = {}
wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(
callback=cb, response=failing_upstream(), hook=_passthrough_hook, request_data=request_data
)
with pytest.raises(HTTPException):
async for _ in wrapped:
pass
assert detail == {"error": "upstream rejected the stream"}
assert request_data == {}
# ---------------------------------------------------------------------------
@ -700,33 +736,83 @@ async def test_post_call_response_headers_hook_swallows_callback_error(proxy_log
assert out == {}
class _StreamBlocker(CustomGuardrail):
def __init__(self, guardrail_name: str = "stream-blocker") -> None:
super().__init__(guardrail_name=guardrail_name, event_hook=GuardrailEventHooks.post_call, default_on=True)
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object]
) -> AsyncGenerator[object, None]:
async for _ in response:
raise HTTPException(status_code=400, detail={"error": "blocked"})
yield # pragma: no cover
class _StreamPasser(CustomGuardrail):
def __init__(self, guardrail_name: str = "stream-passer") -> None:
super().__init__(guardrail_name=guardrail_name, event_hook=GuardrailEventHooks.post_call, default_on=True)
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object]
) -> AsyncGenerator[object, None]:
async for chunk in response:
yield chunk
async def _drain_stream_chain(
proxy_logging: ProxyLogging,
user_api_key_dict: UserAPIKeyAuth,
upstream: AsyncIterator[object],
request_data: dict[str, object],
) -> None:
async for _ in proxy_logging.async_post_call_streaming_iterator_hook(
response=upstream,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
):
pass
async def _failing_provider_stream() -> AsyncGenerator[object, None]:
yield "chunk"
raise RuntimeError("provider connection dropped")
@pytest.mark.asyncio
async def test_stream_guardrail_block_names_the_blocking_guardrail_in_applied_guardrails(
proxy_logging, make_user_api_key_auth, monkeypatch
):
class _StreamBlocker(CustomGuardrail):
def __init__(self) -> None:
super().__init__(guardrail_name="stream-blocker", event_hook=GuardrailEventHooks.post_call, default_on=True)
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object]
) -> AsyncGenerator[object, None]:
async for _ in response:
raise HTTPException(status_code=400, detail={"error": "blocked"})
yield # pragma: no cover
monkeypatch.setattr(litellm, "callbacks", [_StreamBlocker()])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
async def upstream():
yield "chunk"
request_data: dict[str, object] = {"metadata": {}}
with pytest.raises(HTTPException):
async for _ in proxy_logging.async_post_call_streaming_iterator_hook(
response=upstream(),
user_api_key_dict=make_user_api_key_auth(),
request_data=request_data,
):
pass
await _drain_stream_chain(proxy_logging, make_user_api_key_auth(), _one_chunk(), request_data)
assert request_data["metadata"]["applied_guardrails"] == ["stream-blocker"]
@pytest.mark.asyncio
async def test_stream_block_by_inner_guardrail_does_not_name_the_outer_layers(
proxy_logging, make_user_api_key_auth, monkeypatch
):
monkeypatch.setattr(litellm, "callbacks", [_StreamBlocker(), _StreamPasser("outer-a"), _StreamPasser("outer-b")])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
request_data: dict[str, object] = {"metadata": {}}
with pytest.raises(HTTPException) as info:
await _drain_stream_chain(proxy_logging, make_user_api_key_auth(), _one_chunk(), request_data)
assert info.value.detail["guardrail_name"] == "stream-blocker"
assert request_data["metadata"]["applied_guardrails"] == ["stream-blocker"]
@pytest.mark.asyncio
async def test_stream_provider_failure_is_not_attributed_to_any_guardrail(
proxy_logging, make_user_api_key_auth, monkeypatch
):
monkeypatch.setattr(litellm, "callbacks", [_StreamPasser("outer-a"), _StreamPasser("outer-b")])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
request_data: dict[str, object] = {"metadata": {}}
with pytest.raises(RuntimeError, match="provider connection dropped"):
await _drain_stream_chain(proxy_logging, make_user_api_key_auth(), _failing_provider_stream(), request_data)
assert request_data["metadata"] == {}