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 e3511d46544..5d4efc78c35 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -39,6 +39,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"), 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"), + fire_and_forget=_get_config_value(litellm_params, optional_params, "fire_and_forget"), + fire_and_forget_max_inflight=_get_config_value(litellm_params, optional_params, "fire_and_forget_max_inflight"), ) litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py new file mode 100644 index 00000000000..853b5230fae --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/background_dispatch.py @@ -0,0 +1,82 @@ +import asyncio +import contextvars +from collections.abc import Awaitable, Callable +from typing import Final + +from litellm._logging import verbose_proxy_logger + +DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT: Final = 100 + +FIRE_AND_FORGET_POST_TIMEOUT_SECONDS: Final = 30.0 + +FIRE_AND_FORGET_DISPATCHED_REASON: Final = "fire_and_forget dispatched, verdict not read" + +FIRE_AND_FORGET_DROPPED_REASON: Final = "fire_and_forget_max_inflight reached, call dropped" + +FIRE_AND_FORGET_NOT_DISPATCHED_REASON: Final = "fire_and_forget payload could not be built, call not dispatched" + +_DROP_LOG_INTERVAL: Final = 100 + + +def resolve_max_inflight(value: object) -> int: + if value is None: + return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError(f"fire_and_forget_max_inflight must be an int, got {value!r}") + return value + + +class BackgroundDispatcher: + """Runs calls as detached tasks, dropping (and counting) calls once ``max_inflight`` are outstanding.""" + + def __init__(self, *, guardrail_name: str | None, max_inflight: int) -> None: + if max_inflight < 1: + raise ValueError(f"fire_and_forget_max_inflight must be >= 1 (got {max_inflight})") + self._guardrail_name: Final = guardrail_name + self._max_inflight: Final = max_inflight + self._pending: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong refs, asyncio keeps only weak ones + self._dropped: int = 0 + + @property + def pending_count(self) -> int: + return len(self._pending) + + @property + def dropped_count(self) -> int: + return self._dropped + + def dispatch(self, run: Callable[[], Awaitable[None]], *, context: str) -> bool: + if len(self._pending) >= self._max_inflight: + self._dropped += 1 + if self._dropped % _DROP_LOG_INTERVAL == 1: + verbose_proxy_logger.warning( + "Generic Guardrail API (%s, fire_and_forget): dropped %d call(s) so far, " + "%d already in flight (fire_and_forget_max_inflight=%d). %s", + self._guardrail_name, + self._dropped, + len(self._pending), + self._max_inflight, + context, + ) + return False + + task: Final = contextvars.Context().run(asyncio.create_task, self._run_logging_failures(run, context=context)) + self._pending.add(task) + task.add_done_callback(self._pending.discard) + return True + + async def _run_logging_failures(self, run: Callable[[], Awaitable[None]], *, context: str) -> None: + try: + await run() + except Exception as e: # noqa: BLE001 # a detached task has no caller to raise into + verbose_proxy_logger.warning( + "Generic Guardrail API (%s, fire_and_forget) call failed. %s: %s", + self._guardrail_name, + context, + e, + ) + + async def wait_for_pending(self) -> None: + pending: Final = tuple(self._pending) + if pending: + await asyncio.gather(*pending, return_exceptions=True) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 3d1a173635e..8b07d39b75d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -11,6 +11,7 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx +from pydantic import JsonValue from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version @@ -20,6 +21,7 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) @@ -33,6 +35,15 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ) from litellm.types.utils import GenericGuardrailAPIInputs +from .background_dispatch import ( + FIRE_AND_FORGET_DISPATCHED_REASON, + FIRE_AND_FORGET_DROPPED_REASON, + FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + BackgroundDispatcher, + resolve_max_inflight, +) + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -170,6 +181,15 @@ def _structured_rows_to_write_back( ) +def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs: + return GenericGuardrailAPIInputs(**inputs) + + +def _call_context(input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"]) -> str: + call_id: Final = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None + return f"input_type={input_type} litellm_call_id={call_id}" + + class GenericGuardrailAPI(CustomGuardrail): """ Generic Guardrail API integration for LiteLLM. @@ -204,9 +224,15 @@ class GenericGuardrailAPI(CustomGuardrail): streaming_end_of_stream_only: bool | None = None, streaming_sampling_rate: int | None = None, streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None, + fire_and_forget: bool | None = None, + fire_and_forget_max_inflight: int | None = None, + async_handler: AsyncHTTPHandler | None = None, + dispatcher: BackgroundDispatcher | None = None, **kwargs, ): - self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + self.async_handler = async_handler or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) self.headers = headers or {} self.extra_headers = extra_headers or [] @@ -235,9 +261,14 @@ class GenericGuardrailAPI(CustomGuardrail): self.fail_on_error: bool = True if fail_on_error is None else fail_on_error + if fire_and_forget is not None and not isinstance(fire_and_forget, bool): # pyright: ignore[reportUnnecessaryIsInstance] # config extras reach here unvalidated + raise ValueError(f"fire_and_forget must be a bool, got {fire_and_forget!r}") + self.fire_and_forget: bool = fire_and_forget is True + # Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook - # via getattr(guardrail_to_apply, "streaming_*", default). - self.streaming_end_of_stream_only: bool = ( + # via getattr(guardrail_to_apply, "streaming_*", default). Forced on under + # fire_and_forget so a stream dispatches one call, not one per sampled chunk. + self.streaming_end_of_stream_only: bool = self.fire_and_forget or ( False if streaming_end_of_stream_only is None else streaming_end_of_stream_only ) if streaming_sampling_rate is not None and streaming_sampling_rate < 1: @@ -256,6 +287,22 @@ class GenericGuardrailAPI(CustomGuardrail): super().__init__(**kwargs) + self._dispatcher: Final = dispatcher or BackgroundDispatcher( + guardrail_name=self.guardrail_name, + max_inflight=resolve_max_inflight(fire_and_forget_max_inflight), + ) + + if self.fire_and_forget: + verbose_proxy_logger.warning( + "Generic Guardrail API (%s): fire_and_forget=True makes this guardrail observe-only. " + "action=BLOCKED and action=GUARDRAIL_INTERVENED are ignored, fail_on_error=%s and " + "unreachable_fallback=%s cannot block the request, and streaming is forced to " + "end-of-stream observation.", + self.guardrail_name, + self.fail_on_error, + self.unreachable_fallback, + ) + verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base) def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata: @@ -377,6 +424,8 @@ class GenericGuardrailAPI(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"], is_unreachable: bool = True, ) -> GenericGuardrailAPIInputs: + if self.fire_and_forget: + raise error unreachable_fail_open: Final = is_unreachable and self.unreachable_fallback == "fail_open" if unreachable_fail_open or not self.fail_on_error: http_status_code: Final = getattr(getattr(error, "response", None), "status_code", None) @@ -390,6 +439,24 @@ class GenericGuardrailAPI(CustomGuardrail): verbose_proxy_logger.error("Generic Guardrail API: failed to make request: %s", str(error)) raise Exception(f"Generic Guardrail API failed: {error}") + def _dispatch_background_post( + self, + *, + payload: Mapping[str, JsonValue], + headers: Mapping[str, str], + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"], + ) -> bool: + async def _post() -> None: + await self.async_handler.post( + url=self.api_base, + json=dict(payload), + headers=dict(headers), + timeout=FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + ) + + return self._dispatcher.dispatch(_post, context=_call_context(input_type, logging_obj)) + @log_guardrail_information async def apply_guardrail( self, @@ -418,6 +485,31 @@ class GenericGuardrailAPI(CustomGuardrail): Raises: Exception: If the guardrail blocks the request """ + if not self.fire_and_forget: + return await self._apply_guardrail(inputs, request_data, input_type, logging_obj) + try: + return await self._apply_guardrail(inputs, request_data, input_type, logging_obj) + except Exception as e: # noqa: BLE001 # an observe-only guardrail must never fail the request + verbose_proxy_logger.warning( + "Generic Guardrail API (%s, fire_and_forget) call not dispatched. %s: %s", + self.guardrail_name, + _call_context(input_type, logging_obj), + e, + ) + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + request_data=request_data or {}, + guardrail_status="not_run", + ) + return _passthrough_inputs(inputs) + + async def _apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"], + ) -> GenericGuardrailAPIInputs: verbose_proxy_logger.debug("Generic Guardrail API: Applying guardrail to text") # Extract texts and images from inputs @@ -470,14 +562,23 @@ class GenericGuardrailAPI(CustomGuardrail): ) headers: Final = self._build_request_headers() - - # Make the API request # Use mode="json" to ensure all iterables are converted to lists - response: Final = await self.async_handler.post( - url=self.api_base, - json=guardrail_request.model_dump(mode="json"), - headers=headers, - ) + payload: Final = guardrail_request.model_dump(mode="json") + + if self.fire_and_forget: + dispatched: Final = self._dispatch_background_post( + payload=payload, headers=headers, input_type=input_type, logging_obj=logging_obj + ) + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=( + FIRE_AND_FORGET_DISPATCHED_REASON if dispatched else FIRE_AND_FORGET_DROPPED_REASON + ), + request_data=request_data, + guardrail_status="success" if dispatched else "not_run", + ) + return _passthrough_inputs(inputs) + + response: Final = await self.async_handler.post(url=self.api_base, json=payload, headers=headers) response.raise_for_status() response_json: Final = response.json() diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 44e2cc2404f..b7fe8850b93 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -103,6 +103,35 @@ class GenericGuardrailAPIOptionalParams(BaseModel): ), ) + fire_and_forget: bool | None = Field( + default=None, + description=( + "If True, the guardrail HTTP call runs as a background task and the request proceeds " + "without waiting for the response, in every mode (pre_call, during_call, post_call). " + "The guardrail becomes observe-only: action=BLOCKED and action=GUARDRAIL_INTERVENED " + "are ignored, and fail_on_error / unreachable_fallback cannot block the request. The " + "background call has a fixed 30 second timeout. A dispatched call is recorded in the " + "guardrail logs as guardrail_status=success with a response saying the verdict was not " + "read. A call whose payload cannot be built is logged as a warning, passes the request " + "through, and is recorded as guardrail_status=not_run. Also forces " + "streaming_end_of_stream_only=True so a stream sends one call instead of one per " + "sampled chunk. Defaults to False in GenericGuardrailAPI.__init__ when None." + ), + ) + + fire_and_forget_max_inflight: int | None = Field( + default=None, + ge=1, + description=( + "Maximum number of fire_and_forget calls in flight at once for this guardrail, per " + "worker process, so a slow guardrail endpoint cannot pile up background tasks without " + "limit. Calls beyond this limit are dropped and counted, with a rate-limited " + "warning, and recorded in the guardrail logs as guardrail_status=not_run. Must be >= 1, " + "and is validated at startup even when fire_and_forget is off. Only used when " + "fire_and_forget is True. Defaults to 100 in GenericGuardrailAPI.__init__ when None." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py new file mode 100644 index 00000000000..8d84474154f --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_background_dispatch.py @@ -0,0 +1,449 @@ +import asyncio +import contextvars +import json +import logging +from collections.abc import Iterator +from types import SimpleNamespace + +import httpx +import pydantic +import pytest + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GenericGuardrailAPI, + initialize_guardrail, +) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import ( + DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT, + FIRE_AND_FORGET_DISPATCHED_REASON, + FIRE_AND_FORGET_DROPPED_REASON, + FIRE_AND_FORGET_NOT_DISPATCHED_REASON, + FIRE_AND_FORGET_POST_TIMEOUT_SECONDS, + BackgroundDispatcher, +) +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, +) +from litellm.types.guardrails import LitellmParams +from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GenericGuardrailAPIOptionalParams, +) +from litellm.types.utils import Delta, ModelResponseStream + +API_BASE = "https://api.test.guardrail.com" +CLIENT_TIMEOUT_SECONDS = 600.0 + +_request_scoped = contextvars.ContextVar("request_scoped", default=None) + + +class _Endpoint: + """The guardrail server behind a MockTransport. Each request is recorded, then waits on ``gate``.""" + + def __init__(self, *, body=None, status_code=200, error=None, gate_open=True): + self.gate = asyncio.Event() + if gate_open: + self.gate.set() + self.payloads: list[dict] = [] + self.read_timeouts: list[float | None] = [] + self.seen_request_scoped: list[object] = [] + self.completed = 0 + self._body = body or {"action": "NONE"} + self._status_code = status_code + self._error = error + + async def __call__(self, request: httpx.Request) -> httpx.Response: + self.payloads.append(json.loads(request.content)) + self.read_timeouts.append(request.extensions["timeout"]["read"]) + self.seen_request_scoped.append(_request_scoped.get()) + await self.gate.wait() + if self._error is not None: + raise self._error + self.completed += 1 + return httpx.Response(self._status_code, json=self._body) + + def handler(self) -> AsyncHTTPHandler: + return AsyncHTTPHandler(timeout=CLIENT_TIMEOUT_SECONDS, transport=httpx.MockTransport(self)) + + +def _logging_obj(call_id="call-123"): + return SimpleNamespace(litellm_call_id=call_id, litellm_trace_id="trace-123", model_call_details={}) + + +def _guardrail(endpoint, *, name="ff-guardrail", event_hook="pre_call", **options): + return GenericGuardrailAPI( + api_base=API_BASE, + guardrail_name=name, + event_hook=event_hook, + default_on=True, + async_handler=endpoint.handler(), + **options, + ) + + +def _fire_and_forget(endpoint, *, max_inflight=10, **options): + dispatcher = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=max_inflight) + return _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget=True, **options), dispatcher + + +def _request_data(): + return { + "messages": [{"role": "user", "content": "hello"}], + "metadata": {"user_api_key_hash": "hash-1", "user_api_key_team_id": "team-1"}, + } + + +@pytest.fixture +def captured_warnings() -> Iterator[list[logging.LogRecord]]: + records: list[logging.LogRecord] = [] + handler = logging.Handler(level=logging.WARNING) + handler.emit = records.append + previous_level = verbose_proxy_logger.level + verbose_proxy_logger.addHandler(handler) + verbose_proxy_logger.setLevel(logging.WARNING) + yield records + verbose_proxy_logger.removeHandler(handler) + verbose_proxy_logger.setLevel(previous_level) + + +def _messages(records, needle): + return [m for m in (r.getMessage() for r in records) if needle in m] + + +async def test_returns_before_the_post_completes(): + endpoint = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint) + inputs = {"texts": ["hello"], "structured_messages": [{"role": "user", "content": "hello"}]} + + result = await asyncio.wait_for( + guardrail.apply_guardrail( + inputs=inputs, request_data=_request_data(), input_type="request", logging_obj=_logging_obj() + ), + timeout=5, + ) + + assert result == inputs + assert endpoint.completed == 0 + assert dispatcher.pending_count == 1 + + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert endpoint.completed == 1 + assert dispatcher.pending_count == 0 + + +async def test_endpoint_receives_the_same_payload_as_the_awaited_path(): + awaited_endpoint = _Endpoint() + background_endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(background_endpoint) + inputs = {"texts": ["hello"], "images": ["data:image/png;base64,AAAA"], "model": "gpt-4o"} + + for target in (_guardrail(awaited_endpoint), guardrail): + await target.apply_guardrail( + inputs=dict(inputs), request_data=_request_data(), input_type="request", logging_obj=_logging_obj() + ) + await dispatcher.wait_for_pending() + + assert background_endpoint.payloads[0]["texts"] == ["hello"] + assert background_endpoint.payloads[0]["request_data"]["user_api_key_team_id"] == "team-1" + assert background_endpoint.payloads == awaited_endpoint.payloads + + +async def test_background_post_uses_its_own_timeout(): + awaited_endpoint = _Endpoint() + background_endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(background_endpoint) + + for target in (_guardrail(awaited_endpoint), guardrail): + await target.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + + assert background_endpoint.read_timeouts == [FIRE_AND_FORGET_POST_TIMEOUT_SECONDS] + assert awaited_endpoint.read_timeouts == [CLIENT_TIMEOUT_SECONDS] + + +async def test_background_post_does_not_inherit_request_context(): + endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint) + token = _request_scoped.set("request-1") + try: + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + finally: + _request_scoped.reset(token) + await dispatcher.wait_for_pending() + + assert endpoint.completed == 1 + assert endpoint.seen_request_scoped == [None] + + +@pytest.mark.parametrize( + "body", + [ + {"action": "BLOCKED", "blocked_reason": "nope"}, + {"action": "GUARDRAIL_INTERVENED", "texts": ["MASKED"]}, + ], +) +async def test_verdict_is_ignored(body): + endpoint = _Endpoint(body=body) + guardrail, dispatcher = _fire_and_forget(endpoint) + + result = await guardrail.apply_guardrail(inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + + assert result == {"texts": ["my ssn is 123"]} + assert endpoint.completed == 1 + + +@pytest.mark.parametrize( + "endpoint_options", + [ + {"error": httpx.ConnectError("connection refused")}, + {"status_code": 500}, + ], +) +async def test_failing_endpoint_is_logged_not_raised(endpoint_options, captured_warnings): + endpoint = _Endpoint(**endpoint_options) + guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=True, unreachable_fallback="fail_closed") + + result = await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="response", + logging_obj=_logging_obj(call_id="call-failing"), + ) + await dispatcher.wait_for_pending() + + assert result == {"texts": ["hello"]} + failures = _messages(captured_warnings, "call failed") + assert len(failures) == 1 + assert "ff-guardrail" in failures[0] + assert "input_type=response" in failures[0] + assert "litellm_call_id=call-failing" in failures[0] + + +@pytest.mark.parametrize("fail_on_error", [True, False]) +@pytest.mark.parametrize( + ("inputs", "make_request_data"), + [ + ({"texts": ["hi"], "tools": [{"function": {"name": "f"}}]}, _request_data), + ({"texts": ["hi"]}, lambda: {"messages": [], "metadata": None}), + ], + ids=["tool_without_type", "malformed_request_metadata"], +) +async def test_failure_before_dispatch_is_logged_and_passes_through( + inputs, make_request_data, fail_on_error, captured_warnings +): + request_data = make_request_data() + endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=fail_on_error) + + result = await guardrail.apply_guardrail( + inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj("call-bad") + ) + await dispatcher.wait_for_pending() + + assert result == inputs + assert endpoint.payloads == [] + assert _recorded_outcomes(request_data) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)] + warnings = _messages(captured_warnings, "not dispatched") + assert len(warnings) == 1 + assert "litellm_call_id=call-bad" in warnings[0] + + +async def test_inflight_cap_drops_and_counts_excess_calls(captured_warnings): + endpoint = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=2) + + results = [ + await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request") + for i in range(5) + ] + + assert results == [{"texts": [f"t{i}"]} for i in range(5)] + assert dispatcher.pending_count == 2 + assert dispatcher.dropped_count == 3 + assert len(_messages(captured_warnings, "dropped")) == 1 + + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert [p["texts"] for p in endpoint.payloads] == [["t0"], ["t1"]] + + +def _recorded_outcomes(request_data): + entries = request_data["metadata"]["standard_logging_guardrail_information"] + return [(entry["guardrail_status"], entry["guardrail_response"]) for entry in entries] + + +async def test_dispatched_and_dropped_calls_are_recorded(): + endpoint = _Endpoint(gate_open=False) + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) + dispatched, dropped = _request_data(), _request_data() + + for request_data in (dispatched, dropped): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await dispatcher.wait_for_pending() + + assert _recorded_outcomes(dispatched) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] + assert _recorded_outcomes(dropped) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + + +async def test_finished_task_frees_its_slot(): + endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1) + + for i in range(3): + await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request") + await dispatcher.wait_for_pending() + assert dispatcher.pending_count == 0 + + assert dispatcher.dropped_count == 0 + assert endpoint.completed == 3 + + +def _stream_chunks(): + words = ("Hello", " ", "world", "!", " Bye") + return [ + ModelResponseStream( + model="gpt-4", + choices=[ + litellm.StreamingChoices( + index=0, + delta=Delta(role="assistant", content=word), + finish_reason="stop" if i == len(words) - 1 else None, + ) + ], + ) + for i, word in enumerate(words) + ] + + +async def _run_stream(guardrail): + async def stream(): + for chunk in _stream_chunks(): + yield chunk + + return [ + chunk + async for chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/chat/completions"), + response=stream(), + request_data={ + "messages": [{"role": "user", "content": "hi"}], + "guardrail_to_apply": guardrail, + "metadata": {"guardrails": ["ff-guardrail"]}, + }, + ) + ] + + +async def test_stream_dispatches_one_call(): + per_chunk_endpoint = _Endpoint() + await _run_stream(_guardrail(per_chunk_endpoint, event_hook="post_call", streaming_sampling_rate=1)) + + background_endpoint = _Endpoint() + guardrail, dispatcher = _fire_and_forget( + background_endpoint, event_hook="post_call", streaming_end_of_stream_only=False, streaming_sampling_rate=1 + ) + streamed = await _run_stream(guardrail) + await dispatcher.wait_for_pending() + + assert len(per_chunk_endpoint.payloads) > 1 + assert len(streamed) == len(_stream_chunks()) + assert len(background_endpoint.payloads) == 1 + assert guardrail.streaming_end_of_stream_only is True + + +def test_observe_only_warning_only_when_enabled(captured_warnings): + _guardrail(_Endpoint(), name="enforcing") + _guardrail(_Endpoint(), name="observer", fire_and_forget=True) + + warnings = _messages(captured_warnings, "observe-only") + assert len(warnings) == 1 + assert "observer" in warnings[0] + + +@pytest.mark.parametrize("value", ["false", "true", 1]) +def test_non_bool_fire_and_forget_is_rejected(value): + with pytest.raises(ValueError, match="fire_and_forget must be a bool"): + _guardrail(_Endpoint(), fire_and_forget=value) + + +@pytest.mark.parametrize("max_inflight", ["5", 2.5, True]) +def test_non_int_max_inflight_is_rejected(max_inflight): + with pytest.raises(ValueError, match="fire_and_forget_max_inflight must be an int"): + _guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) + + +@pytest.mark.parametrize("max_inflight", [0, -1]) +def test_max_inflight_below_one_is_rejected(max_inflight): + with pytest.raises(ValueError, match="fire_and_forget_max_inflight"): + _guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight) + with pytest.raises(pydantic.ValidationError): + GenericGuardrailAPIOptionalParams(fire_and_forget_max_inflight=max_inflight) + + +async def test_configured_max_inflight_bounds_dispatch(): + endpoint = _Endpoint(gate_open=False) + guardrail = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=1) + first, second = _request_data(), _request_data() + + for request_data in (first, second): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await guardrail._dispatcher.wait_for_pending() + + assert _recorded_outcomes(first) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)] + assert _recorded_outcomes(second) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)] + assert len(endpoint.payloads) == 1 + + +async def test_default_max_inflight_admits_concurrent_calls(): + endpoint = _Endpoint(gate_open=False) + guardrail = _guardrail(endpoint, fire_and_forget=True) + calls = [_request_data() for _ in range(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + 1)] + + for request_data in calls: + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request") + endpoint.gate.set() + await guardrail._dispatcher.wait_for_pending() + + outcomes = [_recorded_outcomes(request_data)[0][0] for request_data in calls] + assert outcomes == ["success"] * DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + ["not_run"] + + +async def test_initialize_guardrail_forwards_fire_and_forget(): + litellm_params = LitellmParams( + guardrail="generic_guardrail_api", mode="pre_call", api_base=API_BASE, default_on=True + ) + litellm_params.fire_and_forget = True + litellm_params.fire_and_forget_max_inflight = 1 + gate = asyncio.Event() + + guardrail = initialize_guardrail(litellm_params, {"guardrail_name": "from-config"}) + try: + assert guardrail.fire_and_forget is True + assert guardrail.streaming_end_of_stream_only is True + assert guardrail._dispatcher.dispatch(gate.wait, context="first") is True + assert guardrail._dispatcher.dispatch(gate.wait, context="second") is False + finally: + gate.set() + await guardrail._dispatcher.wait_for_pending() + litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail) + + +async def test_default_awaits_the_endpoint_and_blocks(): + endpoint = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"}) + guardrail = _guardrail(endpoint) + + assert guardrail.fire_and_forget is False + assert guardrail.streaming_end_of_stream_only is False + with pytest.raises(GuardrailRaisedException, match="nope"): + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request") + assert endpoint.completed == 1