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

View file

@ -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)

View file

@ -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)

View file

@ -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 (

View file

@ -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

View file

@ -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],

View file

@ -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."""

View file

@ -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),
)