From 73904c9f946d026836b175f5d912f7e4d61f8f72 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Mon, 28 Sep 2026 18:12:33 +0300 Subject: [PATCH 1/5] feat(guardrails): add fire_and_forget dispatch to generic_guardrail_api With fire_and_forget enabled the guardrail POST runs as a detached task and apply_guardrail returns the inputs unchanged right away, so the guardrail becomes observe-only. A dispatched call is recorded as success with a response saying the verdict was not read. At most fire_and_forget_max_inflight calls (default 100) are in flight at once per guardrail and worker. Extra calls are dropped, counted with a rate-limited warning, and recorded as guardrail_status not_run. The background POST has a fixed 30 second timeout and starts from an empty context, so it does not keep request-scoped context vars alive Failures in the background call are logged and never raised. A failure while building the payload before dispatch is logged, passes the request through regardless of fail_on_error, and is recorded as not_run with its own reason. A non-bool fire_and_forget or a non-int fire_and_forget_max_inflight is rejected at startup. Enabling it also forces streaming_end_of_stream_only so a stream sends one call --- .../generic_guardrail_api/__init__.py | 2 + .../background_dispatch.py | 82 ++++ .../generic_guardrail_api.py | 120 ++++- .../guardrail_hooks/generic_guardrail_api.py | 29 ++ .../generic_guardrail_api/__init__.py | 0 .../test_background_dispatch.py | 449 ++++++++++++++++++ 6 files changed, 673 insertions(+), 9 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index de389d8a945..c231eb304a2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -40,6 +40,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"), streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"), timeout=litellm_params.timeout, + fire_and_forget=_get_config_value(litellm_params, optional_params, "fire_and_forget"), + fire_and_forget_max_inflight=_get_config_value(litellm_params, optional_params, "fire_and_forget_max_inflight"), ) litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py new file mode 100644 index 00000000000..853b5230fae --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py @@ -0,0 +1,82 @@ +import asyncio +import contextvars +from collections.abc import Awaitable, Callable +from typing import Final + +from litellm._logging import verbose_proxy_logger + +DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT: Final = 100 + +FIRE_AND_FORGET_POST_TIMEOUT_SECONDS: Final = 30.0 + +FIRE_AND_FORGET_DISPATCHED_REASON: Final = "fire_and_forget dispatched, verdict not read" + +FIRE_AND_FORGET_DROPPED_REASON: Final = "fire_and_forget_max_inflight reached, call dropped" + +FIRE_AND_FORGET_NOT_DISPATCHED_REASON: Final = "fire_and_forget payload could not be built, call not dispatched" + +_DROP_LOG_INTERVAL: Final = 100 + + +def resolve_max_inflight(value: object) -> int: + if value is None: + return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError(f"fire_and_forget_max_inflight must be an int, got {value!r}") + return value + + +class BackgroundDispatcher: + """Runs calls as detached tasks, dropping (and counting) calls once ``max_inflight`` are outstanding.""" + + def __init__(self, *, guardrail_name: str | None, max_inflight: int) -> None: + if max_inflight < 1: + raise ValueError(f"fire_and_forget_max_inflight must be >= 1 (got {max_inflight})") + self._guardrail_name: Final = guardrail_name + self._max_inflight: Final = max_inflight + self._pending: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong refs, asyncio keeps only weak ones + self._dropped: int = 0 + + @property + def pending_count(self) -> int: + return len(self._pending) + + @property + def dropped_count(self) -> int: + return self._dropped + + def dispatch(self, run: Callable[[], Awaitable[None]], *, context: str) -> bool: + if len(self._pending) >= self._max_inflight: + self._dropped += 1 + if self._dropped % _DROP_LOG_INTERVAL == 1: + verbose_proxy_logger.warning( + "Generic Guardrail API (%s, fire_and_forget): dropped %d call(s) so far, " + "%d already in flight (fire_and_forget_max_inflight=%d). %s", + self._guardrail_name, + self._dropped, + len(self._pending), + self._max_inflight, + context, + ) + return False + + task: Final = contextvars.Context().run(asyncio.create_task, self._run_logging_failures(run, context=context)) + self._pending.add(task) + task.add_done_callback(self._pending.discard) + return True + + async def _run_logging_failures(self, run: Callable[[], Awaitable[None]], *, context: str) -> None: + try: + await run() + except Exception as e: # noqa: BLE001 # a detached task has no caller to raise into + verbose_proxy_logger.warning( + "Generic Guardrail API (%s, fire_and_forget) call failed. %s: %s", + self._guardrail_name, + context, + e, + ) + + async def wait_for_pending(self) -> None: + pending: Final = tuple(self._pending) + if pending: + await asyncio.gather(*pending, return_exceptions=True) 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 786b65b1cc3..c5965665a97 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 @@ -11,6 +11,7 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx +from pydantic import JsonValue from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version @@ -20,6 +21,7 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) @@ -33,6 +35,15 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ) from litellm.types.utils import GenericGuardrailAPIInputs +from .background_dispatch import ( + FIRE_AND_FORGET_DISPATCHED_REASON, + FIRE_AND_FORGET_DROPPED_REASON, + FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + BackgroundDispatcher, + resolve_max_inflight, +) + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -170,6 +181,15 @@ def _structured_rows_to_write_back( ) +def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs: + return GenericGuardrailAPIInputs(**inputs) + + +def _call_context(input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"]) -> str: + call_id: Final = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None + return f"input_type={input_type} litellm_call_id={call_id}" + + class GenericGuardrailAPI(CustomGuardrail): """ Generic Guardrail API integration for LiteLLM. @@ -204,9 +224,15 @@ class GenericGuardrailAPI(CustomGuardrail): streaming_end_of_stream_only: bool | None = None, streaming_sampling_rate: int | None = None, streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None, + fire_and_forget: bool | None = None, + fire_and_forget_max_inflight: int | None = None, + async_handler: AsyncHTTPHandler | None = None, + dispatcher: BackgroundDispatcher | None = None, **kwargs, ): - self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + self.async_handler = async_handler or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) self.headers = headers or {} self.extra_headers = extra_headers or [] @@ -235,9 +261,14 @@ class GenericGuardrailAPI(CustomGuardrail): self.fail_on_error: bool = True if fail_on_error is None else fail_on_error + if fire_and_forget is not None and not isinstance(fire_and_forget, bool): # pyright: ignore[reportUnnecessaryIsInstance] # config extras reach here unvalidated + raise ValueError(f"fire_and_forget must be a bool, got {fire_and_forget!r}") + self.fire_and_forget: bool = fire_and_forget is True + # Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook - # via getattr(guardrail_to_apply, "streaming_*", default). - self.streaming_end_of_stream_only: bool = ( + # via getattr(guardrail_to_apply, "streaming_*", default). Forced on under + # fire_and_forget so a stream dispatches one call, not one per sampled chunk. + self.streaming_end_of_stream_only: bool = self.fire_and_forget or ( 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: @@ -256,6 +287,22 @@ class GenericGuardrailAPI(CustomGuardrail): super().__init__(**kwargs) + self._dispatcher: Final = dispatcher or BackgroundDispatcher( + guardrail_name=self.guardrail_name, + max_inflight=resolve_max_inflight(fire_and_forget_max_inflight), + ) + + if self.fire_and_forget: + 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.", + self.guardrail_name, + self.fail_on_error, + self.unreachable_fallback, + ) + verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base) def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata: @@ -377,6 +424,8 @@ class GenericGuardrailAPI(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"], is_unreachable: bool = True, ) -> GenericGuardrailAPIInputs: + if self.fire_and_forget: + raise error unreachable_fail_open: Final = is_unreachable and self.unreachable_fallback == "fail_open" if unreachable_fail_open or not self.fail_on_error: http_status_code: Final = getattr(getattr(error, "response", None), "status_code", None) @@ -390,6 +439,24 @@ class GenericGuardrailAPI(CustomGuardrail): verbose_proxy_logger.error("Generic Guardrail API: failed to make request: %s", str(error)) raise Exception(f"Generic Guardrail API failed: {error}") + def _dispatch_background_post( + self, + *, + payload: Mapping[str, JsonValue], + headers: Mapping[str, str], + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"], + ) -> bool: + async def _post() -> None: + await self.async_handler.post( + url=self.api_base, + json=dict(payload), + headers=dict(headers), + timeout=FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + ) + + return self._dispatcher.dispatch(_post, context=_call_context(input_type, logging_obj)) + @log_guardrail_information async def apply_guardrail( self, @@ -418,6 +485,31 @@ class GenericGuardrailAPI(CustomGuardrail): Raises: Exception: If the guardrail blocks the request """ + if not self.fire_and_forget: + return await self._apply_guardrail(inputs, request_data, input_type, logging_obj) + try: + return await self._apply_guardrail(inputs, request_data, input_type, logging_obj) + except Exception as e: # noqa: BLE001 # an observe-only guardrail must never fail the request + verbose_proxy_logger.warning( + "Generic Guardrail API (%s, fire_and_forget) call not dispatched. %s: %s", + self.guardrail_name, + _call_context(input_type, logging_obj), + e, + ) + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + request_data=request_data or {}, + guardrail_status="not_run", + ) + return _passthrough_inputs(inputs) + + async def _apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"], + ) -> GenericGuardrailAPIInputs: verbose_proxy_logger.debug("Generic Guardrail API: Applying guardrail to text") # Extract texts and images from inputs @@ -470,14 +562,24 @@ class GenericGuardrailAPI(CustomGuardrail): ) headers: Final = self._build_request_headers() - - # Make the API request # Use mode="json" to ensure all iterables are converted to lists + payload: Final = guardrail_request.model_dump(mode="json") + + if self.fire_and_forget: + dispatched: Final = self._dispatch_background_post( + payload=payload, headers=headers, input_type=input_type, logging_obj=logging_obj + ) + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=( + FIRE_AND_FORGET_DISPATCHED_REASON if dispatched else FIRE_AND_FORGET_DROPPED_REASON + ), + request_data=request_data, + guardrail_status="success" if dispatched else "not_run", + ) + return _passthrough_inputs(inputs) + response: Final = await self.async_handler.post( - url=self.api_base, - json=guardrail_request.model_dump(mode="json"), - headers=headers, - timeout=self.timeout, + url=self.api_base, json=payload, headers=headers, timeout=self.timeout ) response.raise_for_status() 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 44e2cc2404f..b7fe8850b93 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -103,6 +103,35 @@ class GenericGuardrailAPIOptionalParams(BaseModel): ), ) + fire_and_forget: bool | None = Field( + default=None, + description=( + "If True, the guardrail HTTP call runs as a background task and the request proceeds " + "without waiting for the response, in every mode (pre_call, during_call, post_call). " + "The guardrail becomes observe-only: action=BLOCKED and action=GUARDRAIL_INTERVENED " + "are ignored, and fail_on_error / unreachable_fallback cannot block the request. The " + "background call has a fixed 30 second timeout. A dispatched call is recorded in the " + "guardrail logs as guardrail_status=success with a response saying the verdict was not " + "read. A call whose payload cannot be built is logged as a warning, passes the request " + "through, and is recorded as guardrail_status=not_run. Also forces " + "streaming_end_of_stream_only=True so a stream sends one call instead of one per " + "sampled chunk. Defaults to False in GenericGuardrailAPI.__init__ when None." + ), + ) + + fire_and_forget_max_inflight: int | None = Field( + default=None, + ge=1, + description=( + "Maximum number of fire_and_forget calls in flight at once for this guardrail, per " + "worker process, so a slow guardrail endpoint cannot pile up background tasks without " + "limit. Calls beyond this limit are dropped and counted, with a rate-limited " + "warning, and recorded in the guardrail logs as guardrail_status=not_run. Must be >= 1, " + "and is validated at startup even when fire_and_forget is off. Only used when " + "fire_and_forget is True. Defaults to 100 in GenericGuardrailAPI.__init__ when None." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py new file mode 100644 index 00000000000..8d84474154f --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py @@ -0,0 +1,449 @@ +import asyncio +import contextvars +import json +import logging +from collections.abc import Iterator +from types import SimpleNamespace + +import httpx +import pydantic +import pytest + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GenericGuardrailAPI, + initialize_guardrail, +) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import ( + DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT, + FIRE_AND_FORGET_DISPATCHED_REASON, + FIRE_AND_FORGET_DROPPED_REASON, + FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + BackgroundDispatcher, +) +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, +) +from litellm.types.guardrails import LitellmParams +from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GenericGuardrailAPIOptionalParams, +) +from litellm.types.utils import Delta, ModelResponseStream + +API_BASE = "https://api.test.guardrail.com" +CLIENT_TIMEOUT_SECONDS = 600.0 + +_request_scoped = contextvars.ContextVar("request_scoped", default=None) + + +class _Endpoint: + """The guardrail server behind a MockTransport. Each request is recorded, then waits on ``gate``.""" + + def __init__(self, *, body=None, status_code=200, error=None, gate_open=True): + self.gate = asyncio.Event() + if gate_open: + self.gate.set() + self.payloads: list[dict] = [] + self.read_timeouts: list[float | None] = [] + self.seen_request_scoped: list[object] = [] + self.completed = 0 + self._body = body or {"action": "NONE"} + self._status_code = status_code + self._error = error + + async def __call__(self, request: httpx.Request) -> httpx.Response: + self.payloads.append(json.loads(request.content)) + self.read_timeouts.append(request.extensions["timeout"]["read"]) + self.seen_request_scoped.append(_request_scoped.get()) + await self.gate.wait() + if self._error is not None: + raise self._error + self.completed += 1 + return httpx.Response(self._status_code, json=self._body) + + def handler(self) -> AsyncHTTPHandler: + return AsyncHTTPHandler(timeout=CLIENT_TIMEOUT_SECONDS, transport=httpx.MockTransport(self)) + + +def _logging_obj(call_id="call-123"): + return SimpleNamespace(litellm_call_id=call_id, litellm_trace_id="trace-123", model_call_details={}) + + +def _guardrail(endpoint, *, name="ff-guardrail", event_hook="pre_call", **options): + return GenericGuardrailAPI( + api_base=API_BASE, + guardrail_name=name, + event_hook=event_hook, + default_on=True, + async_handler=endpoint.handler(), + **options, + ) + + +def _fire_and_forget(endpoint, *, max_inflight=10, **options): + dispatcher = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=max_inflight) + return _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget=True, **options), dispatcher + + +def _request_data(): + return { + "messages": [{"role": "user", "content": "hello"}], + "metadata": {"user_api_key_hash": "hash-1", "user_api_key_team_id": "team-1"}, + } + + +@pytest.fixture +def captured_warnings() -> Iterator[list[logging.LogRecord]]: + records: list[logging.LogRecord] = [] + handler = logging.Handler(level=logging.WARNING) + handler.emit = records.append + previous_level = verbose_proxy_logger.level + verbose_proxy_logger.addHandler(handler) + verbose_proxy_logger.setLevel(logging.WARNING) + yield records + verbose_proxy_logger.removeHandler(handler) + verbose_proxy_logger.setLevel(previous_level) + + +def _messages(records, needle): + return [m for m in (r.getMessage() for r in records) if needle in m] + + +async def test_returns_before_the_post_completes(): + endpoint = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint) + inputs = {"texts": ["hello"], "structured_messages": [{"role": "user", "content": "hello"}]} + + result = await asyncio.wait_for( + guardrail.apply_guardrail( + inputs=inputs, request_data=_request_data(), input_type="request", logging_obj=_logging_obj() + ), + timeout=5, + ) + + assert result == inputs + assert endpoint.completed == 0 + assert dispatcher.pending_count == 1 + + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert endpoint.completed == 1 + assert dispatcher.pending_count == 0 + + +async def test_endpoint_receives_the_same_payload_as_the_awaited_path(): + awaited_endpoint = _Endpoint() + background_endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(background_endpoint) + inputs = {"texts": ["hello"], "images": ["data:image/png;base64,AAAA"], "model": "gpt-4o"} + + for target in (_guardrail(awaited_endpoint), guardrail): + await target.apply_guardrail( + inputs=dict(inputs), request_data=_request_data(), input_type="request", logging_obj=_logging_obj() + ) + await dispatcher.wait_for_pending() + + assert background_endpoint.payloads[0]["texts"] == ["hello"] + assert background_endpoint.payloads[0]["request_data"]["user_api_key_team_id"] == "team-1" + assert background_endpoint.payloads == awaited_endpoint.payloads + + +async def test_background_post_uses_its_own_timeout(): + awaited_endpoint = _Endpoint() + background_endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(background_endpoint) + + for target in (_guardrail(awaited_endpoint), guardrail): + await target.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + + assert background_endpoint.read_timeouts == [FIRE_AND_FORGET_POST_TIMEOUT_SECONDS] + assert awaited_endpoint.read_timeouts == [CLIENT_TIMEOUT_SECONDS] + + +async def test_background_post_does_not_inherit_request_context(): + endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint) + token = _request_scoped.set("request-1") + try: + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + finally: + _request_scoped.reset(token) + await dispatcher.wait_for_pending() + + assert endpoint.completed == 1 + assert endpoint.seen_request_scoped == [None] + + +@pytest.mark.parametrize( + "body", + [ + {"action": "BLOCKED", "blocked_reason": "nope"}, + {"action": "GUARDRAIL_INTERVENED", "texts": ["MASKED"]}, + ], +) +async def test_verdict_is_ignored(body): + endpoint = _Endpoint(body=body) + guardrail, dispatcher = _fire_and_forget(endpoint) + + result = await guardrail.apply_guardrail(inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + + assert result == {"texts": ["my ssn is 123"]} + assert endpoint.completed == 1 + + +@pytest.mark.parametrize( + "endpoint_options", + [ + {"error": httpx.ConnectError("connection refused")}, + {"status_code": 500}, + ], +) +async def test_failing_endpoint_is_logged_not_raised(endpoint_options, captured_warnings): + endpoint = _Endpoint(**endpoint_options) + guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=True, unreachable_fallback="fail_closed") + + result = await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="response", + logging_obj=_logging_obj(call_id="call-failing"), + ) + await dispatcher.wait_for_pending() + + assert result == {"texts": ["hello"]} + failures = _messages(captured_warnings, "call failed") + assert len(failures) == 1 + assert "ff-guardrail" in failures[0] + assert "input_type=response" in failures[0] + assert "litellm_call_id=call-failing" in failures[0] + + +@pytest.mark.parametrize("fail_on_error", [True, False]) +@pytest.mark.parametrize( + ("inputs", "make_request_data"), + [ + ({"texts": ["hi"], "tools": [{"function": {"name": "f"}}]}, _request_data), + ({"texts": ["hi"]}, lambda: {"messages": [], "metadata": None}), + ], + ids=["tool_without_type", "malformed_request_metadata"], +) +async def test_failure_before_dispatch_is_logged_and_passes_through( + inputs, make_request_data, fail_on_error, captured_warnings +): + request_data = make_request_data() + endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=fail_on_error) + + result = await guardrail.apply_guardrail( + inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj("call-bad") + ) + await dispatcher.wait_for_pending() + + assert result == inputs + assert endpoint.payloads == [] + assert _recorded_outcomes(request_data) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] + warnings = _messages(captured_warnings, "not dispatched") + assert len(warnings) == 1 + assert "litellm_call_id=call-bad" in warnings[0] + + +async def test_inflight_cap_drops_and_counts_excess_calls(captured_warnings): + endpoint = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=2) + + results = [ + await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request") + for i in range(5) + ] + + assert results == [{"texts": [f"t{i}"]} for i in range(5)] + assert dispatcher.pending_count == 2 + assert dispatcher.dropped_count == 3 + assert len(_messages(captured_warnings, "dropped")) == 1 + + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert [p["texts"] for p in endpoint.payloads] == [["t0"], ["t1"]] + + +def _recorded_outcomes(request_data): + entries = request_data["metadata"]["standard_logging_guardrail_information"] + return [(entry["guardrail_status"], entry["guardrail_response"]) for entry in entries] + + +async def test_dispatched_and_dropped_calls_are_recorded(): + endpoint = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) + dispatched, dropped = _request_data(), _request_data() + + for request_data in (dispatched, dropped): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert _recorded_outcomes(dispatched) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] + assert _recorded_outcomes(dropped) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + + +async def test_finished_task_frees_its_slot(): + endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) + + for i in range(3): + await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + assert dispatcher.pending_count == 0 + + assert dispatcher.dropped_count == 0 + assert endpoint.completed == 3 + + +def _stream_chunks(): + words = ("Hello", " ", "world", "!", " Bye") + return [ + ModelResponseStream( + model="gpt-4", + choices=[ + litellm.StreamingChoices( + index=0, + delta=Delta(role="assistant", content=word), + finish_reason="stop" if i == len(words) - 1 else None, + ) + ], + ) + for i, word in enumerate(words) + ] + + +async def _run_stream(guardrail): + async def stream(): + for chunk in _stream_chunks(): + yield chunk + + return [ + chunk + 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"]}, + }, + ) + ] + + +async def test_stream_dispatches_one_call(): + per_chunk_endpoint = _Endpoint() + await _run_stream(_guardrail(per_chunk_endpoint, event_hook="post_call", streaming_sampling_rate=1)) + + background_endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget( + background_endpoint, event_hook="post_call", streaming_end_of_stream_only=False, streaming_sampling_rate=1 + ) + streamed = await _run_stream(guardrail) + await dispatcher.wait_for_pending() + + assert len(per_chunk_endpoint.payloads) > 1 + assert len(streamed) == len(_stream_chunks()) + assert len(background_endpoint.payloads) == 1 + assert guardrail.streaming_end_of_stream_only is True + + +def test_observe_only_warning_only_when_enabled(captured_warnings): + _guardrail(_Endpoint(), name="enforcing") + _guardrail(_Endpoint(), name="observer", fire_and_forget=True) + + warnings = _messages(captured_warnings, "observe-only") + assert len(warnings) == 1 + assert "observer" in warnings[0] + + +@pytest.mark.parametrize("value", ["false", "true", 1]) +def test_non_bool_fire_and_forget_is_rejected(value): + with pytest.raises(ValueError, match="fire_and_forget must be a bool"): + _guardrail(_Endpoint(), fire_and_forget=value) + + +@pytest.mark.parametrize("max_inflight", ["5", 2.5, True]) +def test_non_int_max_inflight_is_rejected(max_inflight): + with pytest.raises(ValueError, match="fire_and_forget_max_inflight must be an int"): + _guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) + + +@pytest.mark.parametrize("max_inflight", [0, -1]) +def test_max_inflight_below_one_is_rejected(max_inflight): + with pytest.raises(ValueError, match="fire_and_forget_max_inflight"): + _guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) + with pytest.raises(pydantic.ValidationError): + GenericGuardrailAPIOptionalParams(fire_and_forget_max_inflight=max_inflight) + + +async def test_configured_max_inflight_bounds_dispatch(): + endpoint = _Endpoint(gate_open=False) + guardrail = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=1) + first, second = _request_data(), _request_data() + + for request_data in (first, second): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await guardrail._dispatcher.wait_for_pending() + + assert _recorded_outcomes(first) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] + assert _recorded_outcomes(second) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + assert len(endpoint.payloads) == 1 + + +async def test_default_max_inflight_admits_concurrent_calls(): + endpoint = _Endpoint(gate_open=False) + guardrail = _guardrail(endpoint, fire_and_forget=True) + calls = [_request_data() for _ in range(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + 1)] + + for request_data in calls: + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await guardrail._dispatcher.wait_for_pending() + + outcomes = [_recorded_outcomes(request_data)[0][0] for request_data in calls] + assert outcomes == ["success"] * DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + ["not_run"] + + +async def test_initialize_guardrail_forwards_fire_and_forget(): + litellm_params = LitellmParams( + guardrail="generic_guardrail_api", mode="pre_call", api_base=API_BASE, default_on=True + ) + litellm_params.fire_and_forget = True + litellm_params.fire_and_forget_max_inflight = 1 + gate = asyncio.Event() + + guardrail = initialize_guardrail(litellm_params, {"guardrail_name": "from-config"}) + try: + assert guardrail.fire_and_forget is True + assert guardrail.streaming_end_of_stream_only is True + assert guardrail._dispatcher.dispatch(gate.wait, context="first") is True + assert guardrail._dispatcher.dispatch(gate.wait, context="second") is False + finally: + gate.set() + await guardrail._dispatcher.wait_for_pending() + litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail) + + +async def test_default_awaits_the_endpoint_and_blocks(): + endpoint = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"}) + guardrail = _guardrail(endpoint) + + assert guardrail.fire_and_forget is False + assert guardrail.streaming_end_of_stream_only is False + with pytest.raises(GuardrailRaisedException, match="nope"): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + assert endpoint.completed == 1 From f7e934677bc8e8acf781eb1590ab3dfdaa4ae429 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Sat, 3 Oct 2026 14:53:48 +0300 Subject: [PATCH 2/5] fix(guardrails): harden generic_guardrail_api fire_and_forget dispatch Parse fire_and_forget and fire_and_forget_max_inflight the way pydantic parses config values, so "true" and "5" work. A value that cannot be read is ignored with a warning: fire_and_forget falls back to false, so the guardrail keeps enforcing, and fire_and_forget_max_inflight falls back to 100 Force streaming_transform_mode to block_only under fire_and_forget. With incremental_diff the unified guardrail buffered the stream and dispatched one call per chunk, so a stream now sends a single end-of-stream call while the chunks reach the client live The background POST now honors the guardrail's timeout and falls back to 30 seconds when it is unset. It is built from the same URL, headers and payload as the awaited call. A payload that cannot be serialized is recorded as not_run before anything is dispatched Tests that build GenericGuardrailAPI move to a mirror test_generic_guardrail_api.py, the dispatcher and parser tests stay in test_background_dispatch.py, and a shared conftest.py captures proxy warnings --- litellm/integrations/custom_guardrail.py | 2 +- .../background_dispatch.py | 43 +- .../generic_guardrail_api.py | 59 +-- .../guardrail_hooks/generic_guardrail_api.py | 22 +- .../generic_guardrail_api/conftest.py | 24 + .../test_background_dispatch.py | 483 +++--------------- .../test_generic_guardrail_api.py | 454 ++++++++++++++++ 7 files changed, 619 insertions(+), 468 deletions(-) create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/conftest.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 99bb832e26c..b44120b4cbb 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1153,7 +1153,7 @@ class CustomGuardrail(CustomLogger): return False return self.event_hook == event_type.value - def get_guardrail_dynamic_request_body_params(self, request_data: dict) -> dict: + def get_guardrail_dynamic_request_body_params(self, request_data: dict) -> dict[str, object]: """ Returns `extra_body` to be added to the request body for the Guardrail API call diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py index 853b5230fae..247ff2a9b47 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py @@ -1,7 +1,9 @@ import asyncio import contextvars from collections.abc import Awaitable, Callable -from typing import Final +from typing import Annotated, Final + +from pydantic import Field, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger @@ -16,19 +18,46 @@ FIRE_AND_FORGET_DROPPED_REASON: Final = "fire_and_forget_max_inflight reached, c FIRE_AND_FORGET_NOT_DISPATCHED_REASON: Final = "fire_and_forget payload could not be built, call not dispatched" _DROP_LOG_INTERVAL: Final = 100 +_FIRE_AND_FORGET_ADAPTER: Final[TypeAdapter[bool]] = TypeAdapter(bool) +_MAX_INFLIGHT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(Annotated[int, Field(ge=1)]) -def resolve_max_inflight(value: object) -> int: +def fire_and_forget_from_config(value: object) -> bool: + if value is None: + return False + try: + return _FIRE_AND_FORGET_ADAPTER.validate_python(value) + except ValidationError: + verbose_proxy_logger.warning( + "Ignoring fire_and_forget=%r, expected true or false. Awaiting every guardrail call", value + ) + return False + + +def _parsed_max_inflight(value: object) -> int | None: + if isinstance(value, bool): + return None + try: + return _MAX_INFLIGHT_ADAPTER.validate_python(value) + except ValidationError: + return None + + +def max_inflight_from_config(value: object) -> int: if value is None: return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT - if isinstance(value, bool) or not isinstance(value, int): - raise ValueError(f"fire_and_forget_max_inflight must be an int, got {value!r}") - return value + parsed: Final = _parsed_max_inflight(value) + if parsed is None: + verbose_proxy_logger.warning( + "Ignoring fire_and_forget_max_inflight=%r, expected an integer of at least 1. Using %d", + value, + DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT, + ) + return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + return parsed class BackgroundDispatcher: - """Runs calls as detached tasks, dropping (and counting) calls once ``max_inflight`` are outstanding.""" - def __init__(self, *, guardrail_name: str | None, max_inflight: int) -> None: if max_inflight < 1: raise ValueError(f"fire_and_forget_max_inflight must be >= 1 (got {max_inflight})") 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 c5965665a97..ee1b8a72973 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 @@ -11,7 +11,6 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx -from pydantic import JsonValue from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version @@ -25,6 +24,15 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import ( + FIRE_AND_FORGET_DISPATCHED_REASON, + FIRE_AND_FORGET_DROPPED_REASON, + FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + BackgroundDispatcher, + fire_and_forget_from_config, + max_inflight_from_config, +) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( @@ -35,15 +43,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ) from litellm.types.utils import GenericGuardrailAPIInputs -from .background_dispatch import ( - FIRE_AND_FORGET_DISPATCHED_REASON, - FIRE_AND_FORGET_DROPPED_REASON, - FIRE_AND_FORGET_NOT_DISPATCHED_REASON, - FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, - BackgroundDispatcher, - resolve_max_inflight, -) - if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -261,13 +260,10 @@ class GenericGuardrailAPI(CustomGuardrail): self.fail_on_error: bool = True if fail_on_error is None else fail_on_error - if fire_and_forget is not None and not isinstance(fire_and_forget, bool): # pyright: ignore[reportUnnecessaryIsInstance] # config extras reach here unvalidated - raise ValueError(f"fire_and_forget must be a bool, got {fire_and_forget!r}") - self.fire_and_forget: bool = fire_and_forget is True + self.fire_and_forget: bool = fire_and_forget_from_config(fire_and_forget) # Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook - # via getattr(guardrail_to_apply, "streaming_*", default). Forced on under - # fire_and_forget so a stream dispatches one call, not one per sampled chunk. + # via getattr(guardrail_to_apply, "streaming_*", default). self.streaming_end_of_stream_only: bool = self.fire_and_forget or ( False if streaming_end_of_stream_only is None else streaming_end_of_stream_only ) @@ -279,7 +275,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 else streaming_transform_mode + "block_only" if streaming_transform_mode is None or self.fire_and_forget else streaming_transform_mode ) # Set supported event hooks @@ -289,7 +285,7 @@ class GenericGuardrailAPI(CustomGuardrail): self._dispatcher: Final = dispatcher or BackgroundDispatcher( guardrail_name=self.guardrail_name, - max_inflight=resolve_max_inflight(fire_and_forget_max_inflight), + max_inflight=max_inflight_from_config(fire_and_forget_max_inflight), ) if self.fire_and_forget: @@ -297,7 +293,7 @@ class GenericGuardrailAPI(CustomGuardrail): "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.", + "end-of-stream observation in block_only mode.", self.guardrail_name, self.fail_on_error, self.unreachable_fallback, @@ -386,7 +382,7 @@ class GenericGuardrailAPI(CustomGuardrail): def _build_guardrail_return_inputs( self, *, - texts: list, + texts: list[str], images: list[str] | None, tools: list[ChatCompletionToolParam] | None, structured_messages: Sequence[AllMessageValues] | None, @@ -442,18 +438,16 @@ class GenericGuardrailAPI(CustomGuardrail): def _dispatch_background_post( self, *, - payload: Mapping[str, JsonValue], - headers: Mapping[str, str], + guardrail_request: GenericGuardrailAPIRequest, input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"], ) -> bool: + payload: Final = guardrail_request.model_dump(mode="json") + headers: Final = self._build_request_headers() + timeout: Final = FIRE_AND_FORGET_POST_TIMEOUT_SECONDS if self.timeout is None else self.timeout + async def _post() -> None: - await self.async_handler.post( - url=self.api_base, - json=dict(payload), - headers=dict(headers), - timeout=FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, - ) + await self.async_handler.post(url=self.api_base, json=payload, headers=headers, timeout=timeout) return self._dispatcher.dispatch(_post, context=_call_context(input_type, logging_obj)) @@ -498,7 +492,7 @@ class GenericGuardrailAPI(CustomGuardrail): ) self.add_standard_logging_guardrail_information_to_request_data( guardrail_json_response=FIRE_AND_FORGET_NOT_DISPATCHED_REASON, - request_data=request_data or {}, + request_data=request_data, guardrail_status="not_run", ) return _passthrough_inputs(inputs) @@ -561,13 +555,9 @@ class GenericGuardrailAPI(CustomGuardrail): model=model, ) - headers: Final = self._build_request_headers() - # Use mode="json" to ensure all iterables are converted to lists - payload: Final = guardrail_request.model_dump(mode="json") - if self.fire_and_forget: dispatched: Final = self._dispatch_background_post( - payload=payload, headers=headers, input_type=input_type, logging_obj=logging_obj + guardrail_request=guardrail_request, input_type=input_type, logging_obj=logging_obj ) self.add_standard_logging_guardrail_information_to_request_data( guardrail_json_response=( @@ -578,6 +568,9 @@ class GenericGuardrailAPI(CustomGuardrail): ) return _passthrough_inputs(inputs) + headers: Final = self._build_request_headers() + # Use mode="json" to ensure all iterables are converted to lists + payload: Final = guardrail_request.model_dump(mode="json") response: Final = await self.async_handler.post( url=self.api_base, json=payload, headers=headers, timeout=self.timeout ) 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 b7fe8850b93..6ba82f51b1f 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -106,16 +106,10 @@ class GenericGuardrailAPIOptionalParams(BaseModel): fire_and_forget: bool | None = Field( default=None, description=( - "If True, the guardrail HTTP call runs as a background task and the request proceeds " - "without waiting for the response, in every mode (pre_call, during_call, post_call). " - "The guardrail becomes observe-only: action=BLOCKED and action=GUARDRAIL_INTERVENED " - "are ignored, and fail_on_error / unreachable_fallback cannot block the request. The " - "background call has a fixed 30 second timeout. A dispatched call is recorded in the " - "guardrail logs as guardrail_status=success with a response saying the verdict was not " - "read. A call whose payload cannot be built is logged as a warning, passes the request " - "through, and is recorded as guardrail_status=not_run. Also forces " - "streaming_end_of_stream_only=True so a stream sends one call instead of one per " - "sampled chunk. Defaults to False in GenericGuardrailAPI.__init__ when None." + "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." ), ) @@ -123,12 +117,8 @@ class GenericGuardrailAPIOptionalParams(BaseModel): default=None, ge=1, description=( - "Maximum number of fire_and_forget calls in flight at once for this guardrail, per " - "worker process, so a slow guardrail endpoint cannot pile up background tasks without " - "limit. Calls beyond this limit are dropped and counted, with a rate-limited " - "warning, and recorded in the guardrail logs as guardrail_status=not_run. Must be >= 1, " - "and is validated at startup even when fire_and_forget is off. Only used when " - "fire_and_forget is True. Defaults to 100 in GenericGuardrailAPI.__init__ when None." + "Maximum number of fire_and_forget calls in flight at once for this guardrail, per worker. Calls beyond it " + "are dropped and recorded as not_run. Defaults to 100." ), ) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/conftest.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/conftest.py new file mode 100644 index 00000000000..0e477c5fa75 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/conftest.py @@ -0,0 +1,24 @@ +import logging +from collections.abc import Callable, Iterator +from typing import Final + +import pytest + +from litellm._logging import verbose_proxy_logger + + +@pytest.fixture +def warning_messages() -> Iterator[Callable[[str], list[str]]]: + records: Final[list[logging.LogRecord]] = [] # mutable-ok: the handler appends each record + handler: Final = logging.Handler(level=logging.WARNING) + handler.emit = records.append + previous_level: Final = verbose_proxy_logger.level + verbose_proxy_logger.addHandler(handler) + verbose_proxy_logger.setLevel(logging.WARNING) + + def containing(needle: str) -> list[str]: + return [message for message in (record.getMessage() for record in records) if needle in message] + + yield containing + verbose_proxy_logger.removeHandler(handler) + verbose_proxy_logger.setLevel(previous_level) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py index 8d84474154f..30bb7e353c7 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py @@ -1,449 +1,110 @@ import asyncio import contextvars -import json -import logging -from collections.abc import Iterator -from types import SimpleNamespace +from collections.abc import Callable +from typing import Final -import httpx -import pydantic import pytest -import litellm -from litellm._logging import verbose_proxy_logger -from litellm.exceptions import GuardrailRaisedException -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( - GenericGuardrailAPI, - initialize_guardrail, -) from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import ( DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT, - FIRE_AND_FORGET_DISPATCHED_REASON, - FIRE_AND_FORGET_DROPPED_REASON, - FIRE_AND_FORGET_NOT_DISPATCHED_REASON, - FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, BackgroundDispatcher, + fire_and_forget_from_config, + max_inflight_from_config, ) -from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( - UnifiedLLMGuardrails, + +_request_scoped: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar("request_scoped", default=None) + + +@pytest.mark.parametrize( + ("value", "expected"), + [(None, False), (True, True), (False, False), ("true", True), ("false", False), (1, True), (0, False)], ) -from litellm.types.guardrails import LitellmParams -from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( - GenericGuardrailAPIOptionalParams, -) -from litellm.types.utils import Delta, ModelResponseStream - -API_BASE = "https://api.test.guardrail.com" -CLIENT_TIMEOUT_SECONDS = 600.0 - -_request_scoped = contextvars.ContextVar("request_scoped", default=None) +def test_fire_and_forget_accepts_bools_and_their_config_spellings(value: object, expected: bool) -> None: + assert fire_and_forget_from_config(value) is expected -class _Endpoint: - """The guardrail server behind a MockTransport. Each request is recorded, then waits on ``gate``.""" - - def __init__(self, *, body=None, status_code=200, error=None, gate_open=True): - self.gate = asyncio.Event() - if gate_open: - self.gate.set() - self.payloads: list[dict] = [] - self.read_timeouts: list[float | None] = [] - self.seen_request_scoped: list[object] = [] - self.completed = 0 - self._body = body or {"action": "NONE"} - self._status_code = status_code - self._error = error - - async def __call__(self, request: httpx.Request) -> httpx.Response: - self.payloads.append(json.loads(request.content)) - self.read_timeouts.append(request.extensions["timeout"]["read"]) - self.seen_request_scoped.append(_request_scoped.get()) - await self.gate.wait() - if self._error is not None: - raise self._error - self.completed += 1 - return httpx.Response(self._status_code, json=self._body) - - def handler(self) -> AsyncHTTPHandler: - return AsyncHTTPHandler(timeout=CLIENT_TIMEOUT_SECONDS, transport=httpx.MockTransport(self)) +@pytest.mark.parametrize("value", ["maybe", 2, [True]]) +def test_an_unparseable_fire_and_forget_is_ignored_with_a_warning( + value: object, warning_messages: Callable[[str], list[str]] +) -> None: + assert fire_and_forget_from_config(value) is False + assert len(warning_messages("Ignoring fire_and_forget=")) == 1 -def _logging_obj(call_id="call-123"): - return SimpleNamespace(litellm_call_id=call_id, litellm_trace_id="trace-123", model_call_details={}) +@pytest.mark.parametrize(("value", "expected"), [(None, DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT), (5, 5), ("5", 5)]) +def test_max_inflight_accepts_positive_integers(value: object, expected: int) -> None: + assert max_inflight_from_config(value) == expected -def _guardrail(endpoint, *, name="ff-guardrail", event_hook="pre_call", **options): - return GenericGuardrailAPI( - api_base=API_BASE, - guardrail_name=name, - event_hook=event_hook, - default_on=True, - async_handler=endpoint.handler(), - **options, - ) +@pytest.mark.parametrize("value", [0, -1, 2.5, True, "many"]) +def test_an_invalid_max_inflight_falls_back_to_the_default_with_a_warning( + value: object, warning_messages: Callable[[str], list[str]] +) -> None: + assert max_inflight_from_config(value) == DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + assert len(warning_messages("Ignoring fire_and_forget_max_inflight=")) == 1 -def _fire_and_forget(endpoint, *, max_inflight=10, **options): - dispatcher = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=max_inflight) - return _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget=True, **options), dispatcher +def test_a_dispatcher_needs_room_for_at_least_one_call() -> None: + with pytest.raises(ValueError, match="fire_and_forget_max_inflight"): + BackgroundDispatcher(guardrail_name="g", max_inflight=0) -def _request_data(): - return { - "messages": [{"role": "user", "content": "hello"}], - "metadata": {"user_api_key_hash": "hash-1", "user_api_key_team_id": "team-1"}, - } +async def test_calls_beyond_the_cap_are_dropped_counted_and_warned_once( + warning_messages: Callable[[str], list[str]], +) -> None: + gate: Final = asyncio.Event() + dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=2) + dispatched: Final = [dispatcher.dispatch(gate.wait, context=f"call {i}") for i in range(5)] -@pytest.fixture -def captured_warnings() -> Iterator[list[logging.LogRecord]]: - records: list[logging.LogRecord] = [] - handler = logging.Handler(level=logging.WARNING) - handler.emit = records.append - previous_level = verbose_proxy_logger.level - verbose_proxy_logger.addHandler(handler) - verbose_proxy_logger.setLevel(logging.WARNING) - yield records - verbose_proxy_logger.removeHandler(handler) - verbose_proxy_logger.setLevel(previous_level) - - -def _messages(records, needle): - return [m for m in (r.getMessage() for r in records) if needle in m] - - -async def test_returns_before_the_post_completes(): - endpoint = _Endpoint(gate_open=False) - guardrail, dispatcher = _fire_and_forget(endpoint) - inputs = {"texts": ["hello"], "structured_messages": [{"role": "user", "content": "hello"}]} - - result = await asyncio.wait_for( - guardrail.apply_guardrail( - inputs=inputs, request_data=_request_data(), input_type="request", logging_obj=_logging_obj() - ), - timeout=5, - ) - - assert result == inputs - assert endpoint.completed == 0 - assert dispatcher.pending_count == 1 - - endpoint.gate.set() + assert (dispatched, dispatcher.pending_count, dispatcher.dropped_count) == ([True, True, False, False, False], 2, 3) + assert len(warning_messages("dropped")) == 1 + gate.set() await dispatcher.wait_for_pending() - assert endpoint.completed == 1 - assert dispatcher.pending_count == 0 + +async def test_a_finished_call_frees_its_slot() -> None: + dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=1) + + async def finish() -> None: + return None + + for _ in range(3): + assert dispatcher.dispatch(finish, context="call") is True + await dispatcher.wait_for_pending() + + assert (dispatcher.pending_count, dispatcher.dropped_count) == (0, 0) -async def test_endpoint_receives_the_same_payload_as_the_awaited_path(): - awaited_endpoint = _Endpoint() - background_endpoint = _Endpoint() - guardrail, dispatcher = _fire_and_forget(background_endpoint) - inputs = {"texts": ["hello"], "images": ["data:image/png;base64,AAAA"], "model": "gpt-4o"} +async def test_a_failing_call_is_logged_with_its_context_and_not_raised( + warning_messages: Callable[[str], list[str]], +) -> None: + dispatcher: Final = BackgroundDispatcher(guardrail_name="audit", max_inflight=1) - for target in (_guardrail(awaited_endpoint), guardrail): - await target.apply_guardrail( - inputs=dict(inputs), request_data=_request_data(), input_type="request", logging_obj=_logging_obj() - ) + async def fail() -> None: + raise ConnectionError("refused") + + dispatcher.dispatch(fail, context="input_type=response litellm_call_id=call-1") await dispatcher.wait_for_pending() - assert background_endpoint.payloads[0]["texts"] == ["hello"] - assert background_endpoint.payloads[0]["request_data"]["user_api_key_team_id"] == "team-1" - assert background_endpoint.payloads == awaited_endpoint.payloads + assert warning_messages("call failed") == [ + "Generic Guardrail API (audit, fire_and_forget) call failed. " + "input_type=response litellm_call_id=call-1: refused" + ] -async def test_background_post_uses_its_own_timeout(): - awaited_endpoint = _Endpoint() - background_endpoint = _Endpoint() - guardrail, dispatcher = _fire_and_forget(background_endpoint) +async def test_a_dispatched_call_does_not_see_the_request_context() -> None: + dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=1) + seen: Final[list[str | None]] = [] # mutable-ok: records what the background call saw - for target in (_guardrail(awaited_endpoint), guardrail): - await target.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") - await dispatcher.wait_for_pending() + async def record() -> None: + seen.append(_request_scoped.get()) - assert background_endpoint.read_timeouts == [FIRE_AND_FORGET_POST_TIMEOUT_SECONDS] - assert awaited_endpoint.read_timeouts == [CLIENT_TIMEOUT_SECONDS] - - -async def test_background_post_does_not_inherit_request_context(): - endpoint = _Endpoint() - guardrail, dispatcher = _fire_and_forget(endpoint) - token = _request_scoped.set("request-1") + token: Final = _request_scoped.set("request-1") try: - await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + dispatcher.dispatch(record, context="call") finally: _request_scoped.reset(token) await dispatcher.wait_for_pending() - assert endpoint.completed == 1 - assert endpoint.seen_request_scoped == [None] - - -@pytest.mark.parametrize( - "body", - [ - {"action": "BLOCKED", "blocked_reason": "nope"}, - {"action": "GUARDRAIL_INTERVENED", "texts": ["MASKED"]}, - ], -) -async def test_verdict_is_ignored(body): - endpoint = _Endpoint(body=body) - guardrail, dispatcher = _fire_and_forget(endpoint) - - result = await guardrail.apply_guardrail(inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request") - await dispatcher.wait_for_pending() - - assert result == {"texts": ["my ssn is 123"]} - assert endpoint.completed == 1 - - -@pytest.mark.parametrize( - "endpoint_options", - [ - {"error": httpx.ConnectError("connection refused")}, - {"status_code": 500}, - ], -) -async def test_failing_endpoint_is_logged_not_raised(endpoint_options, captured_warnings): - endpoint = _Endpoint(**endpoint_options) - guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=True, unreachable_fallback="fail_closed") - - result = await guardrail.apply_guardrail( - inputs={"texts": ["hello"]}, - request_data={}, - input_type="response", - logging_obj=_logging_obj(call_id="call-failing"), - ) - await dispatcher.wait_for_pending() - - assert result == {"texts": ["hello"]} - failures = _messages(captured_warnings, "call failed") - assert len(failures) == 1 - assert "ff-guardrail" in failures[0] - assert "input_type=response" in failures[0] - assert "litellm_call_id=call-failing" in failures[0] - - -@pytest.mark.parametrize("fail_on_error", [True, False]) -@pytest.mark.parametrize( - ("inputs", "make_request_data"), - [ - ({"texts": ["hi"], "tools": [{"function": {"name": "f"}}]}, _request_data), - ({"texts": ["hi"]}, lambda: {"messages": [], "metadata": None}), - ], - ids=["tool_without_type", "malformed_request_metadata"], -) -async def test_failure_before_dispatch_is_logged_and_passes_through( - inputs, make_request_data, fail_on_error, captured_warnings -): - request_data = make_request_data() - endpoint = _Endpoint() - guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=fail_on_error) - - result = await guardrail.apply_guardrail( - inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj("call-bad") - ) - await dispatcher.wait_for_pending() - - assert result == inputs - assert endpoint.payloads == [] - assert _recorded_outcomes(request_data) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] - warnings = _messages(captured_warnings, "not dispatched") - assert len(warnings) == 1 - assert "litellm_call_id=call-bad" in warnings[0] - - -async def test_inflight_cap_drops_and_counts_excess_calls(captured_warnings): - endpoint = _Endpoint(gate_open=False) - guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=2) - - results = [ - await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request") - for i in range(5) - ] - - assert results == [{"texts": [f"t{i}"]} for i in range(5)] - assert dispatcher.pending_count == 2 - assert dispatcher.dropped_count == 3 - assert len(_messages(captured_warnings, "dropped")) == 1 - - endpoint.gate.set() - await dispatcher.wait_for_pending() - - assert [p["texts"] for p in endpoint.payloads] == [["t0"], ["t1"]] - - -def _recorded_outcomes(request_data): - entries = request_data["metadata"]["standard_logging_guardrail_information"] - return [(entry["guardrail_status"], entry["guardrail_response"]) for entry in entries] - - -async def test_dispatched_and_dropped_calls_are_recorded(): - endpoint = _Endpoint(gate_open=False) - guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) - dispatched, dropped = _request_data(), _request_data() - - for request_data in (dispatched, dropped): - await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") - endpoint.gate.set() - await dispatcher.wait_for_pending() - - assert _recorded_outcomes(dispatched) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] - assert _recorded_outcomes(dropped) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] - - -async def test_finished_task_frees_its_slot(): - endpoint = _Endpoint() - guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) - - for i in range(3): - await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request") - await dispatcher.wait_for_pending() - assert dispatcher.pending_count == 0 - - assert dispatcher.dropped_count == 0 - assert endpoint.completed == 3 - - -def _stream_chunks(): - words = ("Hello", " ", "world", "!", " Bye") - return [ - ModelResponseStream( - model="gpt-4", - choices=[ - litellm.StreamingChoices( - index=0, - delta=Delta(role="assistant", content=word), - finish_reason="stop" if i == len(words) - 1 else None, - ) - ], - ) - for i, word in enumerate(words) - ] - - -async def _run_stream(guardrail): - async def stream(): - for chunk in _stream_chunks(): - yield chunk - - return [ - chunk - 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"]}, - }, - ) - ] - - -async def test_stream_dispatches_one_call(): - per_chunk_endpoint = _Endpoint() - await _run_stream(_guardrail(per_chunk_endpoint, event_hook="post_call", streaming_sampling_rate=1)) - - background_endpoint = _Endpoint() - guardrail, dispatcher = _fire_and_forget( - background_endpoint, event_hook="post_call", streaming_end_of_stream_only=False, streaming_sampling_rate=1 - ) - streamed = await _run_stream(guardrail) - await dispatcher.wait_for_pending() - - assert len(per_chunk_endpoint.payloads) > 1 - assert len(streamed) == len(_stream_chunks()) - assert len(background_endpoint.payloads) == 1 - assert guardrail.streaming_end_of_stream_only is True - - -def test_observe_only_warning_only_when_enabled(captured_warnings): - _guardrail(_Endpoint(), name="enforcing") - _guardrail(_Endpoint(), name="observer", fire_and_forget=True) - - warnings = _messages(captured_warnings, "observe-only") - assert len(warnings) == 1 - assert "observer" in warnings[0] - - -@pytest.mark.parametrize("value", ["false", "true", 1]) -def test_non_bool_fire_and_forget_is_rejected(value): - with pytest.raises(ValueError, match="fire_and_forget must be a bool"): - _guardrail(_Endpoint(), fire_and_forget=value) - - -@pytest.mark.parametrize("max_inflight", ["5", 2.5, True]) -def test_non_int_max_inflight_is_rejected(max_inflight): - with pytest.raises(ValueError, match="fire_and_forget_max_inflight must be an int"): - _guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) - - -@pytest.mark.parametrize("max_inflight", [0, -1]) -def test_max_inflight_below_one_is_rejected(max_inflight): - with pytest.raises(ValueError, match="fire_and_forget_max_inflight"): - _guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) - with pytest.raises(pydantic.ValidationError): - GenericGuardrailAPIOptionalParams(fire_and_forget_max_inflight=max_inflight) - - -async def test_configured_max_inflight_bounds_dispatch(): - endpoint = _Endpoint(gate_open=False) - guardrail = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=1) - first, second = _request_data(), _request_data() - - for request_data in (first, second): - await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") - endpoint.gate.set() - await guardrail._dispatcher.wait_for_pending() - - assert _recorded_outcomes(first) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] - assert _recorded_outcomes(second) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] - assert len(endpoint.payloads) == 1 - - -async def test_default_max_inflight_admits_concurrent_calls(): - endpoint = _Endpoint(gate_open=False) - guardrail = _guardrail(endpoint, fire_and_forget=True) - calls = [_request_data() for _ in range(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + 1)] - - for request_data in calls: - await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") - endpoint.gate.set() - await guardrail._dispatcher.wait_for_pending() - - outcomes = [_recorded_outcomes(request_data)[0][0] for request_data in calls] - assert outcomes == ["success"] * DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + ["not_run"] - - -async def test_initialize_guardrail_forwards_fire_and_forget(): - litellm_params = LitellmParams( - guardrail="generic_guardrail_api", mode="pre_call", api_base=API_BASE, default_on=True - ) - litellm_params.fire_and_forget = True - litellm_params.fire_and_forget_max_inflight = 1 - gate = asyncio.Event() - - guardrail = initialize_guardrail(litellm_params, {"guardrail_name": "from-config"}) - try: - assert guardrail.fire_and_forget is True - assert guardrail.streaming_end_of_stream_only is True - assert guardrail._dispatcher.dispatch(gate.wait, context="first") is True - assert guardrail._dispatcher.dispatch(gate.wait, context="second") is False - finally: - gate.set() - await guardrail._dispatcher.wait_for_pending() - litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail) - - -async def test_default_awaits_the_endpoint_and_blocks(): - endpoint = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"}) - guardrail = _guardrail(endpoint) - - assert guardrail.fire_and_forget is False - assert guardrail.streaming_end_of_stream_only is False - with pytest.raises(GuardrailRaisedException, match="nope"): - await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") - assert endpoint.completed == 1 + assert seen == [None] 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 new file mode 100644 index 00000000000..fce2c0845a6 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py @@ -0,0 +1,454 @@ +import asyncio +import copy +import json +from collections.abc import AsyncIterator, Callable, Mapping +from types import SimpleNamespace +from typing import Final + +import httpx +import pydantic +import pytest + +import litellm +from litellm.exceptions import GuardrailRaisedException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GenericGuardrailAPI, + initialize_guardrail, +) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import ( + DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT, + FIRE_AND_FORGET_DISPATCHED_REASON, + FIRE_AND_FORGET_DROPPED_REASON, + FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + BackgroundDispatcher, +) +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, +) +from litellm.types.guardrails import LitellmParams +from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GenericGuardrailAPIOptionalParams, +) +from litellm.types.utils import Delta, GenericGuardrailAPIInputs, ModelResponseStream + +API_BASE: Final = "https://api.test.guardrail.com" +CLIENT_TIMEOUT_SECONDS: Final = 600.0 +_STREAM_WORDS: Final = ("Hello", " ", "world", "!", " Bye") + + +class _Endpoint: + def __init__( + self, + *, + body: Mapping[str, object] | None = None, + status_code: int = 200, + error: Exception | None = None, + gate_open: bool = True, + ) -> None: + self.gate: Final = asyncio.Event() + if gate_open: + self.gate.set() + self.requests: Final[list[tuple[str, dict[str, str]]]] = [] # mutable-ok: records each request + self.payloads: Final[list[dict[str, object]]] = [] # mutable-ok: records each request body + self.read_timeouts: Final[list[float | None]] = [] # mutable-ok: records each request timeout + self.completed: int = 0 + self._finished: Final = asyncio.Condition() + self._body: Final = dict(body or {"action": "NONE"}) + self._status_code: Final = status_code + self._error: Final = error + + async def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests.append((str(request.url), dict(request.headers))) + self.payloads.append(json.loads(request.content)) + self.read_timeouts.append(request.extensions["timeout"]["read"]) + await self.gate.wait() + async with self._finished: + self.completed += 1 + self._finished.notify_all() + if self._error is not None: + raise self._error + return httpx.Response(self._status_code, json=self._body) + + async def wait_for_completed(self, count: int) -> None: + async with self._finished: + await asyncio.wait_for(self._finished.wait_for(lambda: self.completed >= count), timeout=5) + + def handler(self) -> AsyncHTTPHandler: + return AsyncHTTPHandler(timeout=CLIENT_TIMEOUT_SECONDS, transport=httpx.MockTransport(self)) + + +def _logging_obj(call_id: str = "call-123") -> SimpleNamespace: + return SimpleNamespace(litellm_call_id=call_id, litellm_trace_id="trace-123", model_call_details={}) + + +def _guardrail( + endpoint: _Endpoint, *, name: str = "ff-guardrail", event_hook: str = "pre_call", **options: object +) -> GenericGuardrailAPI: + return GenericGuardrailAPI( + api_base=API_BASE, + guardrail_name=name, + event_hook=event_hook, + default_on=True, + async_handler=endpoint.handler(), + **options, + ) + + +def _fire_and_forget( + endpoint: _Endpoint, *, max_inflight: int = 10, **options: object +) -> tuple[GenericGuardrailAPI, BackgroundDispatcher]: + dispatcher: Final = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=max_inflight) + return _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget=True, **options), dispatcher + + +def _request_data() -> dict[str, object]: # mutable-ok: apply_guardrail records entries into it + return { + "messages": [{"role": "user", "content": "hello"}], + "metadata": {"user_api_key_hash": "hash-1", "user_api_key_team_id": "team-1"}, + } + + +def _recorded_outcomes(request_data: Mapping[str, object]) -> list[tuple[str, str]]: + metadata: Final = request_data["metadata"] + assert isinstance(metadata, dict) + return [ + (entry["guardrail_status"], entry["guardrail_response"]) + for entry in metadata["standard_logging_guardrail_information"] + ] + + +async def test_returns_before_the_post_completes() -> None: + endpoint: Final = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint) + inputs: Final = GenericGuardrailAPIInputs( + texts=["hello"], structured_messages=[{"role": "user", "content": "hello"}] + ) + + result: Final = await asyncio.wait_for( + guardrail.apply_guardrail( + inputs=inputs, request_data=_request_data(), input_type="request", logging_obj=_logging_obj() + ), + timeout=5, + ) + + assert (result, endpoint.completed, dispatcher.pending_count) == (inputs, 0, 1) + endpoint.gate.set() + await dispatcher.wait_for_pending() + assert (endpoint.completed, dispatcher.pending_count) == (1, 0) + + +async def test_background_post_reaches_the_same_url_with_the_same_headers_and_payload() -> None: + awaited_endpoint: Final = _Endpoint() + background_endpoint: Final = _Endpoint() + auth: Final = {"api_key": "audit-key", "headers": {"x-static": "static-value"}} + guardrail, dispatcher = _fire_and_forget(background_endpoint, **auth) + inputs: Final = GenericGuardrailAPIInputs(texts=["hello"], images=["data:image/png;base64,AAAA"], model="gpt-4o") + + for target in (_guardrail(awaited_endpoint, **auth), guardrail): + await target.apply_guardrail( + inputs=GenericGuardrailAPIInputs(**inputs), + request_data=_request_data(), + input_type="request", + logging_obj=_logging_obj(), + ) + await dispatcher.wait_for_pending() + + assert background_endpoint.requests == awaited_endpoint.requests + assert background_endpoint.payloads == awaited_endpoint.payloads + url, headers = background_endpoint.requests[0] + assert (url, headers["x-api-key"], headers["x-static"]) == ( + f"{API_BASE}/beta/litellm_basic_guardrail_api", + "audit-key", + "static-value", + ) + + +@pytest.mark.parametrize( + ("configured_timeout", "expected_background_timeout"), + [(None, FIRE_AND_FORGET_POST_TIMEOUT_SECONDS), (5.0, 5.0)], +) +async def test_background_post_honors_the_configured_timeout( + configured_timeout: float | None, expected_background_timeout: float +) -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, timeout=configured_timeout) + + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + + assert endpoint.read_timeouts == [expected_background_timeout] + + +@pytest.mark.parametrize( + "body", + [{"action": "BLOCKED", "blocked_reason": "nope"}, {"action": "GUARDRAIL_INTERVENED", "texts": ["MASKED"]}], +) +async def test_the_endpoint_verdict_is_ignored(body: Mapping[str, object]) -> None: + endpoint: Final = _Endpoint(body=body) + guardrail, dispatcher = _fire_and_forget(endpoint) + + result: Final = await guardrail.apply_guardrail( + inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request" + ) + await dispatcher.wait_for_pending() + + assert (result, endpoint.completed) == ({"texts": ["my ssn is 123"]}, 1) + + +@pytest.mark.parametrize( + "endpoint_options", + [{"error": httpx.ConnectError("connection refused")}, {"status_code": 500}], +) +async def test_a_failing_endpoint_is_logged_not_raised( + endpoint_options: Mapping[str, object], warning_messages: Callable[[str], list[str]] +) -> None: + endpoint: Final = _Endpoint(**endpoint_options) + guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=True, unreachable_fallback="fail_closed") + + result: Final = await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="response", + logging_obj=_logging_obj(call_id="call-failing"), + ) + await dispatcher.wait_for_pending() + + assert result == {"texts": ["hello"]} + failures: Final = warning_messages("call failed") + assert len(failures) == 1 + assert "ff-guardrail" in failures[0] + assert "input_type=response litellm_call_id=call-failing" in failures[0] + + +@pytest.mark.parametrize("fail_on_error", [True, False]) +@pytest.mark.parametrize( + ("inputs", "request_data"), + [ + ({"texts": ["hi"], "tools": [{"function": {"name": "f"}}]}, _request_data()), + ({"texts": ["hi"]}, {"messages": [], "metadata": None}), + ], + ids=["tool_without_type", "malformed_request_metadata"], +) +async def test_a_failure_before_dispatch_is_logged_and_passes_through( + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], # mutable-ok: apply_guardrail records entries into it + fail_on_error: bool, + warning_messages: Callable[[str], list[str]], +) -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=fail_on_error) + request: Final = copy.deepcopy(request_data) + + result: Final = await guardrail.apply_guardrail( + inputs=inputs, request_data=request, input_type="request", logging_obj=_logging_obj("call-bad") + ) + await dispatcher.wait_for_pending() + + assert (result, endpoint.payloads) == (inputs, []) + assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] + warnings: Final = warning_messages("not dispatched") + assert len(warnings) == 1 + assert "litellm_call_id=call-bad" in warnings[0] + + +async def test_a_payload_that_cannot_be_serialized_is_recorded_as_not_dispatched() -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, additional_provider_specific_params={"bad": object()}) + request: Final = _request_data() + + result: Final = await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, request_data=request, input_type="request" + ) + await dispatcher.wait_for_pending() + + assert (result, endpoint.payloads) == ({"texts": ["hello"]}, []) + assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] + + +async def test_dispatched_calls_are_recorded_as_success_and_dropped_ones_as_not_run() -> None: + endpoint: Final = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) + dispatched, dropped = _request_data(), _request_data() + + for request_data in (dispatched, dropped): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert _recorded_outcomes(dispatched) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] + assert _recorded_outcomes(dropped) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + + +async def test_a_configured_max_inflight_bounds_dispatch() -> None: + endpoint: Final = _Endpoint(gate_open=False) + guardrail: Final = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=1) + first, second = _request_data(), _request_data() + + for request_data in (first, second): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await endpoint.wait_for_completed(1) + + assert _recorded_outcomes(first) + _recorded_outcomes(second) == [ + ("success", FIRE_AND_FORGET_DISPATCHED_REASON), + ("not_run", FIRE_AND_FORGET_DROPPED_REASON), + ] + assert len(endpoint.payloads) == 1 + + +@pytest.mark.parametrize("max_inflight", [None, 0, "many"]) +async def test_an_unset_or_invalid_max_inflight_admits_the_default_number_of_calls(max_inflight: object) -> None: + endpoint: Final = _Endpoint(gate_open=False) + guardrail: Final = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) + calls: Final = [_request_data() for _ in range(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + 1)] + + for request_data in calls: + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await endpoint.wait_for_completed(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT) + + outcomes: Final = [_recorded_outcomes(request_data)[0][0] for request_data in calls] + assert outcomes == ["success"] * DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + ["not_run"] + + +def _stream_chunks() -> list[ModelResponseStream]: + return [ + ModelResponseStream( + model="gpt-4", + choices=[ + litellm.StreamingChoices( + index=0, + delta=Delta(role="assistant", content=word), + finish_reason="stop" if i == len(_STREAM_WORDS) - 1 else None, + ) + ], + ) + for i, word in enumerate(_STREAM_WORDS) + ] + + +async def _streamed_texts(guardrail: GenericGuardrailAPI) -> list[str]: + async def stream() -> AsyncIterator[ModelResponseStream]: + for chunk in _stream_chunks(): + 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"]}, + }, + ) + ] + + +@pytest.mark.parametrize("streaming_transform_mode", ["block_only", "incremental_diff"]) +async def test_a_stream_is_emitted_live_and_sends_one_call_with_the_whole_text( + streaming_transform_mode: str, +) -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget( + endpoint, + event_hook="post_call", + streaming_end_of_stream_only=False, + streaming_sampling_rate=1, + streaming_transform_mode=streaming_transform_mode, + ) + + streamed: Final = await _streamed_texts(guardrail) + await dispatcher.wait_for_pending() + + assert streamed == list(_STREAM_WORDS), "every chunk must be emitted as it arrives" + assert [payload["texts"] for payload in endpoint.payloads] == [["Hello world! Bye"]] + + +async def test_an_awaited_stream_is_checked_per_sampled_chunk() -> None: + endpoint: Final = _Endpoint() + + await _streamed_texts(_guardrail(endpoint, event_hook="post_call", streaming_sampling_rate=1)) + + assert len(endpoint.payloads) > 1 + + +def test_the_observe_only_warning_is_logged_only_when_enabled(warning_messages: Callable[[str], list[str]]) -> None: + _guardrail(_Endpoint(), name="enforcing") + _guardrail(_Endpoint(), name="observer", fire_and_forget=True) + + warnings: Final = warning_messages("observe-only") + assert len(warnings) == 1 + assert "observer" in warnings[0] + + +@pytest.mark.parametrize("value", ["false", "maybe", 2]) +async def test_a_quoted_false_or_unparseable_fire_and_forget_keeps_the_guardrail_enforcing(value: object) -> None: + endpoint: Final = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"}) + guardrail: Final = _guardrail(endpoint, fire_and_forget=value) + + with pytest.raises(GuardrailRaisedException, match="nope"): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + + +async def test_a_quoted_true_fire_and_forget_turns_on_observe_only() -> None: + endpoint: Final = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"}) + dispatcher: Final = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=1) + quoted: Final = _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget="true") + + result: Final = await quoted.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + + assert (result, endpoint.completed) == ({"texts": ["hello"]}, 1) + + +@pytest.mark.parametrize("max_inflight", [0, -1]) +def test_the_config_form_rejects_a_max_inflight_below_one(max_inflight: int) -> None: + with pytest.raises(pydantic.ValidationError): + GenericGuardrailAPIOptionalParams(fire_and_forget_max_inflight=max_inflight) + + +async def test_initialize_guardrail_forwards_fire_and_forget_and_max_inflight() -> None: + endpoint: Final = _Endpoint(gate_open=False) + guardrail: Final = initialize_guardrail( + LitellmParams( + guardrail="generic_guardrail_api", + mode="pre_call", + api_base=API_BASE, + default_on=True, + fire_and_forget=True, + fire_and_forget_max_inflight=1, + ), + {"guardrail_name": "from-config"}, + ) + guardrail.async_handler = endpoint.handler() + first, second = _request_data(), _request_data() + + try: + for request_data in (first, second): + await asyncio.wait_for( + guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request"), + timeout=5, + ) + endpoint.gate.set() + await endpoint.wait_for_completed(1) + finally: + litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail) + + assert _recorded_outcomes(first) + _recorded_outcomes(second) == [ + ("success", FIRE_AND_FORGET_DISPATCHED_REASON), + ("not_run", FIRE_AND_FORGET_DROPPED_REASON), + ] + + +async def test_by_default_the_endpoint_is_awaited_and_can_block() -> None: + endpoint: Final = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"}) + guardrail: Final = _guardrail(endpoint) + + with pytest.raises(GuardrailRaisedException, match="nope"): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + assert endpoint.completed == 1 From 09922b74cda414a77667cb982b582c002d4e3b81 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Sat, 3 Oct 2026 15:48:31 +0300 Subject: [PATCH 3/5] perf(guardrails): build fire_and_forget payloads only when a slot is free BackgroundDispatcher.dispatch now takes a prepare callable and calls it only after the in-flight cap check passes. A call dropped by fire_and_forget_max_inflight no longer serializes its payload on the request path, which matters most for large or image-heavy requests while the endpoint is slow A payload that cannot be serialized is still recorded as not_run before anything is dispatched --- .../background_dispatch.py | 5 +++-- .../generic_guardrail_api.py | 16 ++++++++++------ .../test_background_dispatch.py | 19 +++++++++++++------ .../test_generic_guardrail_api.py | 16 ++++++++++++++++ 4 files changed, 42 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py index 247ff2a9b47..ea48c1c22f2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py @@ -74,7 +74,7 @@ class BackgroundDispatcher: def dropped_count(self) -> int: return self._dropped - def dispatch(self, run: Callable[[], Awaitable[None]], *, context: str) -> bool: + def dispatch(self, prepare: Callable[[], Callable[[], Awaitable[object]]], *, context: str) -> bool: if len(self._pending) >= self._max_inflight: self._dropped += 1 if self._dropped % _DROP_LOG_INTERVAL == 1: @@ -89,12 +89,13 @@ class BackgroundDispatcher: ) return False + run: Final = prepare() task: Final = contextvars.Context().run(asyncio.create_task, self._run_logging_failures(run, context=context)) self._pending.add(task) task.add_done_callback(self._pending.discard) return True - async def _run_logging_failures(self, run: Callable[[], Awaitable[None]], *, context: str) -> None: + async def _run_logging_failures(self, run: Callable[[], Awaitable[object]], *, context: str) -> None: try: await run() except Exception as e: # noqa: BLE001 # a detached task has no caller to raise into 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 ee1b8a72973..b50e9adbb0c 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 @@ -7,7 +7,7 @@ import fnmatch import os -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx @@ -442,14 +442,18 @@ class GenericGuardrailAPI(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"], ) -> bool: - payload: Final = guardrail_request.model_dump(mode="json") - headers: Final = self._build_request_headers() timeout: Final = FIRE_AND_FORGET_POST_TIMEOUT_SECONDS if self.timeout is None else self.timeout - async def _post() -> None: - await self.async_handler.post(url=self.api_base, json=payload, headers=headers, timeout=timeout) + def _prepare() -> Callable[[], Awaitable[None]]: + payload: Final = guardrail_request.model_dump(mode="json") + headers: Final = self._build_request_headers() - return self._dispatcher.dispatch(_post, context=_call_context(input_type, logging_obj)) + async def _post() -> None: + await self.async_handler.post(url=self.api_base, json=payload, headers=headers, timeout=timeout) + + return _post + + return self._dispatcher.dispatch(_prepare, context=_call_context(input_type, logging_obj)) @log_guardrail_information async def apply_guardrail( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py index 30bb7e353c7..9f660dae6df 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py @@ -1,6 +1,7 @@ import asyncio import contextvars -from collections.abc import Callable +from collections.abc import Awaitable, Callable +from functools import partial from typing import Final import pytest @@ -49,15 +50,21 @@ def test_a_dispatcher_needs_room_for_at_least_one_call() -> None: BackgroundDispatcher(guardrail_name="g", max_inflight=0) -async def test_calls_beyond_the_cap_are_dropped_counted_and_warned_once( +async def test_calls_beyond_the_cap_are_dropped_unprepared_counted_and_warned_once( warning_messages: Callable[[str], list[str]], ) -> None: gate: Final = asyncio.Event() dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=2) + prepared: Final[list[int]] = [] # mutable-ok: records which calls were prepared - dispatched: Final = [dispatcher.dispatch(gate.wait, context=f"call {i}") for i in range(5)] + def prepare(call: int) -> Callable[[], Awaitable[bool]]: + prepared.append(call) + return gate.wait + + dispatched: Final = [dispatcher.dispatch(partial(prepare, i), context=f"call {i}") for i in range(5)] assert (dispatched, dispatcher.pending_count, dispatcher.dropped_count) == ([True, True, False, False, False], 2, 3) + assert prepared == [0, 1] assert len(warning_messages("dropped")) == 1 gate.set() await dispatcher.wait_for_pending() @@ -70,7 +77,7 @@ async def test_a_finished_call_frees_its_slot() -> None: return None for _ in range(3): - assert dispatcher.dispatch(finish, context="call") is True + assert dispatcher.dispatch(lambda: finish, context="call") is True await dispatcher.wait_for_pending() assert (dispatcher.pending_count, dispatcher.dropped_count) == (0, 0) @@ -84,7 +91,7 @@ async def test_a_failing_call_is_logged_with_its_context_and_not_raised( async def fail() -> None: raise ConnectionError("refused") - dispatcher.dispatch(fail, context="input_type=response litellm_call_id=call-1") + dispatcher.dispatch(lambda: fail, context="input_type=response litellm_call_id=call-1") await dispatcher.wait_for_pending() assert warning_messages("call failed") == [ @@ -102,7 +109,7 @@ async def test_a_dispatched_call_does_not_see_the_request_context() -> None: token: Final = _request_scoped.set("request-1") try: - dispatcher.dispatch(record, context="call") + dispatcher.dispatch(lambda: record, context="call") finally: _request_scoped.reset(token) await dispatcher.wait_for_pending() 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 fce2c0845a6..909318966f1 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 @@ -268,6 +268,22 @@ async def test_a_payload_that_cannot_be_serialized_is_recorded_as_not_dispatched assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] +async def test_a_call_dropped_by_the_cap_never_builds_its_payload() -> None: + endpoint: Final = _Endpoint() + guardrail, dispatcher = _fire_and_forget( + endpoint, max_inflight=1, additional_provider_specific_params={"bad": object()} + ) + gate: Final = asyncio.Event() + dispatcher.dispatch(lambda: gate.wait, context="occupies the only slot") + request: Final = _request_data() + + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request, input_type="request") + gate.set() + await dispatcher.wait_for_pending() + + assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + + async def test_dispatched_calls_are_recorded_as_success_and_dropped_ones_as_not_run() -> None: endpoint: Final = _Endpoint(gate_open=False) guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) From de065327e9f12bdfa5a1c3e805f3e8de514e70b8 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Sat, 3 Oct 2026 16:38:16 +0300 Subject: [PATCH 4/5] 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")) == ([], [], []) From de597a5228e815d1885610f99301fc02d8335540 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Sat, 3 Oct 2026 17:04:11 +0300 Subject: [PATCH 5/5] test(guardrails): cover an observe-only stream that no handler understands --- .../unified_guardrail/test_unified_guardrail.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) 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 index f2247b3d032..dfb49a48a06 100644 --- 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 @@ -53,7 +53,7 @@ async def _upstream( def _guarded_stream( - observer: _Observer, upstream: AsyncIterator[ModelResponseStream], *, route: str | None = "/chat/completions" + observer: _Observer, upstream: AsyncIterator[object], *, 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), @@ -96,6 +96,19 @@ async def test_a_failing_observer_never_breaks_the_stream(warning_messages: Call ] +async def test_a_stream_no_handler_understands_is_forwarded_unobserved( + warning_messages: Callable[[str], list[str]], +) -> None: + observer: Final = _Observer() + + async def raw_bytes() -> AsyncIterator[bytes]: + yield b"data: raw\n\n" + + sent: Final = [chunk async for chunk in _guarded_stream(observer, raw_bytes(), route=None)] + + assert (sent, observer.observed, warning_messages("observe-only stream check")) == ([b"data: raw\n\n"], [], []) + + async def test_an_empty_stream_is_not_observed(warning_messages: Callable[[str], list[str]]) -> None: observer: Final = _Observer()