mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
09922b74cd
commit
de065327e9
7 changed files with 240 additions and 23 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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")) == ([], [], [])
|
||||
Loading…
Add table
Reference in a new issue