mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
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:
parent
9fa85c5da5
commit
73dea4c567
3 changed files with 223 additions and 64 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue