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
This commit is contained in:
Caduri Katzav 2026-10-01 17:04:22 +03:00
parent 790d8a1df6
commit 290ec6f5c9
5 changed files with 75 additions and 14 deletions

View file

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

View file

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

View file

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

View file

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

View file

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