fix(guardrails): harden generic_guardrail_api fire_and_forget dispatch

Parse fire_and_forget and fire_and_forget_max_inflight the way pydantic parses config values, so "true" and "5" work. A value that cannot be read is ignored with a warning: fire_and_forget falls back to false, so the guardrail keeps enforcing, and fire_and_forget_max_inflight falls back to 100

Force streaming_transform_mode to block_only under fire_and_forget. With incremental_diff the unified guardrail buffered the stream and dispatched one call per chunk, so a stream now sends a single end-of-stream call while the chunks reach the client live

The background POST now honors the guardrail's timeout and falls back to 30 seconds when it is unset. It is built from the same URL, headers and payload as the awaited call. A payload that cannot be serialized is recorded as not_run before anything is dispatched

Tests that build GenericGuardrailAPI move to a mirror test_generic_guardrail_api.py, the dispatcher and parser tests stay in test_background_dispatch.py, and a shared conftest.py captures proxy warnings
This commit is contained in:
Caduri Katzav 2026-10-03 14:53:48 +03:00
parent 73904c9f94
commit f7e934677b
7 changed files with 619 additions and 468 deletions

View file

@ -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

View file

@ -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})")

View file

@ -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
)

View file

@ -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."
),
)

View file

@ -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)

View file

@ -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]

View file

@ -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