diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 67e2173ecd0..72a5fb3c19f 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -75,6 +75,14 @@ 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: + """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: @@ -1158,7 +1166,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 @@ -1644,9 +1652,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, @@ -1670,6 +1683,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) @@ -1684,9 +1698,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, @@ -1706,6 +1725,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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index de389d8a945..0611ce9f6f2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -40,6 +40,7 @@ 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, + guardrail_information_scope=_get_config_value(litellm_params, optional_params, "guardrail_information_scope"), ) litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 7bb41b7586b..af0dcde523d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -18,17 +18,27 @@ from litellm.exceptions import GuardrailRaisedException, Timeout from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, + skip_guardrail_success_record, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.record_scope import ( + Caller, + RecordScope, + guardrail_information_scope_from_config, + returned_unchanged, + session_id_of, +) 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 ( GenericGuardrailAPIMetadata, GenericGuardrailAPIRequest, GenericGuardrailAPIResponse, + GuardrailInformationScope, GuardrailToolParam, ) from litellm.types.utils import GenericGuardrailAPIInputs @@ -204,9 +214,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 +265,8 @@ class GenericGuardrailAPI(CustomGuardrail): "block_only" if streaming_transform_mode is None else streaming_transform_mode ) + self._record_scope: Final = RecordScope(guardrail_information_scope_from_config(guardrail_information_scope)) + # Set supported event hooks kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) @@ -339,7 +355,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, @@ -499,7 +515,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, @@ -523,6 +539,23 @@ class GenericGuardrailAPI(CustomGuardrail): except Exception as e: return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj, is_unreachable=False) + if ( + guardrail_response.action == "NONE" + and not self._record_scope.records_every_allow + and returned_unchanged(inputs, return_inputs) + and not self._record_scope.should_record_allow( + session_id=session_id_of(request_data), + caller=Caller( + key_hash=user_metadata.get("user_api_key_hash"), + team_id=user_metadata.get("user_api_key_team_id"), + user_id=user_metadata.get("user_api_key_user_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 ( diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py new file mode 100644 index 00000000000..f51f8cf832c --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -0,0 +1,102 @@ +import hashlib +import json +from typing import Final, Literal, NamedTuple + +from pydantic import ConfigDict, TypeAdapter, ValidationError +from pydantic_core import to_jsonable_python +from typing_extensions import assert_never + +from litellm._logging import verbose_proxy_logger +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.integrations.custom_guardrail import get_session_id_from_request_data +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") +_REQUEST_DATA_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) +_SCOPE_ADAPTER: Final[TypeAdapter[GuardrailInformationScope]] = TypeAdapter( + GuardrailInformationScope, config=ConfigDict(title="guardrail_information_scope") +) + + +class Caller(NamedTuple): + key_hash: str | None + team_id: str | None + user_id: str | None + + +def session_id_of(request_data: object) -> str | None: + return get_session_id_from_request_data(_REQUEST_DATA_ADAPTER.validate_python(request_data)) + + +def guardrail_information_scope_from_config(value: object) -> GuardrailInformationScope: + if value is None: + return DEFAULT_GUARDRAIL_INFORMATION_SCOPE + try: + return _SCOPE_ADAPTER.validate_python(value) + except ValidationError: + verbose_proxy_logger.warning( + "Ignoring guardrail_information_scope=%r, expected per_call, per_session or off. Recording every call", + value, + ) + return DEFAULT_GUARDRAIL_INFORMATION_SCOPE + + +def _jsonable(value: object) -> object: + return to_jsonable_python(value, fallback=repr, bytes_mode="base64") + + +def _session_key(caller: Caller, session_id: str, input_type: Literal["request", "response"]) -> str: + return hashlib.sha256(json.dumps((*caller, session_id, input_type)).encode()).hexdigest() + + +def _sent_value(sent: GenericGuardrailAPIInputs, key: str) -> object: + return sent.get(key) if key in sent else () + + +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_value(sent, 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, + ) + + @property + def records_every_allow(self) -> bool: + return self._scope == "per_call" + + def should_record_allow( + self, + *, + session_id: str | None, + caller: Caller, + 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(_session_key(caller, session_id, input_type)) + return assert_never(self._scope) + + 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 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 44e2cc2404f..b990e0d5cc5 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -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,17 @@ class GenericGuardrailAPIOptionalParams(BaseModel): ), ) + guardrail_information_scope: GuardrailInformationScope | None = Field( + default=None, + description=( + "How often a call that allows the content unchanged records its guardrail entry in spend logs, OTEL and " + "logging callbacks. 'per_call' (default) records every call. 'per_session' records the first unchanged " + "allow of each session once per side, keyed by the caller and the session id, and records every call " + "that has no session id. 'off' records none. Blocks, rewrites, errors and fail-open passthroughs are " + "always recorded." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 7bfdfb00faf..f9a638b88d1 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -5,10 +5,12 @@ from unittest.mock import AsyncMock import pytest +from litellm.exceptions import GuardrailRaisedException 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 +2209,96 @@ class TestRecordsOwnGuardrailInformation: assert _guardrail_entries(request_data) == [] +def _skip_success_record_then(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs: + 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]: + 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.""" @@ -2217,8 +2309,6 @@ class _UndecoratedGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: - from litellm.exceptions import GuardrailRaisedException - if any("forbidden" in text for text in inputs.get("texts") or []): raise GuardrailRaisedException(guardrail_name=self.guardrail_name, message="Content blocked") return inputs @@ -2273,8 +2363,6 @@ class TestUndecoratedApplyGuardrailIsLogged: @pytest.mark.asyncio async def test_undecorated_block_is_recorded_and_reraised(self): - from litellm.exceptions import GuardrailRaisedException - guardrail = _UndecoratedGuardrail(guardrail_name="docs-style") request_data: dict = {"model": "gpt-4o"} diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py new file mode 100644 index 00000000000..86a48ef144d --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py @@ -0,0 +1,444 @@ +from collections.abc import Callable, Iterator, Mapping +from types import MappingProxyType +from typing import Final, Literal + +import httpx +import pytest + +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, initialize_guardrail +from litellm.types.guardrails import LitellmParams +from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GuardrailInformationScope, +) +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", [["success"], ["success"]]), ("per_session", [["success"], []]), ("off", [[], []])], +) +async def test_chat_completion_turns_through_the_openai_handler_follow_the_scope( + scope: GuardrailInformationScope, expected: list[list[str]] +) -> None: + guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope=scope) + + async def chat_turn() -> list[str]: + 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 _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.asyncio +@pytest.mark.parametrize("misspelled", ["per-session", "Per_Session", "session"]) +async def test_a_misspelled_scope_is_ignored_and_every_call_is_recorded(misspelled: str) -> None: + guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope=misspelled) + + assert [await _run(guardrail, _turn()) for _ in range(2)] == [["success"], ["success"]] + + +@pytest.mark.asyncio +async def test_a_scope_set_in_the_guardrail_config_dedups_the_session() -> None: + guardrail: Final = initialize_guardrail( + LitellmParams( + guardrail="generic_guardrail_api", + mode="pre_call", + api_base="https://guardrail.test", + guardrail_information_scope="per_session", + ), + {"guardrail_name": "configured-scope"}, + ) + guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(_Endpoint(_allow))) + + assert [await _run(guardrail, _turn()) for _ in range(2)] == [["success"], []] + + +@pytest.mark.asyncio +async def test_a_misspelled_scope_in_the_guardrail_config_still_loads_the_guardrail() -> None: + guardrail: Final = initialize_guardrail( + LitellmParams( + guardrail="generic_guardrail_api", + mode="pre_call", + api_base="https://guardrail.test", + guardrail_information_scope="per-session", + ), + {"guardrail_name": "misspelled-scope"}, + ) + guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(_Endpoint(_allow))) + + assert [await _run(guardrail, _turn()) for _ in range(2)] == [["success"], ["success"]], ( + "a typo in a logging setting must neither drop the guardrail nor its allow entries" + ) + + +@pytest.mark.asyncio +async def test_per_session_keeps_team_mates_without_a_key_hash_apart() -> None: + guardrail: Final = _guardrail(_Endpoint(_allow), guardrail_information_scope="per_session") + team_mates: Final = ( + {"user_api_key_team_id": "team", "user_api_key_user_id": "alice"}, + {"user_api_key_team_id": "team", "user_api_key_user_id": "bob"}, + ) + + assert [await _run(guardrail, _turn("shared-id", identity)) for identity in team_mates] == [["success"]] * 2, ( + "a second user of the same team must get their own first entry for the session" + ) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_record_scope.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_record_scope.py new file mode 100644 index 00000000000..dfd043b55be --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_record_scope.py @@ -0,0 +1,114 @@ +from types import MappingProxyType +from typing import Final + +import pytest + +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.record_scope import ( + Caller, + RecordScope, + returned_unchanged, +) +from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( + GuardrailInformationScope, + GuardrailToolParam, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +_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) + + +@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", caller=Caller(key_hash="hash-a", team_id=None, user_id=None), 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), + ) + + +class _Opaque: + pass + + +def test_a_non_json_value_passed_through_compares_as_unchanged() -> None: + passed_through: Final = _Opaque() + + assert returned_unchanged( + _chat_request(tools=[passed_through]), GenericGuardrailAPIInputs(texts=["hello"], tools=[passed_through]) + ) + + +@pytest.mark.parametrize(("scope", "recorded"), [("per_call", True), ("off", False)]) +def test_per_call_and_off_decide_without_claiming_a_session(scope: GuardrailInformationScope, recorded: bool) -> None: + sessions: Final = InMemoryCache() + record_scope: Final = RecordScope(scope, recorded_sessions=sessions) + caller: Final = Caller(key_hash="hash-a", team_id=None, user_id=None) + + decisions: Final = [ + record_scope.should_record_allow(session_id="s1", caller=caller, input_type="request") for _ in range(2) + ] + + session_still_unclaimed: Final = RecordScope("per_session", recorded_sessions=sessions).should_record_allow( + session_id="s1", caller=caller, input_type="request" + ) + + assert (decisions, session_still_unclaimed) == ([recorded, recorded], True) + + +def test_a_huge_session_id_is_stored_under_a_fixed_size_key() -> None: + sessions: Final = InMemoryCache() + record_scope: Final = RecordScope("per_session", recorded_sessions=sessions) + caller: Final = Caller(key_hash="hash-a", team_id=None, user_id=None) + huge_session_id: Final = "s" * 1_000_000 + + decisions: Final = [ + record_scope.should_record_allow(session_id=huge_session_id, caller=caller, input_type="request") + for _ in range(2) + ] + + assert (decisions, [len(key) for key in sessions.cache_dict]) == ([True, False], [64]), ( + "a caller-chosen session id must not grow the in-memory dedup cache" + )