mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge de597a5228 into be4481779e
This commit is contained in:
commit
0412812987
12 changed files with 1083 additions and 10 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
24
tests/unit/proxy/guardrails/guardrail_hooks/conftest.py
Normal file
24
tests/unit/proxy/guardrails/guardrail_hooks/conftest.py
Normal 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)
|
||||
|
|
@ -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]
|
||||
|
|
@ -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
|
||||
|
|
@ -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")) == ([], [], [])
|
||||
Loading…
Add table
Reference in a new issue