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
This commit is contained in:
Caduri Katzav 2026-10-03 16:38:16 +03:00
parent 09922b74cd
commit de065327e9
7 changed files with 240 additions and 23 deletions

View file

@ -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,

View file

@ -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

View file

@ -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."
),
)

View file

@ -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()

View file

@ -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")) == ([], [], [])