From de065327e9f12bdfa5a1c3e805f3e8de514e70b8 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Sat, 3 Oct 2026 16:38:16 +0300 Subject: [PATCH] fix(guardrails): observe abandoned streams for fire_and_forget guardrails With fire_and_forget the response audit ran only after a stream finished, so a stream the client hung up on, or one that failed upstream, never reached the audit endpoint Guardrails can now set streaming_observe_only. UnifiedLLMGuardrails then forwards every chunk untouched and, when the stream closes for any reason, runs one check over a copy of what reached the client. The check is shielded from cancellation, and a failure in it is logged and never breaks the stream. GenericGuardrailAPI sets the flag from fire_and_forget instead of forcing end-of-stream and block_only, which the new path makes unnecessary. Other guardrails are unchanged --- .../generic_guardrail_api.py | 9 +- .../unified_guardrail/unified_guardrail.py | 81 ++++++++++++++ .../guardrail_hooks/generic_guardrail_api.py | 4 +- .../{generic_guardrail_api => }/conftest.py | 0 .../test_generic_guardrail_api.py | 65 ++++++++--- .../unified_guardrail/__init__.py | 0 .../test_unified_guardrail.py | 104 ++++++++++++++++++ 7 files changed, 240 insertions(+), 23 deletions(-) rename tests/unit/proxy/guardrails/guardrail_hooks/{generic_guardrail_api => }/conftest.py (100%) create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index b50e9adbb0c..4c4c5d0c375 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -264,7 +264,8 @@ class GenericGuardrailAPI(CustomGuardrail): # Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook # via getattr(guardrail_to_apply, "streaming_*", default). - self.streaming_end_of_stream_only: bool = self.fire_and_forget or ( + self.streaming_observe_only: bool = self.fire_and_forget + self.streaming_end_of_stream_only: bool = ( False if streaming_end_of_stream_only is None else streaming_end_of_stream_only ) if streaming_sampling_rate is not None and streaming_sampling_rate < 1: @@ -275,7 +276,7 @@ class GenericGuardrailAPI(CustomGuardrail): # "block_only" (default) drops text rewrites on the streaming path; # "incremental_diff" emits them as synthetic deltas. self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = ( - "block_only" if streaming_transform_mode is None or self.fire_and_forget else streaming_transform_mode + "block_only" if streaming_transform_mode is None else streaming_transform_mode ) # Set supported event hooks @@ -292,8 +293,8 @@ class GenericGuardrailAPI(CustomGuardrail): verbose_proxy_logger.warning( "Generic Guardrail API (%s): fire_and_forget=True makes this guardrail observe-only. " "action=BLOCKED and action=GUARDRAIL_INTERVENED are ignored, fail_on_error=%s and " - "unreachable_fallback=%s cannot block the request, and streaming is forced to " - "end-of-stream observation in block_only mode.", + "unreachable_fallback=%s cannot block the request, and a stream is checked once when it " + "closes, with whatever reached the client.", self.guardrail_name, self.fail_on_error, self.unreachable_fallback, diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 37c1829def4..35243df23ba 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -9,8 +9,10 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint import copy import json from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence +from contextlib import aclosing from typing import TYPE_CHECKING, Any, Final, Protocol +import anyio from fastapi import HTTPException from litellm._logging import verbose_proxy_logger @@ -19,6 +21,7 @@ from litellm.cost_calculator import _infer_call_type from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks @@ -960,6 +963,70 @@ class UnifiedLLMGuardrails(CustomLogger): choices: Final = _chunk_choices(item) return any(getattr(choice, "finish_reason", None) is not None for choice in choices) + async def _stream_then_observe( + self, + *, + guardrail_to_apply: CustomGuardrail, + response: AsyncIterable[object], + request_data: dict[str, object], + user_api_key_dict: UserAPIKeyAuth, + mappings: Mapping[CallTypes, type["BaseTranslation"]], + ) -> AsyncGenerator[object, None]: + streamed: Final[list[object]] = [] # mutable-ok: records chunks as they are forwarded + try: + async for item in response: + streamed.append(item) + yield item + finally: + with anyio.CancelScope(shield=True): + await self._observe_streamed( + streamed=streamed, + guardrail_to_apply=guardrail_to_apply, + request_data=request_data, + user_api_key_dict=user_api_key_dict, + mappings=mappings, + ) + + @staticmethod + async def _observe_streamed( + *, + streamed: list[object], + guardrail_to_apply: CustomGuardrail, + request_data: dict[str, object], + user_api_key_dict: UserAPIKeyAuth, + mappings: Mapping[CallTypes, type["BaseTranslation"]], + ) -> None: + if not streamed: + return + try: + route_call_types: Final = ( + None + if user_api_key_dict.request_route is None + else get_call_types_for_route(user_api_key_dict.request_route) + ) + call_type: Final = ( + route_call_types[0].value + if route_call_types + else _infer_call_type(call_type=None, completion_response=streamed[0]) + ) + handler_cls: Final = None if call_type is None else mappings.get(CallTypes(call_type)) + if handler_cls is None: + return + logging_obj: Final = request_data.get("litellm_logging_obj") + await handler_cls().process_output_streaming_response( + responses_so_far=copy.deepcopy(streamed), + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=logging_obj if isinstance(logging_obj, LiteLLMLoggingObj) else None, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except Exception as e: # noqa: BLE001 # an observe-only guardrail must never break the stream + verbose_proxy_logger.warning( + "UnifiedLLMGuardrails: observe-only stream check for %s failed: %s", + guardrail_to_apply.guardrail_name, + e, + ) + def resolve_streaming_flag(self, guardrail_to_apply: CustomGuardrail | None, name: str, default: object) -> object: """Streaming flag resolution order (later wins): default < guardrail attribute < guardrail_config dict < this callback's optional_params.""" @@ -1051,6 +1118,20 @@ class UnifiedLLMGuardrails(CustomLogger): mappings: Final = load_guardrail_translation_mappings() + if _streaming_flag("streaming_observe_only", False): + async with aclosing( + self._stream_then_observe( + guardrail_to_apply=guardrail_to_apply, + response=response, + request_data=request_data, + user_api_key_dict=user_api_key_dict, + mappings=mappings, + ) + ) as observed: + async for observed_item in observed: + yield observed_item + return + # Streaming text transformation (incremental_diff) diverges enough from the # block_only path that it runs as its own iterator. It requires a route we # can resolve up front to an OpenAI-chat handler (the only supported v1 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 6ba82f51b1f..8a14f580763 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -108,8 +108,8 @@ class GenericGuardrailAPIOptionalParams(BaseModel): description=( "Observe-only mode: the guardrail call runs in the background and the request never waits for it, so " "BLOCKED and GUARDRAIL_INTERVENED answers are ignored. A dispatched call is recorded as success. " - "Streaming sends one end-of-stream call in block_only mode. The background call uses timeout, or " - "30 seconds when unset. Defaults to false." + "A stream sends one call when it closes, with whatever reached the client. The background call uses " + "timeout, or 30 seconds when unset. Defaults to false." ), ) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/conftest.py b/tests/unit/proxy/guardrails/guardrail_hooks/conftest.py similarity index 100% rename from tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/conftest.py rename to tests/unit/proxy/guardrails/guardrail_hooks/conftest.py diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py index 909318966f1..150c81fd5d9 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py @@ -1,7 +1,7 @@ import asyncio import copy import json -from collections.abc import AsyncIterator, Callable, Mapping +from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping from types import SimpleNamespace from typing import Final @@ -346,23 +346,29 @@ def _stream_chunks() -> list[ModelResponseStream]: ] -async def _streamed_texts(guardrail: GenericGuardrailAPI) -> list[str]: - async def stream() -> AsyncIterator[ModelResponseStream]: - for chunk in _stream_chunks(): - yield chunk +async def _upstream(*, fail_after: int | None = None) -> AsyncIterator[ModelResponseStream]: + for i, chunk in enumerate(_stream_chunks()): + if i == fail_after: + raise ConnectionError("upstream reset") + yield chunk - return [ - chunk.choices[0].delta.content or "" - async for chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( - user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/chat/completions"), - response=stream(), - request_data={ - "messages": [{"role": "user", "content": "hi"}], - "guardrail_to_apply": guardrail, - "metadata": {"guardrails": ["ff-guardrail"]}, - }, - ) - ] + +def _guarded_stream( + guardrail: GenericGuardrailAPI, upstream: AsyncIterator[ModelResponseStream] +) -> AsyncGenerator[ModelResponseStream, None]: + return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/chat/completions"), + response=upstream, + request_data={ + "messages": [{"role": "user", "content": "hi"}], + "guardrail_to_apply": guardrail, + "metadata": {"guardrails": ["ff-guardrail"]}, + }, + ) + + +async def _streamed_texts(guardrail: GenericGuardrailAPI) -> list[str]: + return [chunk.choices[0].delta.content or "" async for chunk in _guarded_stream(guardrail, _upstream())] @pytest.mark.parametrize("streaming_transform_mode", ["block_only", "incremental_diff"]) @@ -385,6 +391,31 @@ async def test_a_stream_is_emitted_live_and_sends_one_call_with_the_whole_text( assert [payload["texts"] for payload in endpoint.payloads] == [["Hello world! Bye"]] +async def test_an_abandoned_stream_still_sends_what_reached_the_client() -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, event_hook="post_call") + stream: Final = _guarded_stream(guardrail, _upstream()) + + received: Final = [await anext(stream), await anext(stream)] + await stream.aclose() + await dispatcher.wait_for_pending() + + assert [chunk.choices[0].delta.content for chunk in received] == ["Hello", " "] + assert [payload["texts"] for payload in endpoint.payloads] == [["Hello "]] + + +async def test_a_stream_that_fails_upstream_still_sends_what_reached_the_client() -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, event_hook="post_call") + + with pytest.raises(ConnectionError, match="upstream reset"): + async for _ in _guarded_stream(guardrail, _upstream(fail_after=3)): + pass + await dispatcher.wait_for_pending() + + assert [payload["texts"] for payload in endpoint.payloads] == [["Hello world"]] + + async def test_an_awaited_stream_is_checked_per_sampled_chunk() -> None: endpoint: Final = _Endpoint() diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py new file mode 100644 index 00000000000..f2247b3d032 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py @@ -0,0 +1,104 @@ +import asyncio +from collections.abc import AsyncIterator, Callable +from typing import Final, Literal + +import anyio + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails +from litellm.types.utils import Delta, GenericGuardrailAPIInputs, ModelResponseStream + +_WORDS: Final = ("Hello", " ", "world", "!") + + +class _Observer(CustomGuardrail): + def __init__(self, *, failure: Exception | None = None) -> None: + super().__init__(guardrail_name="observer", event_hook="post_call", default_on=True) + self.streaming_observe_only = True + self.observed: Final[list[list[str]]] = [] # mutable-ok: records each observed call + self._failure: Final = failure + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: object = None, + ) -> GenericGuardrailAPIInputs: + loop: Final = asyncio.get_running_loop() + next_iteration: Final = loop.create_future() + loop.call_soon(next_iteration.set_result, None) + await next_iteration + if self._failure is not None: + raise self._failure + self.observed.append(list(inputs.get("texts", []))) + return inputs + + +def _chunk(word: str) -> ModelResponseStream: + return ModelResponseStream( + model="gpt-4", choices=[litellm.StreamingChoices(index=0, delta=Delta(role="assistant", content=word))] + ) + + +async def _upstream( + *, words: tuple[str, ...] = _WORDS, stall_after: int | None = None +) -> AsyncIterator[ModelResponseStream]: + for i, word in enumerate(words): + if i == stall_after: + await asyncio.Event().wait() + yield _chunk(word) + + +def _guarded_stream( + observer: _Observer, upstream: AsyncIterator[ModelResponseStream], *, route: str | None = "/chat/completions" +) -> AsyncIterator[object]: + return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route=route), + response=upstream, + request_data={ + "messages": [{"role": "user", "content": "hi"}], + "guardrail_to_apply": observer, + "metadata": {"guardrails": ["observer"]}, + }, + ) + + +async def test_a_client_disconnect_mid_stream_is_still_observed() -> None: + observer: Final = _Observer() + two_chunks_sent: Final = asyncio.Event() + sent: Final[list[object]] = [] # mutable-ok: records what reached the client + + async def client() -> None: + async for chunk in _guarded_stream(observer, _upstream(stall_after=2)): + sent.append(chunk) + if len(sent) == 2: + two_chunks_sent.set() + + async with anyio.create_task_group() as requests: + requests.start_soon(client) + await two_chunks_sent.wait() + requests.cancel_scope.cancel() + + assert observer.observed == [["Hello "]] + + +async def test_a_failing_observer_never_breaks_the_stream(warning_messages: Callable[[str], list[str]]) -> None: + observer: Final = _Observer(failure=RuntimeError("observer down")) + + sent: Final = [chunk async for chunk in _guarded_stream(observer, _upstream())] + + assert len(sent) == len(_WORDS) + assert warning_messages("observe-only stream check") == [ + "UnifiedLLMGuardrails: observe-only stream check for observer failed: observer down" + ] + + +async def test_an_empty_stream_is_not_observed(warning_messages: Callable[[str], list[str]]) -> None: + observer: Final = _Observer() + + sent: Final = [chunk async for chunk in _guarded_stream(observer, _upstream(words=()), route=None)] + + assert (sent, observer.observed, warning_messages("observe-only stream check")) == ([], [], [])