diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 67e2173ecd0..4dbf06e424d 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1158,7 +1158,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/__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..ea48c1c22f2 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py @@ -0,0 +1,112 @@ +import asyncio +import contextvars +from collections.abc import Awaitable, Callable +from typing import Annotated, Final + +from pydantic import Field, TypeAdapter, ValidationError + +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 +_FIRE_AND_FORGET_ADAPTER: Final[TypeAdapter[bool]] = TypeAdapter(bool) +_MAX_INFLIGHT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(Annotated[int, Field(ge=1)]) + + +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 + 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: + 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, 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: + 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 + + 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[object]], *, 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 7bb41b7586b..3a24054eaee 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 @@ -20,9 +20,19 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, 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 ( @@ -170,6 +180,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 +223,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,8 +260,11 @@ class GenericGuardrailAPI(CustomGuardrail): self.fail_on_error: bool = True if fail_on_error is None else fail_on_error + 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). + 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 ) @@ -256,6 +284,22 @@ class GenericGuardrailAPI(CustomGuardrail): super().__init__(**kwargs) + self._dispatcher: Final = dispatcher or BackgroundDispatcher( + guardrail_name=self.guardrail_name, + max_inflight=max_inflight_from_config(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 a stream is checked once when it " + "closes, with whatever reached the client.", + 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: @@ -339,7 +383,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, @@ -377,6 +421,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 +436,26 @@ 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, + *, + guardrail_request: GenericGuardrailAPIRequest, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"], + ) -> bool: + timeout: Final = FIRE_AND_FORGET_POST_TIMEOUT_SECONDS if self.timeout is None else self.timeout + + def _prepare() -> Callable[[], Awaitable[None]]: + payload: Final = guardrail_request.model_dump(mode="json") + headers: Final = self._build_request_headers() + + 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( self, @@ -418,6 +484,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, + 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 @@ -469,15 +560,24 @@ class GenericGuardrailAPI(CustomGuardrail): model=model, ) - headers: Final = self._build_request_headers() + if self.fire_and_forget: + dispatched: Final = self._dispatch_background_post( + guardrail_request=guardrail_request, 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) - # Make the API request + 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=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/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 44e2cc2404f..8a14f580763 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,25 @@ class GenericGuardrailAPIOptionalParams(BaseModel): ), ) + fire_and_forget: bool | None = Field( + default=None, + 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. " + "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." + ), + ) + + 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. Calls beyond it " + "are dropped and recorded as not_run. Defaults to 100." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/conftest.py b/tests/unit/proxy/guardrails/guardrail_hooks/conftest.py new file mode 100644 index 00000000000..0e477c5fa75 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/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/__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..9f660dae6df --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py @@ -0,0 +1,117 @@ +import asyncio +import contextvars +from collections.abc import Awaitable, Callable +from functools import partial +from typing import Final + +import pytest + +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import ( + DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT, + BackgroundDispatcher, + fire_and_forget_from_config, + max_inflight_from_config, +) + +_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)], +) +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 + + +@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 + + +@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 + + +@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 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) + + +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 + + 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() + + +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(lambda: finish, context="call") is True + await dispatcher.wait_for_pending() + + assert (dispatcher.pending_count, dispatcher.dropped_count) == (0, 0) + + +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) + + async def fail() -> None: + raise ConnectionError("refused") + + dispatcher.dispatch(lambda: fail, context="input_type=response litellm_call_id=call-1") + await dispatcher.wait_for_pending() + + 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_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 + + async def record() -> None: + seen.append(_request_scoped.get()) + + token: Final = _request_scoped.set("request-1") + try: + dispatcher.dispatch(lambda: record, context="call") + finally: + _request_scoped.reset(token) + await dispatcher.wait_for_pending() + + 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..150c81fd5d9 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py @@ -0,0 +1,501 @@ +import asyncio +import copy +import json +from collections.abc import AsyncGenerator, 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_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) + 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 _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 + + +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"]) +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_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() + + 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 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..dfb49a48a06 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrail/test_unified_guardrail.py @@ -0,0 +1,117 @@ +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[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), + 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_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() + + sent: Final = [chunk async for chunk in _guarded_stream(observer, _upstream(words=()), route=None)] + + assert (sent, observer.observed, warning_messages("observe-only stream check")) == ([], [], [])