This commit is contained in:
Caduri 2026-09-30 10:27:15 -04:00 • committed by GitHub
commit 431f66ac89
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 673 additions and 10 deletions

View file

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

View file

@ -0,0 +1,82 @@
import asyncio
import contextvars
from collections.abc import Awaitable, Callable
from typing import Final
from litellm._logging import verbose_proxy_logger
DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT: Final = 100
FIRE_AND_FORGET_POST_TIMEOUT_SECONDS: Final = 30.0
FIRE_AND_FORGET_DISPATCHED_REASON: Final = "fire_and_forget dispatched, verdict not read"
FIRE_AND_FORGET_DROPPED_REASON: Final = "fire_and_forget_max_inflight reached, call dropped"
FIRE_AND_FORGET_NOT_DISPATCHED_REASON: Final = "fire_and_forget payload could not be built, call not dispatched"
_DROP_LOG_INTERVAL: Final = 100
def resolve_max_inflight(value: object) -> int:
if value is None:
return DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"fire_and_forget_max_inflight must be an int, got {value!r}")
return value
class BackgroundDispatcher:
"""Runs calls as detached tasks, dropping (and counting) calls once ``max_inflight`` are outstanding."""
def __init__(self, *, guardrail_name: str | None, max_inflight: int) -> None:
if max_inflight < 1:
raise ValueError(f"fire_and_forget_max_inflight must be >= 1 (got {max_inflight})")
self._guardrail_name: Final = guardrail_name
self._max_inflight: Final = max_inflight
self._pending: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong refs, asyncio keeps only weak ones
self._dropped: int = 0
@property
def pending_count(self) -> int:
return len(self._pending)
@property
def dropped_count(self) -> int:
return self._dropped
def dispatch(self, run: Callable[[], Awaitable[None]], *, context: str) -> bool:
if len(self._pending) >= self._max_inflight:
self._dropped += 1
if self._dropped % _DROP_LOG_INTERVAL == 1:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s, fire_and_forget): dropped %d call(s) so far, "
"%d already in flight (fire_and_forget_max_inflight=%d). %s",
self._guardrail_name,
self._dropped,
len(self._pending),
self._max_inflight,
context,
)
return False
task: Final = contextvars.Context().run(asyncio.create_task, self._run_logging_failures(run, context=context))
self._pending.add(task)
task.add_done_callback(self._pending.discard)
return True
async def _run_logging_failures(self, run: Callable[[], Awaitable[None]], *, context: str) -> None:
try:
await run()
except Exception as e: # noqa: BLE001 # a detached task has no caller to raise into
verbose_proxy_logger.warning(
"Generic Guardrail API (%s, fire_and_forget) call failed. %s: %s",
self._guardrail_name,
context,
e,
)
async def wait_for_pending(self) -> None:
pending: Final = tuple(self._pending)
if pending:
await asyncio.gather(*pending, return_exceptions=True)

View file

@ -11,6 +11,7 @@ from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
import httpx
from pydantic import JsonValue
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
@ -20,6 +21,7 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
@ -33,6 +35,15 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
)
from litellm.types.utils import GenericGuardrailAPIInputs
from .background_dispatch import (
FIRE_AND_FORGET_DISPATCHED_REASON,
FIRE_AND_FORGET_DROPPED_REASON,
FIRE_AND_FORGET_NOT_DISPATCHED_REASON,
FIRE_AND_FORGET_POST_TIMEOUT_SECONDS,
BackgroundDispatcher,
resolve_max_inflight,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -170,6 +181,15 @@ def _structured_rows_to_write_back(
)
def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(**inputs)
def _call_context(input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"]) -> str:
call_id: Final = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None
return f"input_type={input_type} litellm_call_id={call_id}"
class GenericGuardrailAPI(CustomGuardrail):
"""
Generic Guardrail API integration for LiteLLM.
@ -204,9 +224,15 @@ class GenericGuardrailAPI(CustomGuardrail):
streaming_end_of_stream_only: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
fire_and_forget: bool | None = None,
fire_and_forget_max_inflight: int | None = None,
async_handler: AsyncHTTPHandler | None = None,
dispatcher: BackgroundDispatcher | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.headers = headers or {}
self.extra_headers = extra_headers or []
@ -235,9 +261,14 @@ class GenericGuardrailAPI(CustomGuardrail):
self.fail_on_error: bool = True if fail_on_error is None else fail_on_error
if fire_and_forget is not None and not isinstance(fire_and_forget, bool): # pyright: ignore[reportUnnecessaryIsInstance] # config extras reach here unvalidated
raise ValueError(f"fire_and_forget must be a bool, got {fire_and_forget!r}")
self.fire_and_forget: bool = fire_and_forget is True
# Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook
# via getattr(guardrail_to_apply, "streaming_*", default).
self.streaming_end_of_stream_only: bool = (
# via getattr(guardrail_to_apply, "streaming_*", default). Forced on under
# fire_and_forget so a stream dispatches one call, not one per sampled chunk.
self.streaming_end_of_stream_only: bool = self.fire_and_forget or (
False if streaming_end_of_stream_only is None else streaming_end_of_stream_only
)
if streaming_sampling_rate is not None and streaming_sampling_rate < 1:
@ -256,6 +287,22 @@ class GenericGuardrailAPI(CustomGuardrail):
super().__init__(**kwargs)
self._dispatcher: Final = dispatcher or BackgroundDispatcher(
guardrail_name=self.guardrail_name,
max_inflight=resolve_max_inflight(fire_and_forget_max_inflight),
)
if self.fire_and_forget:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): fire_and_forget=True makes this guardrail observe-only. "
"action=BLOCKED and action=GUARDRAIL_INTERVENED are ignored, fail_on_error=%s and "
"unreachable_fallback=%s cannot block the request, and streaming is forced to "
"end-of-stream observation.",
self.guardrail_name,
self.fail_on_error,
self.unreachable_fallback,
)
verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base)
def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata:
@ -377,6 +424,8 @@ class GenericGuardrailAPI(CustomGuardrail):
logging_obj: Optional["LiteLLMLoggingObj"],
is_unreachable: bool = True,
) -> GenericGuardrailAPIInputs:
if self.fire_and_forget:
raise error
unreachable_fail_open: Final = is_unreachable and self.unreachable_fallback == "fail_open"
if unreachable_fail_open or not self.fail_on_error:
http_status_code: Final = getattr(getattr(error, "response", None), "status_code", None)
@ -390,6 +439,24 @@ class GenericGuardrailAPI(CustomGuardrail):
verbose_proxy_logger.error("Generic Guardrail API: failed to make request: %s", str(error))
raise Exception(f"Generic Guardrail API failed: {error}")
def _dispatch_background_post(
self,
*,
payload: Mapping[str, JsonValue],
headers: Mapping[str, str],
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"],
) -> bool:
async def _post() -> None:
await self.async_handler.post(
url=self.api_base,
json=dict(payload),
headers=dict(headers),
timeout=FIRE_AND_FORGET_POST_TIMEOUT_SECONDS,
)
return self._dispatcher.dispatch(_post, context=_call_context(input_type, logging_obj))
@log_guardrail_information
async def apply_guardrail(
self,
@ -418,6 +485,31 @@ class GenericGuardrailAPI(CustomGuardrail):
Raises:
Exception: If the guardrail blocks the request
"""
if not self.fire_and_forget:
return await self._apply_guardrail(inputs, request_data, input_type, logging_obj)
try:
return await self._apply_guardrail(inputs, request_data, input_type, logging_obj)
except Exception as e: # noqa: BLE001 # an observe-only guardrail must never fail the request
verbose_proxy_logger.warning(
"Generic Guardrail API (%s, fire_and_forget) call not dispatched. %s: %s",
self.guardrail_name,
_call_context(input_type, logging_obj),
e,
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=FIRE_AND_FORGET_NOT_DISPATCHED_REASON,
request_data=request_data or {},
guardrail_status="not_run",
)
return _passthrough_inputs(inputs)
async def _apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"],
) -> GenericGuardrailAPIInputs:
verbose_proxy_logger.debug("Generic Guardrail API: Applying guardrail to text")
# Extract texts and images from inputs
@ -470,14 +562,23 @@ class GenericGuardrailAPI(CustomGuardrail):
)
headers: Final = self._build_request_headers()
# Make the API request
# Use mode="json" to ensure all iterables are converted to lists
response: Final = await self.async_handler.post(
url=self.api_base,
json=guardrail_request.model_dump(mode="json"),
headers=headers,
)
payload: Final = guardrail_request.model_dump(mode="json")
if self.fire_and_forget:
dispatched: Final = self._dispatch_background_post(
payload=payload, headers=headers, input_type=input_type, logging_obj=logging_obj
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=(
FIRE_AND_FORGET_DISPATCHED_REASON if dispatched else FIRE_AND_FORGET_DROPPED_REASON
),
request_data=request_data,
guardrail_status="success" if dispatched else "not_run",
)
return _passthrough_inputs(inputs)
response: Final = await self.async_handler.post(url=self.api_base, json=payload, headers=headers)
response.raise_for_status()
response_json: Final = response.json()

View file

@ -103,6 +103,35 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
fire_and_forget: bool | None = Field(
default=None,
description=(
"If True, the guardrail HTTP call runs as a background task and the request proceeds "
"without waiting for the response, in every mode (pre_call, during_call, post_call). "
"The guardrail becomes observe-only: action=BLOCKED and action=GUARDRAIL_INTERVENED "
"are ignored, and fail_on_error / unreachable_fallback cannot block the request. The "
"background call has a fixed 30 second timeout. A dispatched call is recorded in the "
"guardrail logs as guardrail_status=success with a response saying the verdict was not "
"read. A call whose payload cannot be built is logged as a warning, passes the request "
"through, and is recorded as guardrail_status=not_run. Also forces "
"streaming_end_of_stream_only=True so a stream sends one call instead of one per "
"sampled chunk. Defaults to False in GenericGuardrailAPI.__init__ when None."
),
)
fire_and_forget_max_inflight: int | None = Field(
default=None,
ge=1,
description=(
"Maximum number of fire_and_forget calls in flight at once for this guardrail, per "
"worker process, so a slow guardrail endpoint cannot pile up background tasks without "
"limit. Calls beyond this limit are dropped and counted, with a rate-limited "
"warning, and recorded in the guardrail logs as guardrail_status=not_run. Must be >= 1, "
"and is validated at startup even when fire_and_forget is off. Only used when "
"fire_and_forget is True. Defaults to 100 in GenericGuardrailAPI.__init__ when None."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],

View file

@ -0,0 +1,449 @@
import asyncio
import contextvars
import json
import logging
from collections.abc import Iterator
from types import SimpleNamespace
import httpx
import pydantic
import pytest
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import GuardrailRaisedException
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPI,
initialize_guardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.background_dispatch import (
DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT,
FIRE_AND_FORGET_DISPATCHED_REASON,
FIRE_AND_FORGET_DROPPED_REASON,
FIRE_AND_FORGET_NOT_DISPATCHED_REASON,
FIRE_AND_FORGET_POST_TIMEOUT_SECONDS,
BackgroundDispatcher,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.types.guardrails import LitellmParams
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIOptionalParams,
)
from litellm.types.utils import Delta, ModelResponseStream
API_BASE = "https://api.test.guardrail.com"
CLIENT_TIMEOUT_SECONDS = 600.0
_request_scoped = contextvars.ContextVar("request_scoped", default=None)
class _Endpoint:
"""The guardrail server behind a MockTransport. Each request is recorded, then waits on ``gate``."""
def __init__(self, *, body=None, status_code=200, error=None, gate_open=True):
self.gate = asyncio.Event()
if gate_open:
self.gate.set()
self.payloads: list[dict] = []
self.read_timeouts: list[float | None] = []
self.seen_request_scoped: list[object] = []
self.completed = 0
self._body = body or {"action": "NONE"}
self._status_code = status_code
self._error = error
async def __call__(self, request: httpx.Request) -> httpx.Response:
self.payloads.append(json.loads(request.content))
self.read_timeouts.append(request.extensions["timeout"]["read"])
self.seen_request_scoped.append(_request_scoped.get())
await self.gate.wait()
if self._error is not None:
raise self._error
self.completed += 1
return httpx.Response(self._status_code, json=self._body)
def handler(self) -> AsyncHTTPHandler:
return AsyncHTTPHandler(timeout=CLIENT_TIMEOUT_SECONDS, transport=httpx.MockTransport(self))
def _logging_obj(call_id="call-123"):
return SimpleNamespace(litellm_call_id=call_id, litellm_trace_id="trace-123", model_call_details={})
def _guardrail(endpoint, *, name="ff-guardrail", event_hook="pre_call", **options):
return GenericGuardrailAPI(
api_base=API_BASE,
guardrail_name=name,
event_hook=event_hook,
default_on=True,
async_handler=endpoint.handler(),
**options,
)
def _fire_and_forget(endpoint, *, max_inflight=10, **options):
dispatcher = BackgroundDispatcher(guardrail_name="ff-guardrail", max_inflight=max_inflight)
return _guardrail(endpoint, dispatcher=dispatcher, fire_and_forget=True, **options), dispatcher
def _request_data():
return {
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"user_api_key_hash": "hash-1", "user_api_key_team_id": "team-1"},
}
@pytest.fixture
def captured_warnings() -> Iterator[list[logging.LogRecord]]:
records: list[logging.LogRecord] = []
handler = logging.Handler(level=logging.WARNING)
handler.emit = records.append
previous_level = verbose_proxy_logger.level
verbose_proxy_logger.addHandler(handler)
verbose_proxy_logger.setLevel(logging.WARNING)
yield records
verbose_proxy_logger.removeHandler(handler)
verbose_proxy_logger.setLevel(previous_level)
def _messages(records, needle):
return [m for m in (r.getMessage() for r in records) if needle in m]
async def test_returns_before_the_post_completes():
endpoint = _Endpoint(gate_open=False)
guardrail, dispatcher = _fire_and_forget(endpoint)
inputs = {"texts": ["hello"], "structured_messages": [{"role": "user", "content": "hello"}]}
result = await asyncio.wait_for(
guardrail.apply_guardrail(
inputs=inputs, request_data=_request_data(), input_type="request", logging_obj=_logging_obj()
),
timeout=5,
)
assert result == inputs
assert endpoint.completed == 0
assert dispatcher.pending_count == 1
endpoint.gate.set()
await dispatcher.wait_for_pending()
assert endpoint.completed == 1
assert dispatcher.pending_count == 0
async def test_endpoint_receives_the_same_payload_as_the_awaited_path():
awaited_endpoint = _Endpoint()
background_endpoint = _Endpoint()
guardrail, dispatcher = _fire_and_forget(background_endpoint)
inputs = {"texts": ["hello"], "images": ["data:image/png;base64,AAAA"], "model": "gpt-4o"}
for target in (_guardrail(awaited_endpoint), guardrail):
await target.apply_guardrail(
inputs=dict(inputs), request_data=_request_data(), input_type="request", logging_obj=_logging_obj()
)
await dispatcher.wait_for_pending()
assert background_endpoint.payloads[0]["texts"] == ["hello"]
assert background_endpoint.payloads[0]["request_data"]["user_api_key_team_id"] == "team-1"
assert background_endpoint.payloads == awaited_endpoint.payloads
async def test_background_post_uses_its_own_timeout():
awaited_endpoint = _Endpoint()
background_endpoint = _Endpoint()
guardrail, dispatcher = _fire_and_forget(background_endpoint)
for target in (_guardrail(awaited_endpoint), guardrail):
await target.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
await dispatcher.wait_for_pending()
assert background_endpoint.read_timeouts == [FIRE_AND_FORGET_POST_TIMEOUT_SECONDS]
assert awaited_endpoint.read_timeouts == [CLIENT_TIMEOUT_SECONDS]
async def test_background_post_does_not_inherit_request_context():
endpoint = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint)
token = _request_scoped.set("request-1")
try:
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
finally:
_request_scoped.reset(token)
await dispatcher.wait_for_pending()
assert endpoint.completed == 1
assert endpoint.seen_request_scoped == [None]
@pytest.mark.parametrize(
"body",
[
{"action": "BLOCKED", "blocked_reason": "nope"},
{"action": "GUARDRAIL_INTERVENED", "texts": ["MASKED"]},
],
)
async def test_verdict_is_ignored(body):
endpoint = _Endpoint(body=body)
guardrail, dispatcher = _fire_and_forget(endpoint)
result = await guardrail.apply_guardrail(inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request")
await dispatcher.wait_for_pending()
assert result == {"texts": ["my ssn is 123"]}
assert endpoint.completed == 1
@pytest.mark.parametrize(
"endpoint_options",
[
{"error": httpx.ConnectError("connection refused")},
{"status_code": 500},
],
)
async def test_failing_endpoint_is_logged_not_raised(endpoint_options, captured_warnings):
endpoint = _Endpoint(**endpoint_options)
guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=True, unreachable_fallback="fail_closed")
result = await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
request_data={},
input_type="response",
logging_obj=_logging_obj(call_id="call-failing"),
)
await dispatcher.wait_for_pending()
assert result == {"texts": ["hello"]}
failures = _messages(captured_warnings, "call failed")
assert len(failures) == 1
assert "ff-guardrail" in failures[0]
assert "input_type=response" in failures[0]
assert "litellm_call_id=call-failing" in failures[0]
@pytest.mark.parametrize("fail_on_error", [True, False])
@pytest.mark.parametrize(
("inputs", "make_request_data"),
[
({"texts": ["hi"], "tools": [{"function": {"name": "f"}}]}, _request_data),
({"texts": ["hi"]}, lambda: {"messages": [], "metadata": None}),
],
ids=["tool_without_type", "malformed_request_metadata"],
)
async def test_failure_before_dispatch_is_logged_and_passes_through(
inputs, make_request_data, fail_on_error, captured_warnings
):
request_data = make_request_data()
endpoint = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, fail_on_error=fail_on_error)
result = await guardrail.apply_guardrail(
inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj("call-bad")
)
await dispatcher.wait_for_pending()
assert result == inputs
assert endpoint.payloads == []
assert _recorded_outcomes(request_data) == [("not_run", FIRE_AND_FORGET_NOT_DISPATCHED_REASON)]
warnings = _messages(captured_warnings, "not dispatched")
assert len(warnings) == 1
assert "litellm_call_id=call-bad" in warnings[0]
async def test_inflight_cap_drops_and_counts_excess_calls(captured_warnings):
endpoint = _Endpoint(gate_open=False)
guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=2)
results = [
await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request")
for i in range(5)
]
assert results == [{"texts": [f"t{i}"]} for i in range(5)]
assert dispatcher.pending_count == 2
assert dispatcher.dropped_count == 3
assert len(_messages(captured_warnings, "dropped")) == 1
endpoint.gate.set()
await dispatcher.wait_for_pending()
assert [p["texts"] for p in endpoint.payloads] == [["t0"], ["t1"]]
def _recorded_outcomes(request_data):
entries = request_data["metadata"]["standard_logging_guardrail_information"]
return [(entry["guardrail_status"], entry["guardrail_response"]) for entry in entries]
async def test_dispatched_and_dropped_calls_are_recorded():
endpoint = _Endpoint(gate_open=False)
guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1)
dispatched, dropped = _request_data(), _request_data()
for request_data in (dispatched, dropped):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request")
endpoint.gate.set()
await dispatcher.wait_for_pending()
assert _recorded_outcomes(dispatched) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)]
assert _recorded_outcomes(dropped) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)]
async def test_finished_task_frees_its_slot():
endpoint = _Endpoint()
guardrail, dispatcher = _fire_and_forget(endpoint, max_inflight=1)
for i in range(3):
await guardrail.apply_guardrail(inputs={"texts": [f"t{i}"]}, request_data={}, input_type="request")
await dispatcher.wait_for_pending()
assert dispatcher.pending_count == 0
assert dispatcher.dropped_count == 0
assert endpoint.completed == 3
def _stream_chunks():
words = ("Hello", " ", "world", "!", " Bye")
return [
ModelResponseStream(
model="gpt-4",
choices=[
litellm.StreamingChoices(
index=0,
delta=Delta(role="assistant", content=word),
finish_reason="stop" if i == len(words) - 1 else None,
)
],
)
for i, word in enumerate(words)
]
async def _run_stream(guardrail):
async def stream():
for chunk in _stream_chunks():
yield chunk
return [
chunk
async for chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/chat/completions"),
response=stream(),
request_data={
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["ff-guardrail"]},
},
)
]
async def test_stream_dispatches_one_call():
per_chunk_endpoint = _Endpoint()
await _run_stream(_guardrail(per_chunk_endpoint, event_hook="post_call", streaming_sampling_rate=1))
background_endpoint = _Endpoint()
guardrail, dispatcher = _fire_and_forget(
background_endpoint, event_hook="post_call", streaming_end_of_stream_only=False, streaming_sampling_rate=1
)
streamed = await _run_stream(guardrail)
await dispatcher.wait_for_pending()
assert len(per_chunk_endpoint.payloads) > 1
assert len(streamed) == len(_stream_chunks())
assert len(background_endpoint.payloads) == 1
assert guardrail.streaming_end_of_stream_only is True
def test_observe_only_warning_only_when_enabled(captured_warnings):
_guardrail(_Endpoint(), name="enforcing")
_guardrail(_Endpoint(), name="observer", fire_and_forget=True)
warnings = _messages(captured_warnings, "observe-only")
assert len(warnings) == 1
assert "observer" in warnings[0]
@pytest.mark.parametrize("value", ["false", "true", 1])
def test_non_bool_fire_and_forget_is_rejected(value):
with pytest.raises(ValueError, match="fire_and_forget must be a bool"):
_guardrail(_Endpoint(), fire_and_forget=value)
@pytest.mark.parametrize("max_inflight", ["5", 2.5, True])
def test_non_int_max_inflight_is_rejected(max_inflight):
with pytest.raises(ValueError, match="fire_and_forget_max_inflight must be an int"):
_guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight)
@pytest.mark.parametrize("max_inflight", [0, -1])
def test_max_inflight_below_one_is_rejected(max_inflight):
with pytest.raises(ValueError, match="fire_and_forget_max_inflight"):
_guardrail(_Endpoint(), fire_and_forget=True, fire_and_forget_max_inflight=max_inflight)
with pytest.raises(pydantic.ValidationError):
GenericGuardrailAPIOptionalParams(fire_and_forget_max_inflight=max_inflight)
async def test_configured_max_inflight_bounds_dispatch():
endpoint = _Endpoint(gate_open=False)
guardrail = _guardrail(endpoint, fire_and_forget=True, fire_and_forget_max_inflight=1)
first, second = _request_data(), _request_data()
for request_data in (first, second):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request")
endpoint.gate.set()
await guardrail._dispatcher.wait_for_pending()
assert _recorded_outcomes(first) == [("success", FIRE_AND_FORGET_DISPATCHED_REASON)]
assert _recorded_outcomes(second) == [("not_run", FIRE_AND_FORGET_DROPPED_REASON)]
assert len(endpoint.payloads) == 1
async def test_default_max_inflight_admits_concurrent_calls():
endpoint = _Endpoint(gate_open=False)
guardrail = _guardrail(endpoint, fire_and_forget=True)
calls = [_request_data() for _ in range(DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + 1)]
for request_data in calls:
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type="request")
endpoint.gate.set()
await guardrail._dispatcher.wait_for_pending()
outcomes = [_recorded_outcomes(request_data)[0][0] for request_data in calls]
assert outcomes == ["success"] * DEFAULT_FIRE_AND_FORGET_MAX_INFLIGHT + ["not_run"]
async def test_initialize_guardrail_forwards_fire_and_forget():
litellm_params = LitellmParams(
guardrail="generic_guardrail_api", mode="pre_call", api_base=API_BASE, default_on=True
)
litellm_params.fire_and_forget = True
litellm_params.fire_and_forget_max_inflight = 1
gate = asyncio.Event()
guardrail = initialize_guardrail(litellm_params, {"guardrail_name": "from-config"})
try:
assert guardrail.fire_and_forget is True
assert guardrail.streaming_end_of_stream_only is True
assert guardrail._dispatcher.dispatch(gate.wait, context="first") is True
assert guardrail._dispatcher.dispatch(gate.wait, context="second") is False
finally:
gate.set()
await guardrail._dispatcher.wait_for_pending()
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
async def test_default_awaits_the_endpoint_and_blocks():
endpoint = _Endpoint(body={"action": "BLOCKED", "blocked_reason": "nope"})
guardrail = _guardrail(endpoint)
assert guardrail.fire_and_forget is False
assert guardrail.streaming_end_of_stream_only is False
with pytest.raises(GuardrailRaisedException, match="nope"):
await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={}, input_type="request")
assert endpoint.completed == 1