From 290ec6f5c9d19202969146e3848d726669d6ba0a Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Thu, 1 Oct 2026 17:04:22 +0300 Subject: [PATCH] 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)