From abdef108b8aad51dff6d0cf10ca97574c2551c0c Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Mon, 28 Sep 2026 18:06:43 +0300 Subject: [PATCH 1/6] feat(guardrails): add guardrail_information_scope to generic_guardrail_api Every generic_guardrail_api call appends a guardrail information entry to the request metadata, so a long agent session carries one entry per call in its spend log row. guardrail_information_scope lets operators keep that (per_call, the default), record only the first unchanged allow of a session on each of the request and response sides (per_session), or record no unchanged allows (off). Blocks, rewrites, errors and fail-open passthroughs are recorded under every scope and never claim a session. The per_session dedup key includes the authenticated key hash, falling back to the team id, since the session id comes from the caller. log_guardrail_information gains skip_guardrail_success_record(), backed by its own context variable that only the success branch reads, so skipping a success entry can never drop an error entry. --- litellm/integrations/custom_guardrail.py | 27 +- .../generic_guardrail_api/__init__.py | 1 + .../generic_guardrail_api.py | 27 +- .../generic_guardrail_api/record_scope.py | 56 +++ .../guardrail_hooks/generic_guardrail_api.py | 19 + .../integrations/test_custom_guardrail.py | 95 ++++ .../generic_guardrail_api/__init__.py | 0 .../test_record_scope.py | 433 ++++++++++++++++++ 8 files changed, 654 insertions(+), 4 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_record_scope.py diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 99bb832e26c..d91f44b0171 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -72,6 +72,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: @@ -1639,9 +1650,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, @@ -1665,6 +1681,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) @@ -1679,9 +1696,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, @@ -1701,6 +1723,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 786b65b1cc3..b140a9ed08c 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())) @@ -499,7 +513,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 +537,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 7bfdfb00faf..cfac33efc5c 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/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), + ) From 790d8a1df6a37c4dfe777b138799590990766de9 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Thu, 1 Oct 2026 15:17:56 +0300 Subject: [PATCH 2/6] fix(guardrails): ignore a misspelled guardrail_information_scope A value outside per_call, per_session and off went through unchecked, and the scope check then returned nothing, so a typo such as per-session acted like off and dropped every unchanged allow entry. The value is now validated when the guardrail is built, and an unknown one is ignored with a warning, so the guardrail keeps enforcing and records every call as it does under per_call. The match on the scope now ends in assert_never The apply_guardrail and initialize_guardrail tests move from test_record_scope.py to a test_generic_guardrail_api.py mirror, so a change to generic_guardrail_api.py selects them. New tests cover a scope set in the guardrail config, a misspelled scope on both paths and a non-JSON value passed through. The scope option's description is shorter, the record scope import is absolute and test imports are at the top of their files --- litellm/integrations/custom_guardrail.py | 5 +- .../generic_guardrail_api.py | 11 +- .../generic_guardrail_api/record_scope.py | 21 + .../guardrail_hooks/generic_guardrail_api.py | 16 +- .../integrations/test_custom_guardrail.py | 9 +- .../test_generic_guardrail_api.py | 431 ++++++++++++++++++ .../test_record_scope.py | 382 +--------------- 7 files changed, 478 insertions(+), 397 deletions(-) create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index d91f44b0171..41990b04c44 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -78,10 +78,7 @@ _guardrail_success_record_skipped: Final[contextvars.ContextVar[bool]] = context 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. - """ + """Only the success branch reads this flag, so a later raise still records its error entry""" _guardrail_success_record_skipped.set(True) 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 b140a9ed08c..431918a0687 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 @@ -26,6 +26,11 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.record_scope import ( + RecordScope, + guardrail_information_scope_from_config, + returned_unchanged, +) 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 ( @@ -37,8 +42,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ) 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 @@ -261,9 +264,7 @@ 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 - ) + 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())) 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 index 095bd331345..c418deb1206 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -1,8 +1,11 @@ import json from typing import Final, Literal +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.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GuardrailInformationScope from litellm.types.utils import GenericGuardrailAPIInputs @@ -12,6 +15,22 @@ DEFAULT_GUARDRAIL_INFORMATION_SCOPE: Final[GuardrailInformationScope] = "per_cal _SESSION_CACHE_MAX_ENTRIES: Final = 100_000 _SESSION_CACHE_TTL_SECONDS: Final = 3600 _REWRITABLE_KEYS: Final = ("texts", "images", "tools", "structured_messages") +_SCOPE_ADAPTER: Final[TypeAdapter[GuardrailInformationScope]] = TypeAdapter( + GuardrailInformationScope, config=ConfigDict(title="guardrail_information_scope") +) + + +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: @@ -48,6 +67,8 @@ class RecordScope: return False case "per_session": return session_id is None or self._claim_session(json.dumps([tenant, session_id, input_type])) + case _: + assert_never(self._scope) def _claim_session(self, key: str) -> bool: if self._recorded_sessions.get_cache(key) is 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 c6e6952247b..b990e0d5cc5 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -108,17 +108,11 @@ 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." + "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." ), ) diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index cfac33efc5c..f9a638b88d1 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -5,6 +5,7 @@ from unittest.mock import AsyncMock import pytest +from litellm.exceptions import GuardrailRaisedException from litellm.integrations.custom_guardrail import ( DEFAULT_ADVISORY_MESSAGE, CustomGuardrail, @@ -2209,8 +2210,6 @@ class TestRecordsOwnGuardrailInformation: 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() @@ -2240,8 +2239,6 @@ class _SuccessRecordSkippingGuardrail(CustomGuardrail): 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) @@ -2312,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 @@ -2368,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/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..c845caf83fc --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_generic_guardrail_api.py @@ -0,0 +1,431 @@ +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" + ) 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 index 0521ffa6e0d..1a1e09a4c5f 100644 --- 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 @@ -1,28 +1,18 @@ -from collections.abc import Callable, Iterator, Mapping from types import MappingProxyType -from typing import Final, Literal +from typing import Final -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"]}}} @@ -33,364 +23,6 @@ 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"), [ @@ -431,3 +63,15 @@ def test_undecodable_bytes_compare_without_raising() -> None: 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]) + ) From 290ec6f5c9d19202969146e3848d726669d6ba0a Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Thu, 1 Oct 2026 17:04:22 +0300 Subject: [PATCH 3/6] fix(guardrails): key per_session dedup by user and skip work under per_call A JWT caller with no key hash fell back to the team id alone, so two users of one team sending the same session id shared a dedup slot. The caller key is now the key hash, team id and user id together. Under the default per_call scope the unchanged-content comparison and the session lookup no longer run, since nothing is skipped there The session id is read through a TypeAdapter-validated request mapping, the scope match returns assert_never, and get_guardrail_dynamic_request_body_params is annotated with the dict[str, object] it returns, which also clears unknown-type errors at its callers --- litellm/integrations/custom_guardrail.py | 2 +- .../generic_guardrail_api.py | 23 ++++++++++----- .../generic_guardrail_api/record_scope.py | 28 +++++++++++++++---- .../test_generic_guardrail_api.py | 13 +++++++++ .../test_record_scope.py | 23 ++++++++++++++- 5 files changed, 75 insertions(+), 14 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 41990b04c44..87e98962871 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1161,7 +1161,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 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 431918a0687..b64c2f0664c 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,7 +17,6 @@ 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, ) @@ -27,9 +26,11 @@ from litellm.llms.custom_httpx.http_handler import ( 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 @@ -354,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, @@ -538,11 +539,19 @@ 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, + 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 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 index c418deb1206..5cc8a93192f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -1,5 +1,5 @@ import json -from typing import Final, Literal +from typing import Final, Literal, NamedTuple from pydantic import ConfigDict, TypeAdapter, ValidationError from pydantic_core import to_jsonable_python @@ -7,6 +7,7 @@ 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 @@ -15,11 +16,23 @@ DEFAULT_GUARDRAIL_INFORMATION_SCOPE: Final[GuardrailInformationScope] = "per_cal _SESSION_CACHE_MAX_ENTRIES: Final = 100_000 _SESSION_CACHE_TTL_SECONDS: Final = 3600 _REWRITABLE_KEYS: Final = ("texts", "images", "tools", "structured_messages") +_NOT_SENT: Final[tuple[()]] = () +_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 @@ -41,7 +54,8 @@ def returned_unchanged(sent: GenericGuardrailAPIInputs, returned: GenericGuardra """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 + key not in returned or _jsonable(returned.get(key)) == _jsonable(sent.get(key, _NOT_SENT)) + for key in _REWRITABLE_KEYS ) @@ -53,11 +67,15 @@ class RecordScope: 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, - tenant: str | None, + caller: Caller, input_type: Literal["request", "response"], ) -> bool: match self._scope: @@ -66,9 +84,9 @@ class RecordScope: case "off": return False case "per_session": - return session_id is None or self._claim_session(json.dumps([tenant, session_id, input_type])) + return session_id is None or self._claim_session(json.dumps((*caller, session_id, input_type))) case _: - assert_never(self._scope) + return assert_never(self._scope) def _claim_session(self, key: str) -> bool: if self._recorded_sessions.get_cache(key) is True: 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 index c845caf83fc..86a48ef144d 100644 --- 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 @@ -429,3 +429,16 @@ async def test_a_misspelled_scope_in_the_guardrail_config_still_loads_the_guardr 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 index 1a1e09a4c5f..97a1060a2ec 100644 --- 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 @@ -5,10 +5,12 @@ 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 @@ -48,7 +50,9 @@ def test_per_session_records_again_once_the_session_has_expired() -> None: ) def record() -> bool: - return record_scope.should_record_allow(session_id="s1", tenant="hash-a", input_type="request") + 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 @@ -75,3 +79,20 @@ def test_a_non_json_value_passed_through_compares_as_unchanged() -> None: 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) From a765188dc7125662365eb7c66eb3874b5d5dc288 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Thu, 1 Oct 2026 17:54:56 +0300 Subject: [PATCH 4/6] refactor(guardrails): end the record scope match with a final assert_never CodeQL does not treat a case _ arm as making the match exhaustive, so it flagged a possible implicit None return. The match now covers the three scopes and returns assert_never after it, as other matches in the repo do --- .../guardrail_hooks/generic_guardrail_api/record_scope.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) 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 index 5cc8a93192f..4ed2dcf83a3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -85,8 +85,7 @@ class RecordScope: return False case "per_session": return session_id is None or self._claim_session(json.dumps((*caller, session_id, input_type))) - case _: - return assert_never(self._scope) + return assert_never(self._scope) def _claim_session(self, key: str) -> bool: if self._recorded_sessions.get_cache(key) is True: From 045b13fe2109d3d9295aabedada97325dadbf294 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Thu, 1 Oct 2026 18:12:50 +0300 Subject: [PATCH 5/6] refactor(guardrails): drop the _NOT_SENT default in returned_unchanged A small _sent_value helper returns the sent value when the key is present and an empty tuple otherwise, so no sentinel constant is needed outside constants.py --- .../guardrail_hooks/generic_guardrail_api/record_scope.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) 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 index 4ed2dcf83a3..07b586df7f4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -16,7 +16,6 @@ DEFAULT_GUARDRAIL_INFORMATION_SCOPE: Final[GuardrailInformationScope] = "per_cal _SESSION_CACHE_MAX_ENTRIES: Final = 100_000 _SESSION_CACHE_TTL_SECONDS: Final = 3600 _REWRITABLE_KEYS: Final = ("texts", "images", "tools", "structured_messages") -_NOT_SENT: Final[tuple[()]] = () _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") @@ -50,11 +49,15 @@ def _jsonable(value: object) -> object: return to_jsonable_python(value, fallback=repr, bytes_mode="base64") +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.get(key, _NOT_SENT)) + key not in returned or _jsonable(returned.get(key)) == _jsonable(_sent_value(sent, key)) for key in _REWRITABLE_KEYS ) From a37c292b5f759e0abed9d6a3e529470fb1bd6ae8 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Thu, 1 Oct 2026 18:32:29 +0300 Subject: [PATCH 6/6] fix(guardrails): bound per_session dedup keys with a fixed-size hash The dedup cache keyed sessions by the serialized caller, session id and side, so a caller sending large distinct litellm_session_id values grew it by their full size for an hour. Keys are now the SHA-256 of that tuple, 64 characters each, so the 100k-entry cap bounds the cache's memory --- .../generic_guardrail_api/record_scope.py | 7 ++++++- .../generic_guardrail_api/test_record_scope.py | 16 ++++++++++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) 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 index 07b586df7f4..f51f8cf832c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/record_scope.py @@ -1,3 +1,4 @@ +import hashlib import json from typing import Final, Literal, NamedTuple @@ -49,6 +50,10 @@ 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 () @@ -87,7 +92,7 @@ class RecordScope: case "off": return False case "per_session": - return session_id is None or self._claim_session(json.dumps((*caller, session_id, input_type))) + 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: 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 index 97a1060a2ec..dfd043b55be 100644 --- 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 @@ -96,3 +96,19 @@ def test_per_call_and_off_decide_without_claiming_a_session(scope: GuardrailInfo ) 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" + )