diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index ba1b6e4c10d..dfd9c8041fb 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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: @@ -1629,9 +1640,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, @@ -1655,6 +1671,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) @@ -1669,9 +1686,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, @@ -1691,6 +1713,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 e3511d46544..1637cfacf42 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -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) 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 3d1a173635e..565f6d9e6cc 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 @@ -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 ( 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..095bd331345 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -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 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..c6e6952247b 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,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], diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 4649bddd281..7d167d2941f 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -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.""" diff --git a/tests/unit/proxy/guardrails/__init__.py b/tests/unit/proxy/guardrails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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_record_scope.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_record_scope.py new file mode 100644 index 00000000000..0521ffa6e0d --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_record_scope.py @@ -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), + )