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