mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 800b76c1b4 into b781d157d7
This commit is contained in:
commit
23d81423cb
8 changed files with 654 additions and 4 deletions
|
|
@ -70,6 +70,17 @@ DEFAULT_ADVISORY_MESSAGE: Final = (
|
|||
_guardrail_self_recorded: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar(
|
||||
"litellm_guardrail_self_recorded", default=False
|
||||
)
|
||||
_guardrail_success_record_skipped: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar(
|
||||
"litellm_guardrail_success_record_skipped", default=False
|
||||
)
|
||||
|
||||
|
||||
def skip_guardrail_success_record() -> None:
|
||||
"""Skip the entry ``log_guardrail_information`` would record for this call's normal return.
|
||||
|
||||
Only the success branch reads this flag, so a later raise still records its error entry.
|
||||
"""
|
||||
_guardrail_success_record_skipped.set(True)
|
||||
|
||||
|
||||
def is_guardrail_intervention(e: Exception) -> bool:
|
||||
|
|
@ -1630,9 +1641,14 @@ def log_guardrail_information(func):
|
|||
|
||||
logging_obj: Final = kwargs.get("logging_obj") or request_data.get("litellm_logging_obj")
|
||||
self_recorded_token: Final = _guardrail_self_recorded.set(False)
|
||||
success_record_skipped_token: Final = _guardrail_success_record_skipped.set(False)
|
||||
try:
|
||||
response: Final = await func(*args, **kwargs)
|
||||
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
|
||||
if (
|
||||
self.records_own_guardrail_information
|
||||
or _guardrail_self_recorded.get()
|
||||
or _guardrail_success_record_skipped.get()
|
||||
):
|
||||
return response
|
||||
return self._process_response(
|
||||
response=response,
|
||||
|
|
@ -1656,6 +1672,7 @@ def log_guardrail_information(func):
|
|||
)
|
||||
finally:
|
||||
_guardrail_self_recorded.reset(self_recorded_token)
|
||||
_guardrail_success_record_skipped.reset(success_record_skipped_token)
|
||||
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
|
||||
|
||||
@functools.wraps(func)
|
||||
|
|
@ -1670,9 +1687,14 @@ def log_guardrail_information(func):
|
|||
|
||||
logging_obj: Final = kwargs.get("logging_obj") or request_data.get("litellm_logging_obj")
|
||||
self_recorded_token: Final = _guardrail_self_recorded.set(False)
|
||||
success_record_skipped_token: Final = _guardrail_success_record_skipped.set(False)
|
||||
try:
|
||||
response: Final = func(*args, **kwargs)
|
||||
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
|
||||
if (
|
||||
self.records_own_guardrail_information
|
||||
or _guardrail_self_recorded.get()
|
||||
or _guardrail_success_record_skipped.get()
|
||||
):
|
||||
return response
|
||||
return self._process_response(
|
||||
response=response,
|
||||
|
|
@ -1692,6 +1714,7 @@ def log_guardrail_information(func):
|
|||
)
|
||||
finally:
|
||||
_guardrail_self_recorded.reset(self_recorded_token)
|
||||
_guardrail_success_record_skipped.reset(success_record_skipped_token)
|
||||
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
|
||||
|
||||
@functools.wraps(func)
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
|
||||
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
|
||||
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
|
||||
guardrail_information_scope=_get_config_value(litellm_params, optional_params, "guardrail_information_scope"),
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)
|
||||
|
|
|
|||
|
|
@ -17,9 +17,12 @@ from litellm._version import version as litellm_version
|
|||
from litellm.exceptions import GuardrailRaisedException, Timeout
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
get_session_id_from_request_data,
|
||||
log_guardrail_information,
|
||||
skip_guardrail_success_record,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
|
|
@ -29,10 +32,13 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
|
|||
GenericGuardrailAPIMetadata,
|
||||
GenericGuardrailAPIRequest,
|
||||
GenericGuardrailAPIResponse,
|
||||
GuardrailInformationScope,
|
||||
GuardrailToolParam,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
from .record_scope import DEFAULT_GUARDRAIL_INFORMATION_SCOPE, RecordScope, returned_unchanged
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
|
@ -204,9 +210,13 @@ 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,
|
||||
guardrail_information_scope: GuardrailInformationScope | None = None,
|
||||
async_handler: AsyncHTTPHandler | 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 []
|
||||
|
||||
|
|
@ -251,6 +261,10 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
"block_only" if streaming_transform_mode is None else streaming_transform_mode
|
||||
)
|
||||
|
||||
self._record_scope: Final = RecordScope(
|
||||
DEFAULT_GUARDRAIL_INFORMATION_SCOPE if guardrail_information_scope is None else guardrail_information_scope
|
||||
)
|
||||
|
||||
# Set supported event hooks
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
|
||||
|
|
@ -498,7 +512,7 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
blocked_content=True,
|
||||
)
|
||||
|
||||
return self._build_guardrail_return_inputs(
|
||||
return_inputs: Final = self._build_guardrail_return_inputs(
|
||||
texts=texts,
|
||||
images=images,
|
||||
tools=tools,
|
||||
|
|
@ -522,6 +536,15 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
except Exception as e:
|
||||
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj, is_unreachable=False)
|
||||
|
||||
unchanged_allow: Final = guardrail_response.action == "NONE" and returned_unchanged(inputs, return_inputs)
|
||||
if unchanged_allow and not self._record_scope.should_record_allow(
|
||||
session_id=get_session_id_from_request_data(request_data),
|
||||
tenant=user_metadata.get("user_api_key_hash") or user_metadata.get("user_api_key_team_id"),
|
||||
input_type=input_type,
|
||||
):
|
||||
skip_guardrail_success_record()
|
||||
return return_inputs
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type["GuardrailConfigModel"] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
|
|
|
|||
|
|
@ -0,0 +1,56 @@
|
|||
import json
|
||||
from typing import Final, Literal
|
||||
|
||||
from pydantic_core import to_jsonable_python
|
||||
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GuardrailInformationScope
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
DEFAULT_GUARDRAIL_INFORMATION_SCOPE: Final[GuardrailInformationScope] = "per_call"
|
||||
|
||||
_SESSION_CACHE_MAX_ENTRIES: Final = 100_000
|
||||
_SESSION_CACHE_TTL_SECONDS: Final = 3600
|
||||
_REWRITABLE_KEYS: Final = ("texts", "images", "tools", "structured_messages")
|
||||
|
||||
|
||||
def _jsonable(value: object) -> object:
|
||||
return to_jsonable_python(value, fallback=repr, bytes_mode="base64")
|
||||
|
||||
|
||||
def returned_unchanged(sent: GenericGuardrailAPIInputs, returned: GenericGuardrailAPIInputs) -> bool:
|
||||
"""The return builder always sets texts and sets other rewritable keys only when passing them through or
|
||||
rewriting them, so a key missing from ``returned`` is unchanged and missing texts were sent as ``[]``."""
|
||||
return all(
|
||||
key not in returned or _jsonable(returned.get(key)) == _jsonable(sent.get(key, [])) for key in _REWRITABLE_KEYS
|
||||
)
|
||||
|
||||
|
||||
class RecordScope:
|
||||
def __init__(self, scope: GuardrailInformationScope, *, recorded_sessions: InMemoryCache | None = None) -> None:
|
||||
self._scope: Final = scope
|
||||
self._recorded_sessions: Final = recorded_sessions or InMemoryCache(
|
||||
max_size_in_memory=_SESSION_CACHE_MAX_ENTRIES,
|
||||
default_ttl=_SESSION_CACHE_TTL_SECONDS,
|
||||
)
|
||||
|
||||
def should_record_allow(
|
||||
self,
|
||||
*,
|
||||
session_id: str | None,
|
||||
tenant: str | None,
|
||||
input_type: Literal["request", "response"],
|
||||
) -> bool:
|
||||
match self._scope:
|
||||
case "per_call":
|
||||
return True
|
||||
case "off":
|
||||
return False
|
||||
case "per_session":
|
||||
return session_id is None or self._claim_session(json.dumps([tenant, session_id, input_type]))
|
||||
|
||||
def _claim_session(self, key: str) -> bool:
|
||||
if self._recorded_sessions.get_cache(key) is True:
|
||||
return False
|
||||
self._recorded_sessions.set_cache(key, True)
|
||||
return True
|
||||
|
|
@ -11,6 +11,8 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall
|
||||
|
||||
GuardrailInformationScope = Literal["per_call", "per_session", "off"]
|
||||
|
||||
|
||||
class GuardrailToolParam(BaseModel):
|
||||
"""A tool forwarded verbatim to the guardrail for inspection.
|
||||
|
|
@ -103,6 +105,23 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
|
|||
),
|
||||
)
|
||||
|
||||
guardrail_information_scope: GuardrailInformationScope | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"How often a call that allows the content unchanged records its guardrail information entry (spend "
|
||||
"logs, OTEL, logging callbacks). 'per_call' (default) records every call, so with pre_call and "
|
||||
"post_call a session grows by a request and a response entry per turn, plus one per sampled "
|
||||
"streaming check. 'per_session' records the first unchanged allow of each session once on the "
|
||||
"request side and once on the response side. Sessions are keyed by the authenticated key hash, "
|
||||
"falling back to the team id (callers with neither share one namespace), and the session id "
|
||||
"(litellm_session_id or metadata.session_id). A call without a session id records as 'per_call'. "
|
||||
"Seen sessions are remembered per proxy process for an hour. 'off' records no unchanged allows. "
|
||||
"Blocks, rewrites (GUARDRAIL_INTERVENED, or returned content that differs from what was sent), "
|
||||
"errors, fail-open passthroughs and not_run entries are recorded under every scope and never count "
|
||||
"as a session's first call. Defaults to 'per_call' when None."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class GenericGuardrailAPIConfigModel(
|
||||
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.integrations.custom_guardrail import (
|
|||
DEFAULT_ADVISORY_MESSAGE,
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
skip_guardrail_success_record,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
|
||||
|
|
@ -2207,6 +2208,100 @@ class TestRecordsOwnGuardrailInformation:
|
|||
assert _guardrail_entries(request_data) == []
|
||||
|
||||
|
||||
def _skip_success_record_then(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs:
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
||||
texts: Final = inputs.get("texts") or []
|
||||
if "skip" in texts:
|
||||
skip_guardrail_success_record()
|
||||
if "raise" in texts:
|
||||
raise GuardrailRaisedException(guardrail_name="skipper", message="blocked")
|
||||
return inputs
|
||||
|
||||
|
||||
class _SuccessRecordSkippingGuardrail(CustomGuardrail):
|
||||
@log_guardrail_information
|
||||
async def check_async(self, inputs: GenericGuardrailAPIInputs, request_data: dict) -> GenericGuardrailAPIInputs:
|
||||
return _skip_success_record_then(inputs)
|
||||
|
||||
@log_guardrail_information
|
||||
def check_sync(self, inputs: GenericGuardrailAPIInputs, request_data: dict) -> GenericGuardrailAPIInputs:
|
||||
return _skip_success_record_then(inputs)
|
||||
|
||||
@log_guardrail_information
|
||||
async def outer_check_async(
|
||||
self, inputs: GenericGuardrailAPIInputs, request_data: dict
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return await self.check_async(inputs=inputs, request_data=request_data)
|
||||
|
||||
@log_guardrail_information
|
||||
def outer_check_sync(self, inputs: GenericGuardrailAPIInputs, request_data: dict) -> GenericGuardrailAPIInputs:
|
||||
return self.check_sync(inputs=inputs, request_data=request_data)
|
||||
|
||||
|
||||
async def _recorded_after_check(branch: Literal["async", "sync"], texts: list[str]) -> list[str]:
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
||||
guardrail: Final = _SuccessRecordSkippingGuardrail(guardrail_name="skipper")
|
||||
request_data: Final[dict] = {"metadata": {}}
|
||||
inputs: Final = GenericGuardrailAPIInputs(texts=texts)
|
||||
try:
|
||||
if branch == "async":
|
||||
await guardrail.check_async(inputs=inputs, request_data=request_data)
|
||||
else:
|
||||
guardrail.check_sync(inputs=inputs, request_data=request_data)
|
||||
except GuardrailRaisedException:
|
||||
pass
|
||||
return [entry["guardrail_status"] for entry in _guardrail_entries(request_data)]
|
||||
|
||||
|
||||
class TestSkipGuardrailSuccessRecord:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("branch", ["async", "sync"])
|
||||
@pytest.mark.parametrize(
|
||||
("texts", "expected"),
|
||||
[
|
||||
(["ok"], ["success"]),
|
||||
(["skip"], []),
|
||||
(["skip", "raise"], ["guardrail_intervened"]),
|
||||
],
|
||||
)
|
||||
async def test_skips_only_the_success_entry(
|
||||
self, branch: Literal["async", "sync"], texts: list[str], expected: list[str]
|
||||
) -> None:
|
||||
assert await _recorded_after_check(branch, texts) == expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("branch", ["async", "sync"])
|
||||
async def test_skip_does_not_leak_into_the_next_call(self, branch: Literal["async", "sync"]) -> None:
|
||||
skipped: Final = await _recorded_after_check(branch, ["skip"])
|
||||
|
||||
assert (skipped, await _recorded_after_check(branch, ["ok"])) == ([], ["success"])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("branch", ["async", "sync"])
|
||||
async def test_skip_outside_a_guardrail_call_does_not_leak_into_it(self, branch: Literal["async", "sync"]) -> None:
|
||||
async def skip_then_check() -> list[str]:
|
||||
skip_guardrail_success_record()
|
||||
return await _recorded_after_check(branch, ["ok"])
|
||||
|
||||
assert await asyncio.create_task(skip_then_check()) == ["success"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("branch", ["async", "sync"])
|
||||
async def test_inner_skip_does_not_skip_the_outer_entry(self, branch: Literal["async", "sync"]) -> None:
|
||||
guardrail: Final = _SuccessRecordSkippingGuardrail(guardrail_name="skipper")
|
||||
request_data: Final[dict] = {"metadata": {}}
|
||||
inputs: Final = GenericGuardrailAPIInputs(texts=["skip"])
|
||||
|
||||
if branch == "async":
|
||||
await guardrail.outer_check_async(inputs=inputs, request_data=request_data)
|
||||
else:
|
||||
guardrail.outer_check_sync(inputs=inputs, request_data=request_data)
|
||||
|
||||
assert [entry["guardrail_status"] for entry in _guardrail_entries(request_data)] == ["success"]
|
||||
|
||||
|
||||
class _UndecoratedGuardrail(CustomGuardrail):
|
||||
"""apply_guardrail written like the docs example: no @log_guardrail_information."""
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,433 @@
|
|||
from collections.abc import Callable, Iterator, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.record_scope import (
|
||||
RecordScope,
|
||||
returned_unchanged,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GuardrailInformationScope,
|
||||
GuardrailToolParam,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
Responder = Callable[[httpx.Request], httpx.Response]
|
||||
|
||||
_KEY_A: Final = MappingProxyType({"user_api_key_hash": "hash-a"})
|
||||
_USER_ROW: Final = MappingProxyType({"role": "user", "content": "hello"})
|
||||
_TOOL: Final = MappingProxyType(
|
||||
{"type": "function", "function": {"name": "lookup", "parameters": {"type": "object", "required": ["q"]}}}
|
||||
)
|
||||
|
||||
|
||||
def _chat_request(**extra: object) -> GenericGuardrailAPIInputs:
|
||||
return GenericGuardrailAPIInputs(texts=["hello"], structured_messages=[dict(_USER_ROW)], model="gpt-test", **extra)
|
||||
|
||||
|
||||
def _chat_response() -> GenericGuardrailAPIInputs:
|
||||
return GenericGuardrailAPIInputs(texts=["hi there"], model="gpt-test")
|
||||
|
||||
|
||||
def _respond_with(body: Mapping[str, object]) -> Responder:
|
||||
return lambda request: httpx.Response(200, json=dict(body), request=request)
|
||||
|
||||
|
||||
_allow: Final = _respond_with({"action": "NONE"})
|
||||
_echo_texts: Final = _respond_with({"action": "NONE", "texts": ["hello"]})
|
||||
_echo_rows: Final = _respond_with({"action": "NONE", "structured_messages": [dict(_USER_ROW)]})
|
||||
_block: Final = _respond_with({"action": "BLOCKED", "blocked_reason": "policy"})
|
||||
_mask: Final = _respond_with({"action": "GUARDRAIL_INTERVENED", "texts": ["[REDACTED]"]})
|
||||
_rewrite_texts: Final = _respond_with({"action": "NONE", "texts": ["[REDACTED]"]})
|
||||
_rewrite_rows: Final = _respond_with({"action": "NONE", "structured_messages": [{"role": "user", "content": "[x]"}]})
|
||||
_intervene_without_rewrite: Final = _respond_with({"action": "GUARDRAIL_INTERVENED"})
|
||||
|
||||
|
||||
def _unreachable(request: httpx.Request) -> httpx.Response:
|
||||
raise httpx.ConnectError("connection refused", request=request)
|
||||
|
||||
|
||||
def _unavailable(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(503, request=request)
|
||||
|
||||
|
||||
def _in_order(*responders: Responder) -> Responder:
|
||||
remaining: Final[Iterator[Responder]] = iter(responders)
|
||||
return lambda request: next(remaining)(request)
|
||||
|
||||
|
||||
class _Endpoint:
|
||||
def __init__(self, respond: Responder) -> None:
|
||||
self.requests: Final[list[httpx.Request]] = [] # mutable-ok: records what the endpoint received
|
||||
self._respond: Final = respond
|
||||
|
||||
def __call__(self, request: httpx.Request) -> httpx.Response:
|
||||
self.requests.append(request)
|
||||
return self._respond(request)
|
||||
|
||||
|
||||
def _guardrail(endpoint: _Endpoint, **options: object) -> GenericGuardrailAPI:
|
||||
return GenericGuardrailAPI(
|
||||
api_base="https://guardrail.test",
|
||||
guardrail_name="scoped-guardrail",
|
||||
event_hook="pre_call",
|
||||
async_handler=AsyncHTTPHandler(transport=httpx.MockTransport(endpoint)),
|
||||
**options,
|
||||
)
|
||||
|
||||
|
||||
def _turn(
|
||||
session_id: str | None = "session-1",
|
||||
metadata: Mapping[str, str] = _KEY_A,
|
||||
) -> dict[str, object]: # mutable-ok: the decorator appends entries into the request data
|
||||
base: Final = {"metadata": dict(metadata)}
|
||||
return base if session_id is None else {**base, "litellm_session_id": session_id}
|
||||
|
||||
|
||||
def _recorded(request_data: Mapping[str, object]) -> list[str]:
|
||||
metadata: Final = request_data.get("metadata")
|
||||
assert isinstance(metadata, dict)
|
||||
return [entry["guardrail_status"] for entry in metadata.get("standard_logging_guardrail_information") or []]
|
||||
|
||||
|
||||
async def _run(
|
||||
guardrail: GenericGuardrailAPI,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"] = "request",
|
||||
inputs: GenericGuardrailAPIInputs | None = None,
|
||||
) -> list[str]:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=_chat_request() if inputs is None else inputs, request_data=request_data, input_type=input_type
|
||||
)
|
||||
return _recorded(request_data)
|
||||
|
||||
|
||||
async def _run_blocked(guardrail: GenericGuardrailAPI, request_data: dict[str, object]) -> list[str]:
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail.apply_guardrail(inputs=_chat_request(), request_data=request_data, input_type="request")
|
||||
return _recorded(request_data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"options", [{}, {"guardrail_information_scope": None}, {"guardrail_information_scope": "per_call"}]
|
||||
)
|
||||
async def test_per_call_records_every_call_of_a_session(options: dict[str, object]) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), **options)
|
||||
|
||||
assert [await _run(guardrail, _turn()) for _ in range(3)] == [["success"]] * 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_records_the_first_call_only_but_still_runs_every_call() -> None:
|
||||
endpoint: Final = _Endpoint(_allow)
|
||||
guardrail: Final = _guardrail(endpoint, guardrail_information_scope="per_session")
|
||||
|
||||
assert [await _run(guardrail, _turn()) for _ in range(3)] == [["success"], [], []]
|
||||
assert len(endpoint.requests) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scope", ["per_session", "off"])
|
||||
@pytest.mark.parametrize(
|
||||
("respond", "inputs"),
|
||||
[
|
||||
(_allow, _chat_request()),
|
||||
(_allow, _chat_request(images=[])),
|
||||
(_echo_texts, _chat_request()),
|
||||
(_echo_rows, _chat_request()),
|
||||
(_allow, _chat_request(tools=[dict(_TOOL)])),
|
||||
(_respond_with({"action": "NONE", "tools": [dict(_TOOL)]}), _chat_request(tools=[dict(_TOOL)])),
|
||||
],
|
||||
ids=["allow", "empty-images", "echoed-texts", "echoed-rows", "passed-through-tools", "echoed-tools"],
|
||||
)
|
||||
async def test_unchanged_chat_allows_are_deduped_after_the_first_turn(
|
||||
scope: GuardrailInformationScope, respond: Responder, inputs: GenericGuardrailAPIInputs
|
||||
) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(respond), guardrail_information_scope=scope)
|
||||
|
||||
turns: Final = [await _run(guardrail, _turn(), inputs=inputs) for _ in range(2)]
|
||||
|
||||
assert turns[1] == []
|
||||
assert turns[0] == (["success"] if scope == "per_session" else [])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("scope", "expected"), [("per_session", [["success"], []]), ("off", [[], []])])
|
||||
async def test_tool_use_only_responses_without_texts_are_deduped(
|
||||
scope: GuardrailInformationScope, expected: list[list[str]]
|
||||
) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope=scope)
|
||||
tool_use_only: Final = GenericGuardrailAPIInputs(
|
||||
tool_calls=[{"id": "call-1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}],
|
||||
model="gpt-test",
|
||||
)
|
||||
|
||||
assert [await _run(guardrail, _turn(), "response", tool_use_only) for _ in range(2)] == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("scope", "expected"), [("per_call", [1, 1]), ("per_session", [1, 0]), ("off", [0, 0])])
|
||||
async def test_chat_completion_turns_through_the_openai_handler_follow_the_scope(
|
||||
scope: GuardrailInformationScope, expected: list[int]
|
||||
) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope=scope)
|
||||
|
||||
async def chat_turn() -> int:
|
||||
data: Final = {
|
||||
"model": "gpt-test",
|
||||
"messages": [{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}],
|
||||
"tools": [dict(_TOOL)],
|
||||
"metadata": dict(_KEY_A),
|
||||
"litellm_session_id": "chat-session",
|
||||
}
|
||||
await OpenAIChatCompletionsHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
return len(_recorded(data))
|
||||
|
||||
assert [await chat_turn() for _ in range(2)] == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_records_the_first_call_of_each_session() -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
assert [await _run(guardrail, _turn(session_id)) for session_id in ("s1", "s1", "s2", "s2")] == [
|
||||
["success"],
|
||||
[],
|
||||
["success"],
|
||||
[],
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_records_the_first_request_and_the_first_response_of_a_session() -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
async def run_turn() -> list[str]:
|
||||
request_data: Final = _turn()
|
||||
await _run(guardrail, request_data, "request")
|
||||
return await _run(guardrail, request_data, "response", _chat_response())
|
||||
|
||||
first_turn: Final = await run_turn()
|
||||
second_turn: Final = await run_turn()
|
||||
|
||||
assert (first_turn, second_turn) == (["success", "success"], [])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("identity_field", ["user_api_key_hash", "user_api_key_team_id"])
|
||||
async def test_per_session_dedups_per_authenticated_caller(identity_field: str) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
statuses: Final = [
|
||||
await _run(guardrail, _turn("shared-id", {identity_field: tenant}))
|
||||
for tenant in ("tenant-a", "tenant-b", "tenant-a", "tenant-b")
|
||||
]
|
||||
|
||||
assert statuses == [["success"], ["success"], [], []]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_key_hash_takes_precedence_over_team_id() -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
same_team: Final = (
|
||||
{"user_api_key_hash": "hash-a", "user_api_key_team_id": "team"},
|
||||
{"user_api_key_hash": "hash-b", "user_api_key_team_id": "team"},
|
||||
)
|
||||
|
||||
assert [await _run(guardrail, _turn("shared-id", identity)) for identity in same_team] == [["success"]] * 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_dedups_callers_without_a_key_or_team_apart_from_authenticated_ones() -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
statuses: Final = [await _run(guardrail, _turn("shared-id", identity)) for identity in ({}, {}, _KEY_A, _KEY_A)]
|
||||
|
||||
assert statuses == [["success"], [], ["success"], []]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_without_a_session_id_records_every_call() -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
assert [await _run(guardrail, _turn(session_id=None)) for _ in range(3)] == [["success"]] * 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_session_reads_the_session_id_from_metadata() -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
statuses: Final = [
|
||||
await _run(guardrail, _turn(None, {"user_api_key_hash": "hash-a", "session_id": "meta-session"}))
|
||||
for _ in range(2)
|
||||
]
|
||||
|
||||
assert statuses == [["success"], []]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_off_still_runs_every_call() -> None:
|
||||
endpoint: Final = _Endpoint(_allow)
|
||||
guardrail: Final = _guardrail(endpoint, guardrail_information_scope="off")
|
||||
|
||||
assert [await _run(guardrail, _turn(session_id)) for session_id in ("s1", "s1", None)] == [[], [], []]
|
||||
assert len(endpoint.requests) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scope", ["per_session", "off"])
|
||||
@pytest.mark.parametrize(
|
||||
("respond", "inputs"),
|
||||
[
|
||||
(_mask, _chat_request()),
|
||||
(_rewrite_texts, _chat_request()),
|
||||
(_rewrite_rows, _chat_request()),
|
||||
(_intervene_without_rewrite, _chat_request()),
|
||||
(_respond_with({"action": "NONE", "images": ["data:image/png;base64,Yg=="]}), _chat_request(images=["a"])),
|
||||
(
|
||||
_respond_with({"action": "NONE", "tools": [{"type": "function", "function": {"name": "other"}}]}),
|
||||
_chat_request(tools=[dict(_TOOL)]),
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"mask",
|
||||
"rewritten-texts",
|
||||
"rewritten-rows",
|
||||
"intervention-without-rewrite",
|
||||
"rewritten-images",
|
||||
"rewritten-tools",
|
||||
],
|
||||
)
|
||||
async def test_rewrites_and_interventions_are_recorded_under_every_scope(
|
||||
scope: GuardrailInformationScope, respond: Responder, inputs: GenericGuardrailAPIInputs
|
||||
) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(respond), guardrail_information_scope=scope)
|
||||
|
||||
assert [await _run(guardrail, _turn(), inputs=inputs) for _ in range(2)] == [["success"]] * 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scope", ["per_session", "off"])
|
||||
async def test_blocks_are_recorded_under_every_scope(scope: GuardrailInformationScope) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_block), guardrail_information_scope=scope)
|
||||
|
||||
assert [await _run_blocked(guardrail, _turn()) for _ in range(2)] == [["guardrail_intervened"]] * 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scope", ["per_session", "off"])
|
||||
async def test_endpoint_failures_are_recorded_under_every_scope(scope: GuardrailInformationScope) -> None:
|
||||
guardrail: Final = _guardrail(_Endpoint(_unreachable), guardrail_information_scope=scope)
|
||||
request_data: Final = _turn()
|
||||
|
||||
with pytest.raises(Exception, match="Generic Guardrail API failed"):
|
||||
await guardrail.apply_guardrail(inputs=_chat_request(), request_data=request_data, input_type="request")
|
||||
|
||||
assert _recorded(request_data) == ["guardrail_failed_to_respond"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scope", ["per_session", "off"])
|
||||
@pytest.mark.parametrize("respond", [_unreachable, _unavailable])
|
||||
async def test_fail_open_passthroughs_are_recorded_under_every_scope(
|
||||
scope: GuardrailInformationScope, respond: Responder
|
||||
) -> None:
|
||||
guardrail: Final = _guardrail(
|
||||
_Endpoint(respond), guardrail_information_scope=scope, unreachable_fallback="fail_open"
|
||||
)
|
||||
|
||||
assert [await _run(guardrail, _turn()) for _ in range(2)] == [["success"]] * 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_block_does_not_claim_the_session() -> None:
|
||||
guardrail: Final = _guardrail(
|
||||
_Endpoint(_in_order(_block, _allow, _allow)), guardrail_information_scope="per_session"
|
||||
)
|
||||
|
||||
blocked: Final = await _run_blocked(guardrail, _turn())
|
||||
allowed: Final = [await _run(guardrail, _turn()) for _ in range(2)]
|
||||
|
||||
assert (blocked, allowed) == (["guardrail_intervened"], [["success"], []])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("first", [_mask, _intervene_without_rewrite])
|
||||
async def test_a_rewrite_or_intervention_does_not_claim_the_session(first: Responder) -> None:
|
||||
guardrail: Final = _guardrail(
|
||||
_Endpoint(_in_order(first, _allow, _allow)), guardrail_information_scope="per_session"
|
||||
)
|
||||
|
||||
assert [await _run(guardrail, _turn()) for _ in range(3)] == [["success"], ["success"], []]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_fail_open_passthrough_does_not_claim_the_session() -> None:
|
||||
guardrail: Final = _guardrail(
|
||||
_Endpoint(_in_order(_unavailable, _allow, _allow)),
|
||||
guardrail_information_scope="per_session",
|
||||
unreachable_fallback="fail_open",
|
||||
)
|
||||
|
||||
assert [await _run(guardrail, _turn()) for _ in range(3)] == [["success"], ["success"], []]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_each_guardrail_instance_records_its_own_first_call_of_a_session() -> None:
|
||||
first: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
second: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session")
|
||||
|
||||
statuses: Final = [await _run(guardrail, _turn("shared-id")) for guardrail in (first, second, first, second)]
|
||||
|
||||
assert statuses == [["success"], ["success"], [], []]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("returned_tools", "unchanged"),
|
||||
[
|
||||
([GuardrailToolParam.model_validate(dict(_TOOL))], True),
|
||||
([GuardrailToolParam.model_validate({"type": "function", "function": {"name": "other"}})], False),
|
||||
],
|
||||
)
|
||||
def test_tools_returned_as_models_compare_by_content(returned_tools: list[GuardrailToolParam], unchanged: bool) -> None:
|
||||
tool_with_a_tuple: Final = {
|
||||
"type": "function",
|
||||
"function": {"name": "lookup", "parameters": {"type": "object", "required": ("q",)}},
|
||||
}
|
||||
sent: Final = _chat_request(tools=[tool_with_a_tuple])
|
||||
returned: Final = GenericGuardrailAPIInputs(texts=["hello"], tools=returned_tools)
|
||||
|
||||
assert returned_unchanged(sent, returned) is unchanged
|
||||
|
||||
|
||||
def test_per_session_records_again_once_the_session_has_expired() -> None:
|
||||
now: Final = [0.0] # mutable-ok: fake clock advanced by the test
|
||||
record_scope: Final = RecordScope(
|
||||
"per_session", recorded_sessions=InMemoryCache(default_ttl=60, clock=lambda: now[0])
|
||||
)
|
||||
|
||||
def record() -> bool:
|
||||
return record_scope.should_record_allow(session_id="s1", tenant="hash-a", input_type="request")
|
||||
|
||||
first, within_ttl = record(), record()
|
||||
now[0] = 61.0
|
||||
|
||||
assert (first, within_ttl, record()) == (True, False, True)
|
||||
|
||||
|
||||
def test_undecodable_bytes_compare_without_raising() -> None:
|
||||
rows: Final = [{"role": "user", "content": b"\xff\xfe"}]
|
||||
|
||||
assert returned_unchanged(
|
||||
GenericGuardrailAPIInputs(texts=["hello"], structured_messages=rows),
|
||||
GenericGuardrailAPIInputs(texts=["hello"], structured_messages=rows),
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue