This commit is contained in:
Caduri 2026-10-05 00:58:04 -04:00 • committed by GitHub
commit 0412812987
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1083 additions and 10 deletions

View file

@ -1158,7 +1158,7 @@ class CustomGuardrail(CustomLogger):
return False
return self.event_hook == event_type.value
def get_guardrail_dynamic_request_body_params(self, request_data: dict) -> dict:
def get_guardrail_dynamic_request_body_params(self, request_data: dict) -> dict[str, object]:
"""
Returns `extra_body` to be added to the request body for the Guardrail API call

View file

@ -40,6 +40,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
timeout=litellm_params.timeout,
fire_and_forget=_get_config_value(litellm_params, optional_params, "fire_and_forget"),
fire_and_forget_max_inflight=_get_config_value(litellm_params, optional_params, "fire_and_forget_max_inflight"),
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -0,0 +1,112 @@
import asyncio
import contextvars
from collections.abc import Awaitable, Callable
from typing import Annotated, Final
from pydantic import Field, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT: Final = 100
FIRE_AND_FORGET_POST_TIMEOUT_SECONDS: Final = 30.0
FIRE_AND_FORGET_DISPATCHED_REASON: Final = "fire_and_forget dispatched, verdict not read"
FIRE_AND_FORGET_DROPPED_REASON: Final = "fire_and_forget_max_inflight reached, call dropped"
FIRE_AND_FORGET_NOT_DISPATCHED_REASON: Final = "fire_and_forget payload could not be built, call not dispatched"
_DROP_LOG_INTERVAL: Final = 100
_FIRE_AND_FORGET_ADAPTER: Final[TypeAdapter[bool]] = TypeAdapter(bool)
_MAX_INFLIGHT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(Annotated[int, Field(ge=1)])
def fire_and_forget_from_config(value: object) -> bool:
if value is None:
return False
try:
return _FIRE_AND_FORGET_ADAPTER.validate_python(value)
except ValidationError:
verbose_proxy_logger.warning(
"Ignoring fire_and_forget=%r, expected true or false. Awaiting every guardrail call", value
)
return False
def _parsed_max_inflight(value: object) -> int | None:
if isinstance(value, bool):
return None
try:
return _MAX_INFLIGHT_ADAPTER.validate_python(value)
except ValidationError:
return None
def max_inflight_from_config(value: object) -> int:
if value is None:
return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT
parsed: Final = _parsed_max_inflight(value)
if parsed is None:
verbose_proxy_logger.warning(
"Ignoring fire_and_forget_max_inflight=%r, expected an integer of at least 1. Using %d",
value,
DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT,
)
return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT
return parsed
class BackgroundDispatcher:
def __init__(self, *, guardrail_name: str | None, max_inflight: int) -> None:
if max_inflight < 1:
raise ValueError(f"fire_and_forget_max_inflight must be >= 1 (got {max_inflight})")
self._guardrail_name: Final = guardrail_name
self._max_inflight: Final = max_inflight
self._pending: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong refs, asyncio keeps only weak ones
self._dropped: int = 0
@property
def pending_count(self) -> int:
return len(self._pending)
@property
def dropped_count(self) -> int:
return self._dropped
def dispatch(self, prepare: Callable[[], Callable[[], Awaitable[object]]], *, context: str) -> bool:
if len(self._pending) >= self._max_inflight:
self._dropped += 1
if self._dropped % _DROP_LOG_INTERVAL == 1:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s, fire_and_forget): dropped %d call(s) so far, "
"%d already in flight (fire_and_forget_max_inflight=%d). %s",
self._guardrail_name,
self._dropped,
len(self._pending),
self._max_inflight,
context,
)
return False
run: Final = prepare()
task: Final = contextvars.Context().run(asyncio.create_task, self._run_logging_failures(run, context=context))
self._pending.add(task)
task.add_done_callback(self._pending.discard)
return True
async def _run_logging_failures(self, run: Callable[[], Awaitable[object]], *, context: str) -> None:
try:
await run()
except Exception as e: # noqa: BLE001 # a detached task has no caller to raise into
verbose_proxy_logger.warning(
"Generic Guardrail API (%s, fire_and_forget) call failed. %s: %s",
self._guardrail_name,
context,
e,
)
async def wait_for_pending(self) -> None:
pending: Final = tuple(self._pending)
if pending:
await asyncio.gather(*pending, return_exceptions=True)

View file

@ -7,7 +7,7 @@
import fnmatch
import os
from collections.abc import Mapping, Sequence
from collections.abc import Awaitable, Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
import httpx
@ -20,9 +20,19 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import (
FIRE_AND_FORGET_DISPATCHED_REASON,
FIRE_AND_FORGET_DROPPED_REASON,
FIRE_AND_FORGET_NOT_DISPATCHED_REASON,
FIRE_AND_FORGET_POST_TIMEOUT_SECONDS,
BackgroundDispatcher,
fire_and_forget_from_config,
max_inflight_from_config,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
@ -170,6 +180,15 @@ def _structured_rows_to_write_back(
)
def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(**inputs)
def _call_context(input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"]) -> str:
call_id: Final = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None
return f"input_type={input_type} litellm_call_id={call_id}"
class GenericGuardrailAPI(CustomGuardrail):
"""
Generic Guardrail API integration for LiteLLM.
@ -204,9 +223,15 @@ class GenericGuardrailAPI(CustomGuardrail):
streaming_end_of_stream_only: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
fire_and_forget: bool | None = None,
fire_and_forget_max_inflight: int | None = None,
async_handler: AsyncHTTPHandler | None = None,
dispatcher: BackgroundDispatcher | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.headers = headers or {}
self.extra_headers = extra_headers or []
@ -235,8 +260,11 @@ class GenericGuardrailAPI(CustomGuardrail):
self.fail_on_error: bool = True if fail_on_error is None else fail_on_error
self.fire_and_forget: bool = fire_and_forget_from_config(fire_and_forget)
# Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook
# via getattr(guardrail_to_apply, "streaming_*", default).
self.streaming_observe_only: bool = self.fire_and_forget
self.streaming_end_of_stream_only: bool = (
False if streaming_end_of_stream_only is None else streaming_end_of_stream_only
)
@ -256,6 +284,22 @@ class GenericGuardrailAPI(CustomGuardrail):
super().__init__(**kwargs)
self._dispatcher: Final = dispatcher or BackgroundDispatcher(
guardrail_name=self.guardrail_name,
max_inflight=max_inflight_from_config(fire_and_forget_max_inflight),
)
if self.fire_and_forget:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): fire_and_forget=True makes this guardrail observe-only. "
"action=BLOCKED and action=GUARDRAIL_INTERVENED are ignored, fail_on_error=%s and "
"unreachable_fallback=%s cannot block the request, and a stream is checked once when it "
"closes, with whatever reached the client.",
self.guardrail_name,
self.fail_on_error,
self.unreachable_fallback,
)
verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base)
def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata:
@ -339,7 +383,7 @@ class GenericGuardrailAPI(CustomGuardrail):
def _build_guardrail_return_inputs(
self,
*,
texts: list,
texts: list[str],
images: list[str] | None,
tools: list[ChatCompletionToolParam] | None,
structured_messages: Sequence[AllMessageValues] | None,
@ -377,6 +421,8 @@ class GenericGuardrailAPI(CustomGuardrail):
logging_obj: Optional["LiteLLMLoggingObj"],
is_unreachable: bool = True,
) -> GenericGuardrailAPIInputs:
if self.fire_and_forget:
raise error
unreachable_fail_open: Final = is_unreachable and self.unreachable_fallback == "fail_open"
if unreachable_fail_open or not self.fail_on_error:
http_status_code: Final = getattr(getattr(error, "response", None), "status_code", None)
@ -390,6 +436,26 @@ class GenericGuardrailAPI(CustomGuardrail):
verbose_proxy_logger.error("Generic Guardrail API: failed to make request: %s", str(error))
raise Exception(f"Generic Guardrail API failed: {error}")
def _dispatch_background_post(
self,
*,
guardrail_request: GenericGuardrailAPIRequest,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"],
) -> bool:
timeout: Final = FIRE_AND_FORGET_POST_TIMEOUT_SECONDS if self.timeout is None else self.timeout
def _prepare() -> Callable[[], Awaitable[None]]:
payload: Final = guardrail_request.model_dump(mode="json")
headers: Final = self._build_request_headers()
async def _post() -> None:
await self.async_handler.post(url=self.api_base, json=payload, headers=headers, timeout=timeout)
return _post
return self._dispatcher.dispatch(_prepare, context=_call_context(input_type, logging_obj))
@log_guardrail_information
async def apply_guardrail(
self,
@ -418,6 +484,31 @@ class GenericGuardrailAPI(CustomGuardrail):
Raises:
Exception: If the guardrail blocks the request
"""
if not self.fire_and_forget:
return await self._apply_guardrail(inputs, request_data, input_type, logging_obj)
try:
return await self._apply_guardrail(inputs, request_data, input_type, logging_obj)
except Exception as e: # noqa: BLE001 # an observe-only guardrail must never fail the request
verbose_proxy_logger.warning(
"Generic Guardrail API (%s, fire_and_forget) call not dispatched. %s: %s",
self.guardrail_name,
_call_context(input_type, logging_obj),
e,
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=FIRE_AND_FORGET_NOT_DISPATCHED_REASON,
request_data=request_data,
guardrail_status="not_run",
)
return _passthrough_inputs(inputs)
async def _apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"],
) -> GenericGuardrailAPIInputs:
verbose_proxy_logger.debug("Generic Guardrail API: Applying guardrail to text")
# Extract texts and images from inputs
@ -469,15 +560,24 @@ class GenericGuardrailAPI(CustomGuardrail):
model=model,
)
headers: Final = self._build_request_headers()
if self.fire_and_forget:
dispatched: Final = self._dispatch_background_post(
guardrail_request=guardrail_request, input_type=input_type, logging_obj=logging_obj
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=(
FIRE_AND_FORGET_DISPATCHED_REASON if dispatched else FIRE_AND_FORGET_DROPPED_REASON
),
request_data=request_data,
guardrail_status="success" if dispatched else "not_run",
)
return _passthrough_inputs(inputs)
# Make the API request
headers: Final = self._build_request_headers()
# Use mode="json" to ensure all iterables are converted to lists
payload: Final = guardrail_request.model_dump(mode="json")
response: Final = await self.async_handler.post(
url=self.api_base,
json=guardrail_request.model_dump(mode="json"),
headers=headers,
timeout=self.timeout,
url=self.api_base, json=payload, headers=headers, timeout=self.timeout
)
response.raise_for_status()

View file

@ -9,8 +9,10 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
import copy
import json
from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence
from contextlib import aclosing
from typing import TYPE_CHECKING, Any, Final, Protocol
import anyio
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
@ -19,6 +21,7 @@ from litellm.cost_calculator import _infer_call_type
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
@ -960,6 +963,70 @@ class UnifiedLLMGuardrails(CustomLogger):
choices: Final = _chunk_choices(item)
return any(getattr(choice, "finish_reason", None) is not None for choice in choices)
async def _stream_then_observe(
self,
*,
guardrail_to_apply: CustomGuardrail,
response: AsyncIterable[object],
request_data: dict[str, object],
user_api_key_dict: UserAPIKeyAuth,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
) -> AsyncGenerator[object, None]:
streamed: Final[list[object]] = [] # mutable-ok: records chunks as they are forwarded
try:
async for item in response:
streamed.append(item)
yield item
finally:
with anyio.CancelScope(shield=True):
await self._observe_streamed(
streamed=streamed,
guardrail_to_apply=guardrail_to_apply,
request_data=request_data,
user_api_key_dict=user_api_key_dict,
mappings=mappings,
)
@staticmethod
async def _observe_streamed(
*,
streamed: list[object],
guardrail_to_apply: CustomGuardrail,
request_data: dict[str, object],
user_api_key_dict: UserAPIKeyAuth,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
) -> None:
if not streamed:
return
try:
route_call_types: Final = (
None
if user_api_key_dict.request_route is None
else get_call_types_for_route(user_api_key_dict.request_route)
)
call_type: Final = (
route_call_types[0].value
if route_call_types
else _infer_call_type(call_type=None, completion_response=streamed[0])
)
handler_cls: Final = None if call_type is None else mappings.get(CallTypes(call_type))
if handler_cls is None:
return
logging_obj: Final = request_data.get("litellm_logging_obj")
await handler_cls().process_output_streaming_response(
responses_so_far=copy.deepcopy(streamed),
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=logging_obj if isinstance(logging_obj, LiteLLMLoggingObj) else None,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
except Exception as e: # noqa: BLE001 # an observe-only guardrail must never break the stream
verbose_proxy_logger.warning(
"UnifiedLLMGuardrails: observe-only stream check for %s failed: %s",
guardrail_to_apply.guardrail_name,
e,
)
def resolve_streaming_flag(self, guardrail_to_apply: CustomGuardrail | None, name: str, default: object) -> object:
"""Streaming flag resolution order (later wins): default < guardrail
attribute < guardrail_config dict < this callback's optional_params."""
@ -1051,6 +1118,20 @@ class UnifiedLLMGuardrails(CustomLogger):
mappings: Final = load_guardrail_translation_mappings()
if _streaming_flag("streaming_observe_only", False):
async with aclosing(
self._stream_then_observe(
guardrail_to_apply=guardrail_to_apply,
response=response,
request_data=request_data,
user_api_key_dict=user_api_key_dict,
mappings=mappings,
)
) as observed:
async for observed_item in observed:
yield observed_item
return
# Streaming text transformation (incremental_diff) diverges enough from the
# block_only path that it runs as its own iterator. It requires a route we
# can resolve up front to an OpenAI-chat handler (the only supported v1

View file

@ -103,6 +103,25 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
fire_and_forget: bool | None = Field(
default=None,
description=(
"Observe-only mode: the guardrail call runs in the background and the request never waits for it, so "
"BLOCKED and GUARDRAIL_INTERVENED answers are ignored. A dispatched call is recorded as success. "
"A stream sends one call when it closes, with whatever reached the client. The background call uses "
"timeout, or 30 seconds when unset. Defaults to false."
),
)
fire_and_forget_max_inflight: int | None = Field(
default=None,
ge=1,
description=(
"Maximum number of fire_and_forget calls in flight at once for this guardrail, per worker. Calls beyond it "
"are dropped and recorded as not_run. Defaults to 100."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],

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

@ -0,0 +1,117 @@
import asyncio
import contextvars
from collections.abc import Awaitable, Callable
from functools import partial
from typing import Final
import pytest
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import (
DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT,
BackgroundDispatcher,
fire_and_forget_from_config,
max_inflight_from_config,
)
_request_scoped: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar("request_scoped", default=None)
@pytest.mark.parametrize(
("value", "expected"),
[(None, False), (True, True), (False, False), ("true", True), ("false", False), (1, True), (0, False)],
)
def test_fire_and_forget_accepts_bools_and_their_config_spellings(value: object, expected: bool) -> None:
assert fire_and_forget_from_config(value) is expected
@pytest.mark.parametrize("value", ["maybe", 2, [True]])
def test_an_unparseable_fire_and_forget_is_ignored_with_a_warning(
value: object, warning_messages: Callable[[str], list[str]]
) -> None:
assert fire_and_forget_from_config(value) is False
assert len(warning_messages("Ignoring fire_and_forget=")) == 1
@pytest.mark.parametrize(("value", "expected"), [(None, DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT), (5, 5), ("5", 5)])
def test_max_inflight_accepts_positive_integers(value: object, expected: int) -> None:
assert max_inflight_from_config(value) == expected
@pytest.mark.parametrize("value", [0, -1, 2.5, True, "many"])
def test_an_invalid_max_inflight_falls_back_to_the_default_with_a_warning(
value: object, warning_messages: Callable[[str], list[str]]
) -> None:
assert max_inflight_from_config(value) == DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT
assert len(warning_messages("Ignoring fire_and_forget_max_inflight=")) == 1
def test_a_dispatcher_needs_room_for_at_least_one_call() -> None:
with pytest.raises(ValueError, match="fire_and_forget_max_inflight"):
BackgroundDispatcher(guardrail_name="g", max_inflight=0)
async def test_calls_beyond_the_cap_are_dropped_unprepared_counted_and_warned_once(
warning_messages: Callable[[str], list[str]],
) -> None:
gate: Final = asyncio.Event()
dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=2)
prepared: Final[list[int]] = [] # mutable-ok: records which calls were prepared
def prepare(call: int) -> Callable[[], Awaitable[bool]]:
prepared.append(call)
return gate.wait
dispatched: Final = [dispatcher.dispatch(partial(prepare, i), context=f"call {i}") for i in range(5)]
assert (dispatched, dispatcher.pending_count, dispatcher.dropped_count) == ([True, True, False, False, False], 2, 3)
assert prepared == [0, 1]
assert len(warning_messages("dropped")) == 1
gate.set()
await dispatcher.wait_for_pending()
async def test_a_finished_call_frees_its_slot() -> None:
dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=1)
async def finish() -> None:
return None
for _ in range(3):
assert dispatcher.dispatch(lambda: finish, context="call") is True
await dispatcher.wait_for_pending()
assert (dispatcher.pending_count, dispatcher.dropped_count) == (0, 0)
async def test_a_failing_call_is_logged_with_its_context_and_not_raised(
warning_messages: Callable[[str], list[str]],
) -> None:
dispatcher: Final = BackgroundDispatcher(guardrail_name="audit", max_inflight=1)
async def fail() -> None:
raise ConnectionError("refused")
dispatcher.dispatch(lambda: fail, context="input_type=response litellm_call_id=call-1")
await dispatcher.wait_for_pending()
assert warning_messages("call failed") == [
"Generic Guardrail API (audit, fire_and_forget) call failed. "
"input_type=response litellm_call_id=call-1: refused"
]
async def test_a_dispatched_call_does_not_see_the_request_context() -> None:
dispatcher: Final = BackgroundDispatcher(guardrail_name="g", max_inflight=1)
seen: Final[list[str | None]] = [] # mutable-ok: records what the background call saw
async def record() -> None:
seen.append(_request_scoped.get())
token: Final = _request_scoped.set("request-1")
try:
dispatcher.dispatch(lambda: record, context="call")
finally:
_request_scoped.reset(token)
await dispatcher.wait_for_pending()
assert seen == [None]

View file

@ -0,0 +1,501 @@
import asyncio
import copy
import json
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping
from types import SimpleNamespace
from typing import Final
import httpx
import pydantic
import pytest
import litellm
from litellm.exceptions import GuardrailRaisedException
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPI,
initialize_guardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import (
DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT,
FIRE_AND_FORGET_DISPATCHED_REASON,
FIRE_AND_FORGET_DROPPED_REASON,
FIRE_AND_FORGET_NOT_DISPATCHED_REASON,
FIRE_AND_FORGET_POST_TIMEOUT_SECONDS,
BackgroundDispatcher,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.types.guardrails import LitellmParams
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIOptionalParams,
)
from litellm.types.utils import Delta, GenericGuardrailAPIInputs, ModelResponseStream
API_BASE: Final = "https://api.test.guardrail.com"
CLIENT_TIMEOUT_SECONDS: Final = 600.0
_STREAM_WORDS: Final = ("Hello", " ", "world", "!", " Bye")
class _Endpoint:
def __init__(
self,
*,
body: Mapping[str, object] | None = None,
status_code: int = 200,
error: Exception | None = None,
gate_open: bool = True,
) -> None:
self.gate: Final = asyncio.Event()
if gate_open:
self.gate.set()
self.requests: Final[list[tuple[str, dict[str, str]]]] = [] # mutable-ok: records each request
self.payloads: Final[list[dict[str, object]]] = [] # mutable-ok: records each request body
self.read_timeouts: Final[list[float | None]] = [] # mutable-ok: records each request timeout
self.completed: int = 0
self._finished: Final = asyncio.Condition()
self._body: Final = dict(body or {"action": "NONE"})
self._status_code: Final = status_code
self._error: Final = error
async def __call__(self, request: httpx.Request) -> httpx.Response:
self.requests.append((str(request.url), dict(request.headers)))
self.payloads.append(json.loads(request.content))
self.read_timeouts.append(request.extensions["timeout"]["read"])
await self.gate.wait()
async with self._finished:
self.completed += 1
self._finished.notify_all()
if self._error is not None:
raise self._error
return httpx.Response(self._status_code, json=self._body)
async def wait_for_completed(self, count: int) -> None:
async with self._finished:
await asyncio.wait_for(self._finished.wait_for(lambda: self.completed >= count), timeout=5)
def handler(self) -> AsyncHTTPHandler:
return AsyncHTTPHandler(timeout=CLIENT_TIMEOUT_SECONDS, transport=httpx.MockTransport(self))
def _logging_obj(call_id: str = "call-123") -> SimpleNamespace:
return SimpleNamespace(litellm_call_id=call_id, litellm_trace_id="trace-123", model_call_details={})
def _guardrail(
endpoint: _Endpoint, *, name: str = "ff-guardrail", event_hook: str = "pre_call", **options: object
) -> GenericGuardrailAPI:
return GenericGuardrailAPI(
api_base=API_BASE,
guardrail_name=name,
event_hook=event_hook,
default_on=True,
async_handler=endpoint.handler(),
**options,
)
def _fire_and_forget(
endpoint: _Endpoint, *, max_inflight: int = 10, **options: object
) -> tuple[GenericGuardrailAPI, BackgroundDispatcher]:
dispatcher: Final = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=max_inflight)
return _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget=True, **options), dispatcher
def _request_data() -> dict[str, object]: # mutable-ok: apply_guardrail records entries into it
return {
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"user_api_key_hash": "hash-1", "user_api_key_team_id": "team-1"},
}
def _recorded_outcomes(request_data: Mapping[str, object]) -> list[tuple[str, str]]:
metadata: Final = request_data["metadata"]
assert isinstance(metadata, dict)
return [
(entry["guardrail_status"], entry["guardrail_response"])
for entry in metadata["standard_logging_guardrail_information"]
]
async def test_returns_before_the_post_completes() -> None:
endpoint: Final = _Endpoint(gate_open=False)
guardrail, dispatcher = _fire_and_forget(endpoint)
inputs: Final = GenericGuardrailAPIInputs(
texts=["hello"], structured_messages=[{"role": "user", "content": "hello"}]
)
result: Final = await asyncio.wait_for(
guardrail.apply_guardrail(
inputs=inputs, request_data=_request_data(), input_type="request", logging_obj=_logging_obj()
),
timeout=5,
)
assert (result, endpoint.completed, dispatcher.pending_count) == (inputs, 0, 1)
endpoint.gate.set()
await dispatcher.wait_for_pending()
assert (endpoint.completed, dispatcher.pending_count) == (1, 0)
async def test_background_post_reaches_the_same_url_with_the_same_headers_and_payload() -> None:
awaited_endpoint: Final = _Endpoint()
background_endpoint: Final = _Endpoint()
auth: Final = {"api_key": "audit-key", "headers": {"x-static": "static-value"}}
guardrail, dispatcher = _fire_and_forget(background_endpoint, **auth)
inputs: Final = GenericGuardrailAPIInputs(texts=["hello"], images=["data:image/png;base64,AAAA"], model="gpt-4o")
for target in (_guardrail(awaited_endpoint, **auth), guardrail):
await target.apply_guardrail(
inputs=GenericGuardrailAPIInputs(**inputs),
request_data=_request_data(),
input_type="request",
logging_obj=_logging_obj(),
)
await dispatcher.wait_for_pending()
assert background_endpoint.requests == awaited_endpoint.requests
assert background_endpoint.payloads == awaited_endpoint.payloads
url, headers = background_endpoint.requests[0]
assert (url, headers["x-api-key"], headers["x-static"]) == (
f"{API_BASE}/beta/litellm_basic_guardrail_api",
"audit-key",
"static-value",
)
@pytest.mark.parametrize(
("configured_timeout", "expected_background_timeout"),
[(None, FIRE_AND_FORGET_POST_TIMEOUT_SECONDS), (5.0, 5.0)],
)
async def test_background_post_honors_the_configured_timeout(
configured_timeout: float | None, expected_background_timeout: float
) -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, timeout=configured_timeout)
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
await dispatcher.wait_for_pending()
assert endpoint.read_timeouts == [expected_background_timeout]
@pytest.mark.parametrize(
"body",
[{"action": "BLOCKED", "blocked_reason": "nope"}, {"action": "GUARDRAIL_INTERVENED", "texts": ["MASKED"]}],
)
async def test_the_endpoint_verdict_is_ignored(body: Mapping[str, object]) -> None:
endpoint: Final = _Endpoint(body=body)
guardrail, dispatcher = _fire_and_forget(endpoint)
result: Final = await guardrail.apply_guardrail(
inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request"
)
await dispatcher.wait_for_pending()
assert (result, endpoint.completed) == ({"texts": ["my ssn is 123"]}, 1)
@pytest.mark.parametrize(
"endpoint_options",
[{"error": httpx.ConnectError("connection refused")}, {"status_code": 500}],
)
async def test_a_failing_endpoint_is_logged_not_raised(
endpoint_options: Mapping[str, object], warning_messages: Callable[[str], list[str]]
) -> None:
endpoint: Final = _Endpoint(**endpoint_options)
guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=True, unreachable_fallback="fail_closed")
result: Final = await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
request_data={},
input_type="response",
logging_obj=_logging_obj(call_id="call-failing"),
)
await dispatcher.wait_for_pending()
assert result == {"texts": ["hello"]}
failures: Final = warning_messages("call failed")
assert len(failures) == 1
assert "ff-guardrail" in failures[0]
assert "input_type=response litellm_call_id=call-failing" in failures[0]
@pytest.mark.parametrize("fail_on_error", [True, False])
@pytest.mark.parametrize(
("inputs", "request_data"),
[
({"texts": ["hi"], "tools": [{"function": {"name": "f"}}]}, _request_data()),
({"texts": ["hi"]}, {"messages": [], "metadata": None}),
],
ids=["tool_without_type", "malformed_request_metadata"],
)
async def test_a_failure_before_dispatch_is_logged_and_passes_through(
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object], # mutable-ok: apply_guardrail records entries into it
fail_on_error: bool,
warning_messages: Callable[[str], list[str]],
) -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=fail_on_error)
request: Final = copy.deepcopy(request_data)
result: Final = await guardrail.apply_guardrail(
inputs=inputs, request_data=request, input_type="request", logging_obj=_logging_obj("call-bad")
)
await dispatcher.wait_for_pending()
assert (result, endpoint.payloads) == (inputs, [])
assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)]
warnings: Final = warning_messages("not dispatched")
assert len(warnings) == 1
assert "litellm_call_id=call-bad" in warnings[0]
async def test_a_payload_that_cannot_be_serialized_is_recorded_as_not_dispatched() -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, additional_provider_specific_params={"bad": object()})
request: Final = _request_data()
result: Final = await guardrail.apply_guardrail(
inputs={"texts": ["hello"]}, request_data=request, input_type="request"
)
await dispatcher.wait_for_pending()
assert (result, endpoint.payloads) == ({"texts": ["hello"]}, [])
assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)]
async def test_a_call_dropped_by_the_cap_never_builds_its_payload() -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(
endpoint, max_inflight=1, additional_provider_specific_params={"bad": object()}
)
gate: Final = asyncio.Event()
dispatcher.dispatch(lambda: gate.wait, context="occupies the only slot")
request: Final = _request_data()
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request, input_type="request")
gate.set()
await dispatcher.wait_for_pending()
assert _recorded_outcomes(request) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)]
async def test_dispatched_calls_are_recorded_as_success_and_dropped_ones_as_not_run() -> None:
endpoint: Final = _Endpoint(gate_open=False)
guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1)
dispatched, dropped = _request_data(), _request_data()
for request_data in (dispatched, dropped):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request")
endpoint.gate.set()
await dispatcher.wait_for_pending()
assert _recorded_outcomes(dispatched) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)]
assert _recorded_outcomes(dropped) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)]
async def test_a_configured_max_inflight_bounds_dispatch() -> None:
endpoint: Final = _Endpoint(gate_open=False)
guardrail: Final = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=1)
first, second = _request_data(), _request_data()
for request_data in (first, second):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request")
endpoint.gate.set()
await endpoint.wait_for_completed(1)
assert _recorded_outcomes(first) + _recorded_outcomes(second) == [
("success", FIRE_AND_FORGET_DISPATCHED_REASON),
("not_run", FIRE_AND_FORGET_DROPPED_REASON),
]
assert len(endpoint.payloads) == 1
@pytest.mark.parametrize("max_inflight", [None, 0, "many"])
async def test_an_unset_or_invalid_max_inflight_admits_the_default_number_of_calls(max_inflight: object) -> None:
endpoint: Final = _Endpoint(gate_open=False)
guardrail: Final = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=max_inflight)
calls: Final = [_request_data() for _ in range(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + 1)]
for request_data in calls:
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request")
endpoint.gate.set()
await endpoint.wait_for_completed(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT)
outcomes: Final = [_recorded_outcomes(request_data)[0][0] for request_data in calls]
assert outcomes == ["success"] * DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + ["not_run"]
def _stream_chunks() -> list[ModelResponseStream]:
return [
ModelResponseStream(
model="gpt-4",
choices=[
litellm.StreamingChoices(
index=0,
delta=Delta(role="assistant", content=word),
finish_reason="stop" if i == len(_STREAM_WORDS) - 1 else None,
)
],
)
for i, word in enumerate(_STREAM_WORDS)
]
async def _upstream(*, fail_after: int | None = None) -> AsyncIterator[ModelResponseStream]:
for i, chunk in enumerate(_stream_chunks()):
if i == fail_after:
raise ConnectionError("upstream reset")
yield chunk
def _guarded_stream(
guardrail: GenericGuardrailAPI, upstream: AsyncIterator[ModelResponseStream]
) -> AsyncGenerator[ModelResponseStream, None]:
return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/chat/completions"),
response=upstream,
request_data={
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["ff-guardrail"]},
},
)
async def _streamed_texts(guardrail: GenericGuardrailAPI) -> list[str]:
return [chunk.choices[0].delta.content or "" async for chunk in _guarded_stream(guardrail, _upstream())]
@pytest.mark.parametrize("streaming_transform_mode", ["block_only", "incremental_diff"])
async def test_a_stream_is_emitted_live_and_sends_one_call_with_the_whole_text(
streaming_transform_mode: str,
) -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(
endpoint,
event_hook="post_call",
streaming_end_of_stream_only=False,
streaming_sampling_rate=1,
streaming_transform_mode=streaming_transform_mode,
)
streamed: Final = await _streamed_texts(guardrail)
await dispatcher.wait_for_pending()
assert streamed == list(_STREAM_WORDS), "every chunk must be emitted as it arrives"
assert [payload["texts"] for payload in endpoint.payloads] == [["Hello world! Bye"]]
async def test_an_abandoned_stream_still_sends_what_reached_the_client() -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, event_hook="post_call")
stream: Final = _guarded_stream(guardrail, _upstream())
received: Final = [await anext(stream), await anext(stream)]
await stream.aclose()
await dispatcher.wait_for_pending()
assert [chunk.choices[0].delta.content for chunk in received] == ["Hello", " "]
assert [payload["texts"] for payload in endpoint.payloads] == [["Hello "]]
async def test_a_stream_that_fails_upstream_still_sends_what_reached_the_client() -> None:
endpoint: Final = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, event_hook="post_call")
with pytest.raises(ConnectionError, match="upstream reset"):
async for _ in _guarded_stream(guardrail, _upstream(fail_after=3)):
pass
await dispatcher.wait_for_pending()
assert [payload["texts"] for payload in endpoint.payloads] == [["Hello world"]]
async def test_an_awaited_stream_is_checked_per_sampled_chunk() -> None:
endpoint: Final = _Endpoint()
await _streamed_texts(_guardrail(endpoint, event_hook="post_call", streaming_sampling_rate=1))
assert len(endpoint.payloads) > 1
def test_the_observe_only_warning_is_logged_only_when_enabled(warning_messages: Callable[[str], list[str]]) -> None:
_guardrail(_Endpoint(), name="enforcing")
_guardrail(_Endpoint(), name="observer", fire_and_forget=True)
warnings: Final = warning_messages("observe-only")
assert len(warnings) == 1
assert "observer" in warnings[0]
@pytest.mark.parametrize("value", ["false", "maybe", 2])
async def test_a_quoted_false_or_unparseable_fire_and_forget_keeps_the_guardrail_enforcing(value: object) -> None:
endpoint: Final = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"})
guardrail: Final = _guardrail(endpoint, fire_and_forget=value)
with pytest.raises(GuardrailRaisedException, match="nope"):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
async def test_a_quoted_true_fire_and_forget_turns_on_observe_only() -> None:
endpoint: Final = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"})
dispatcher: Final = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=1)
quoted: Final = _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget="true")
result: Final = await quoted.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
await dispatcher.wait_for_pending()
assert (result, endpoint.completed) == ({"texts": ["hello"]}, 1)
@pytest.mark.parametrize("max_inflight", [0, -1])
def test_the_config_form_rejects_a_max_inflight_below_one(max_inflight: int) -> None:
with pytest.raises(pydantic.ValidationError):
GenericGuardrailAPIOptionalParams(fire_and_forget_max_inflight=max_inflight)
async def test_initialize_guardrail_forwards_fire_and_forget_and_max_inflight() -> None:
endpoint: Final = _Endpoint(gate_open=False)
guardrail: Final = initialize_guardrail(
LitellmParams(
guardrail="generic_guardrail_api",
mode="pre_call",
api_base=API_BASE,
default_on=True,
fire_and_forget=True,
fire_and_forget_max_inflight=1,
),
{"guardrail_name": "from-config"},
)
guardrail.async_handler = endpoint.handler()
first, second = _request_data(), _request_data()
try:
for request_data in (first, second):
await asyncio.wait_for(
guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request"),
timeout=5,
)
endpoint.gate.set()
await endpoint.wait_for_completed(1)
finally:
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
assert _recorded_outcomes(first) + _recorded_outcomes(second) == [
("success", FIRE_AND_FORGET_DISPATCHED_REASON),
("not_run", FIRE_AND_FORGET_DROPPED_REASON),
]
async def test_by_default_the_endpoint_is_awaited_and_can_block() -> None:
endpoint: Final = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"})
guardrail: Final = _guardrail(endpoint)
with pytest.raises(GuardrailRaisedException, match="nope"):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
assert endpoint.completed == 1

View file

@ -0,0 +1,117 @@
import asyncio
from collections.abc import AsyncIterator, Callable
from typing import Final, Literal
import anyio
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails
from litellm.types.utils import Delta, GenericGuardrailAPIInputs, ModelResponseStream
_WORDS: Final = ("Hello", " ", "world", "!")
class _Observer(CustomGuardrail):
def __init__(self, *, failure: Exception | None = None) -> None:
super().__init__(guardrail_name="observer", event_hook="post_call", default_on=True)
self.streaming_observe_only = True
self.observed: Final[list[list[str]]] = [] # mutable-ok: records each observed call
self._failure: Final = failure
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: object = None,
) -> GenericGuardrailAPIInputs:
loop: Final = asyncio.get_running_loop()
next_iteration: Final = loop.create_future()
loop.call_soon(next_iteration.set_result, None)
await next_iteration
if self._failure is not None:
raise self._failure
self.observed.append(list(inputs.get("texts", [])))
return inputs
def _chunk(word: str) -> ModelResponseStream:
return ModelResponseStream(
model="gpt-4", choices=[litellm.StreamingChoices(index=0, delta=Delta(role="assistant", content=word))]
)
async def _upstream(
*, words: tuple[str, ...] = _WORDS, stall_after: int | None = None
) -> AsyncIterator[ModelResponseStream]:
for i, word in enumerate(words):
if i == stall_after:
await asyncio.Event().wait()
yield _chunk(word)
def _guarded_stream(
observer: _Observer, upstream: AsyncIterator[object], *, route: str | None = "/chat/completions"
) -> AsyncIterator[object]:
return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route=route),
response=upstream,
request_data={
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": observer,
"metadata": {"guardrails": ["observer"]},
},
)
async def test_a_client_disconnect_mid_stream_is_still_observed() -> None:
observer: Final = _Observer()
two_chunks_sent: Final = asyncio.Event()
sent: Final[list[object]] = [] # mutable-ok: records what reached the client
async def client() -> None:
async for chunk in _guarded_stream(observer, _upstream(stall_after=2)):
sent.append(chunk)
if len(sent) == 2:
two_chunks_sent.set()
async with anyio.create_task_group() as requests:
requests.start_soon(client)
await two_chunks_sent.wait()
requests.cancel_scope.cancel()
assert observer.observed == [["Hello "]]
async def test_a_failing_observer_never_breaks_the_stream(warning_messages: Callable[[str], list[str]]) -> None:
observer: Final = _Observer(failure=RuntimeError("observer down"))
sent: Final = [chunk async for chunk in _guarded_stream(observer, _upstream())]
assert len(sent) == len(_WORDS)
assert warning_messages("observe-only stream check") == [
"UnifiedLLMGuardrails: observe-only stream check for observer failed: observer down"
]
async def test_a_stream_no_handler_understands_is_forwarded_unobserved(
warning_messages: Callable[[str], list[str]],
) -> None:
observer: Final = _Observer()
async def raw_bytes() -> AsyncIterator[bytes]:
yield b"data: raw\n\n"
sent: Final = [chunk async for chunk in _guarded_stream(observer, raw_bytes(), route=None)]
assert (sent, observer.observed, warning_messages("observe-only stream check")) == ([b"data: raw\n\n"], [], [])
async def test_an_empty_stream_is_not_observed(warning_messages: Callable[[str], list[str]]) -> None:
observer: Final = _Observer()
sent: Final = [chunk async for chunk in _guarded_stream(observer, _upstream(words=()), route=None)]
assert (sent, observer.observed, warning_messages("observe-only stream check")) == ([], [], [])