mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
790d8a1df6
commit
290ec6f5c9
5 changed files with 75 additions and 14 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue