mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
73904c9f94
commit
f7e934677b
7 changed files with 619 additions and 468 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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})")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue