feat(guardrails): add key alias and team skip filters to generic_guardrail_api

skip_if_key_alias_in and skip_if_team_id_in let an admin exempt specific
virtual keys or teams from a generic_guardrail_api guardrail. On the proxy
the match runs on the user_api_key_* identity the auth layer resolved, on the
request and response hooks independently, and a match sends nothing to the
guardrail endpoint. MCP tool results are still scanned for exempt keys. Key
aliases are picked by whoever creates or edits the key, so the option docs
point security-relevant exemptions at skip_if_team_id_in

A skipped call is logged as not_run with the option that matched, so it does
not count as a passed guardrail. A non-string identity value never matches,
and a bare string option is rejected so a YAML typo cannot turn into a set of
single characters

GenericGuardrailAPI also takes an optional async_handler so tests can inject
the HTTP client
This commit is contained in:
Caduri Katzav 2026-09-29 12:00:37 +03:00
parent 8fe7b4f00f
commit 779041bf26
9 changed files with 509 additions and 3 deletions

View file

@ -39,6 +39,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
skip_if_key_alias_in=_get_config_value(litellm_params, optional_params, "skip_if_key_alias_in"),
skip_if_team_id_in=_get_config_value(litellm_params, optional_params, "skip_if_team_id_in"),
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -0,0 +1,15 @@
import re
from collections.abc import Sequence
def config_values(raw: Sequence[str] | None, *, option_name: str) -> tuple[str, ...]:
if isinstance(raw, str):
raise ValueError(f"{option_name} must be a list of strings, got the single string {raw!r}")
return tuple(raw or ())
def compile_patterns(raw: Sequence[str] | None, *, option_name: str) -> tuple[re.Pattern[str], ...]:
try:
return tuple(re.compile(pattern) for pattern in config_values(raw, option_name=option_name))
except re.error as e:
raise ValueError(f"{option_name} contains an invalid regex: {e}") from e

View file

@ -20,9 +20,13 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.identity_filter import (
IdentitySkipFilter,
)
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 (
@ -170,6 +174,10 @@ def _structured_rows_to_write_back(
)
def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(**inputs)
class GenericGuardrailAPI(CustomGuardrail):
"""
Generic Guardrail API integration for LiteLLM.
@ -204,9 +212,14 @@ 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,
skip_if_key_alias_in: Sequence[str] | None = None,
skip_if_team_id_in: Sequence[str] | 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 +264,8 @@ class GenericGuardrailAPI(CustomGuardrail):
"block_only" if streaming_transform_mode is None else streaming_transform_mode
)
self.identity_skip_filter: Final = IdentitySkipFilter.from_config(skip_if_key_alias_in, skip_if_team_id_in)
# Set supported event hooks
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -431,6 +446,19 @@ class GenericGuardrailAPI(CustomGuardrail):
if request_data is None:
request_data = {}
user_metadata: Final = self._extract_user_api_key_metadata(request_data)
skip_option: Final = self.identity_skip_filter.matched_option(user_metadata)
if skip_option is not None:
verbose_proxy_logger.debug(
"Generic Guardrail API: skipping exempt caller per %s (input_type=%s)", skip_option, input_type
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=f"skipped: {skip_option}",
request_data=request_data,
guardrail_status="not_run",
)
return _passthrough_inputs(inputs)
request_body: Final = request_data.get("body") or {}
# Merge additional provider specific params from config and dynamic params
@ -441,8 +469,6 @@ class GenericGuardrailAPI(CustomGuardrail):
if dynamic_params:
additional_params.update(dynamic_params)
# Extract user API key metadata
user_metadata: Final = self._extract_user_api_key_metadata(request_data)
extra_allowlist = {h.lower() for h in self.extra_headers if isinstance(h, str)} if self.extra_headers else None
inbound_headers: Final = _extract_inbound_headers(
request_data=request_data,

View file

@ -0,0 +1,40 @@
from collections.abc import Sequence
from dataclasses import dataclass
from typing import Literal, TypeAlias
from typing_extensions import Self
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.config_parsing import config_values
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIMetadata,
)
IdentitySkipOption: TypeAlias = Literal["skip_if_key_alias_in", "skip_if_team_id_in"]
def _is_listed(value: object, listed: frozenset[str]) -> bool:
return isinstance(value, str) and value in listed
@dataclass(frozen=True, slots=True)
class IdentitySkipFilter:
key_aliases: frozenset[str] = frozenset()
team_ids: frozenset[str] = frozenset()
@classmethod
def from_config(
cls,
skip_if_key_alias_in: Sequence[str] | None,
skip_if_team_id_in: Sequence[str] | None,
) -> Self:
return cls(
key_aliases=frozenset(config_values(skip_if_key_alias_in, option_name="skip_if_key_alias_in")),
team_ids=frozenset(config_values(skip_if_team_id_in, option_name="skip_if_team_id_in")),
)
def matched_option(self, metadata: GenericGuardrailAPIMetadata) -> IdentitySkipOption | None:
if _is_listed(metadata.get("user_api_key_alias"), self.key_aliases):
return "skip_if_key_alias_in"
if _is_listed(metadata.get("user_api_key_team_id"), self.team_ids):
return "skip_if_team_id_in"
return None

View file

@ -103,6 +103,32 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
skip_if_key_alias_in: tuple[str, ...] | None = Field(
default=None,
description=(
"Skip the guardrail on the request and response hooks when the calling virtual key's "
"alias is in this list: nothing is sent to the guardrail endpoint and the call is "
"logged as not_run. On the proxy the match trusts the alias the auth layer resolved "
"for the key, never the request body. That alias is chosen by whoever creates or "
"edits the key, and by default internal users can create keys and rename their own, "
"so any of them can claim an alias that is unused or has been freed. List only "
"aliases held by admin-owned keys, and use skip_if_team_id_in for exemptions that "
"must hold. Outside the proxy the caller builds request_data, so the match is only "
"as trustworthy as that code. MCP tool results (post_mcp_call) are still scanned."
),
)
skip_if_team_id_in: tuple[str, ...] | None = Field(
default=None,
description=(
"Skip the guardrail for calls from a key whose team id is in this list, with the "
"same behavior as skip_if_key_alias_in. On the proxy the match trusts the team the "
"auth layer resolved for the key. Team ids are unique, and by default only admins "
"create teams and manage their members, so prefer this option when an exemption "
"must hold."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],

View file

View file

@ -0,0 +1,397 @@
import json
from dataclasses import dataclass, field
import httpx
import pytest
from starlette.requests import Request
import litellm
from litellm import ModelResponse
from litellm.caching.dual_cache import DualCache
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPI,
initialize_guardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.identity_filter import IdentitySkipFilter
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import ProxyConfig
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIOptionalParams,
)
from litellm.types.utils import CallTypes, Choices, Message
SCANNED = "[scanned]"
EXEMPT = {"skip_if_key_alias_in": ("batch-worker",), "skip_if_team_id_in": ("team-exempt",)}
EXEMPT_IDENTITIES = pytest.mark.parametrize(
"identity",
[{"user_api_key_alias": "batch-worker"}, {"user_api_key_team_id": "team-exempt"}],
ids=["alias", "team"],
)
@dataclass
class _GuardrailEndpoint:
received: list[dict] = field(default_factory=list)
def __call__(self, request: httpx.Request) -> httpx.Response:
self.received.append(json.loads(request.content))
return httpx.Response(200, json={"action": "GUARDRAIL_INTERVENED", "texts": [SCANNED]})
def _make_guardrail(
endpoint: _GuardrailEndpoint,
event_hook: tuple[GuardrailEventHooks, ...] | None = (GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call),
**options,
) -> GenericGuardrailAPI:
return GenericGuardrailAPI(
api_base="https://guardrail.test",
guardrail_name="identity-skip-test",
event_hook=None if event_hook is None else list(event_hook),
default_on=True,
async_handler=AsyncHTTPHandler(transport=httpx.MockTransport(endpoint)),
**options,
)
async def _apply(endpoint: _GuardrailEndpoint, request_data: dict, input_type: str = "request", **options) -> dict:
return await _make_guardrail(endpoint, **options).apply_guardrail(
inputs={"texts": ["hello"]}, request_data=request_data, input_type=input_type
)
@pytest.mark.asyncio
@pytest.mark.parametrize("input_type", ["request", "response"])
@EXEMPT_IDENTITIES
@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"])
async def test_exempt_caller_is_not_sent(input_type, identity, bucket):
endpoint = _GuardrailEndpoint()
result = await _apply(endpoint, {bucket: dict(identity)}, input_type, **EXEMPT)
assert endpoint.received == []
assert result == {"texts": ["hello"]}
@pytest.mark.asyncio
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_other_caller_is_scanned(input_type):
endpoint = _GuardrailEndpoint()
identity = {"user_api_key_alias": "prod-app", "user_api_key_team_id": "team-prod"}
result = await _apply(endpoint, {"litellm_metadata": identity}, input_type, **EXEMPT)
assert [sent["request_data"]["user_api_key_alias"] for sent in endpoint.received] == ["prod-app"]
assert result["texts"] == [SCANNED]
def _recorded_outcomes(request_data: dict) -> list[tuple[str, object]]:
_, bucket = get_or_create_metadata_bucket(request_data)
return [
(entry["guardrail_status"], entry["guardrail_response"])
for entry in bucket.get("standard_logging_guardrail_information", [])
]
@pytest.mark.asyncio
@pytest.mark.parametrize("input_type", ["request", "response"])
@pytest.mark.parametrize(
("identity", "option"),
[
({"user_api_key_alias": "batch-worker"}, "skip_if_key_alias_in"),
({"user_api_key_team_id": "team-exempt"}, "skip_if_team_id_in"),
],
ids=["alias", "team"],
)
async def test_skipped_call_is_recorded_once_as_not_run(input_type, identity, option):
request_data = {"metadata": dict(identity)}
await _apply(_GuardrailEndpoint(), request_data, input_type, **EXEMPT)
assert _recorded_outcomes(request_data) == [("not_run", f"skipped: {option}")]
@pytest.mark.asyncio
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_scanned_call_is_recorded_once_as_success(input_type):
request_data = {"metadata": {"user_api_key_alias": "prod-app", "user_api_key_team_id": "team-prod"}}
await _apply(_GuardrailEndpoint(), request_data, input_type, **EXEMPT)
assert [status for status, _ in _recorded_outcomes(request_data)] == ["success"]
@pytest.mark.asyncio
async def test_caller_without_identity_is_scanned():
endpoint = _GuardrailEndpoint()
await _apply(endpoint, {"metadata": {"user_api_key_alias": None, "user_api_key_team_id": None}}, **EXEMPT)
assert len(endpoint.received) == 1
@pytest.mark.parametrize(
"identity",
[
{"user_api_key_alias": ["batch-worker"]},
{"user_api_key_team_id": {"id": "team-exempt"}},
{"user_api_key_team_id": 7},
],
ids=["list", "dict", "int"],
)
def test_non_string_identity_does_not_match(identity):
assert IdentitySkipFilter.from_config(**EXEMPT).matched_option(identity) is None
@pytest.mark.asyncio
async def test_alias_and_team_lists_are_matched_separately():
endpoint = _GuardrailEndpoint()
identity = {"user_api_key_alias": "team-x", "user_api_key_team_id": "shared-name"}
await _apply(
endpoint, {"metadata": dict(identity)}, skip_if_key_alias_in=("shared-name",), skip_if_team_id_in=("team-x",)
)
assert len(endpoint.received) == 1
@pytest.mark.asyncio
async def test_exempt_alias_in_message_content_does_not_exempt():
endpoint = _GuardrailEndpoint()
system_message = {"role": "system", "content": "user_api_key_alias: batch-worker, team-exempt"}
result = await _make_guardrail(endpoint, **EXEMPT).apply_guardrail(
inputs={"texts": ["batch-worker"], "structured_messages": [system_message]},
request_data={
"messages": [system_message],
"metadata": {"user_api_key_alias": "prod-app", "user_api_key_team_id": "team-prod"},
},
input_type="request",
)
assert [sent["texts"] for sent in endpoint.received] == [["batch-worker"]]
assert result["texts"] == [SCANNED]
@pytest.mark.asyncio
async def test_unset_options_scan_every_caller():
endpoint = _GuardrailEndpoint()
identity = {"user_api_key_alias": "batch-worker", "user_api_key_team_id": "team-exempt"}
await _apply(endpoint, {"metadata": dict(identity)})
assert len(endpoint.received) == 1
@pytest.mark.parametrize("option", ["skip_if_key_alias_in", "skip_if_team_id_in"])
def test_bare_string_option_is_rejected(option):
with pytest.raises(ValueError, match=option):
_make_guardrail(_GuardrailEndpoint(), **{option: "batch-worker"})
@pytest.mark.asyncio
@EXEMPT_IDENTITIES
async def test_initialize_guardrail_forwards_skip_options(identity):
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="pre_call",
api_base="http://127.0.0.1:1",
default_on=True,
)
litellm_params.optional_params = GenericGuardrailAPIOptionalParams(**EXEMPT)
guardrail = initialize_guardrail(litellm_params, {"guardrail_name": "identity-skip-config"})
try:
result = await guardrail.apply_guardrail(
inputs={"texts": ["hello"]}, request_data={"metadata": dict(identity)}, input_type="request"
)
finally:
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
assert result == {"texts": ["hello"]}
ROUTES = pytest.mark.parametrize(
("route", "call_type", "body"),
[
("/v1/chat/completions", "acompletion", {"messages": [{"role": "user", "content": "hello"}]}),
("/v1/messages", "anthropic_messages", {"messages": [{"role": "user", "content": "hello"}], "max_tokens": 16}),
("/v1/responses", "aresponses", {"input": "hello"}),
],
ids=["chat", "messages", "responses"],
)
FORGED = {"user_api_key_alias": "batch-worker", "user_api_key_team_id": "team-exempt"}
def _request(route: str) -> Request:
return Request(
{
"type": "http",
"method": "POST",
"path": route,
"root_path": "",
"scheme": "http",
"query_string": b"",
"headers": [(b"content-type", b"application/json")],
"client": ("127.0.0.1", 1234),
"server": ("localhost", 4000),
}
)
async def _proxy_pre_call(route: str, body: dict, key: UserAPIKeyAuth) -> dict:
return await add_litellm_data_to_request(
data={"model": "gpt-4o", **body},
request=_request(route),
user_api_key_dict=key,
proxy_config=ProxyConfig(),
general_settings={},
)
@pytest.mark.asyncio
@ROUTES
async def test_body_supplied_identity_does_not_exempt_pre_call(route, call_type, body):
endpoint = _GuardrailEndpoint()
key = UserAPIKeyAuth(api_key="hashed", key_alias="prod-app", team_id="team-prod", request_route=route)
data = await _proxy_pre_call(route, {**body, "metadata": dict(FORGED), "litellm_metadata": dict(FORGED)}, key)
data["guardrail_to_apply"] = _make_guardrail(endpoint, **EXEMPT)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=call_type
)
assert [
(sent["request_data"]["user_api_key_alias"], sent["request_data"]["user_api_key_team_id"])
for sent in endpoint.received
] == [("prod-app", "team-prod")]
@pytest.mark.asyncio
@ROUTES
@pytest.mark.parametrize("identity", [{"key_alias": "batch-worker"}, {"team_id": "team-exempt"}], ids=["alias", "team"])
async def test_authenticated_exempt_key_is_skipped_pre_call(route, call_type, body, identity):
endpoint = _GuardrailEndpoint()
key = UserAPIKeyAuth(api_key="hashed", request_route=route, **identity)
data = await _proxy_pre_call(route, body, key)
data["guardrail_to_apply"] = _make_guardrail(endpoint, **EXEMPT)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=call_type
)
assert endpoint.received == []
@pytest.mark.asyncio
@pytest.mark.parametrize(
("key_alias", "body_metadata", "expected_aliases"),
[("prod-app", FORGED, ["prod-app"]), ("batch-worker", {}, [])],
ids=["forged-body", "exempt-key"],
)
async def test_post_call_uses_authenticated_identity(key_alias, body_metadata, expected_aliases):
route = "/v1/chat/completions"
endpoint = _GuardrailEndpoint()
key = UserAPIKeyAuth(api_key="hashed", key_alias=key_alias, team_id="team-prod", request_route=route)
body = {
"messages": [{"role": "user", "content": "hello"}],
"metadata": dict(body_metadata),
"litellm_metadata": dict(body_metadata),
}
data = await _proxy_pre_call(route, body, key)
data["guardrail_to_apply"] = _make_guardrail(endpoint, **EXEMPT)
response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="hi there"))])
await UnifiedLLMGuardrails().async_post_call_success_hook(data=data, user_api_key_dict=key, response=response)
assert [sent["request_data"]["user_api_key_alias"] for sent in endpoint.received] == expected_aliases
PROD_KEY = UserAPIKeyAuth(api_key="sk-prod", key_alias="prod-app", team_id="team-prod")
BARE_KEY = UserAPIKeyAuth(api_key="sk-bare")
EXEMPT_KEYS = pytest.mark.parametrize(
"key",
[
UserAPIKeyAuth(api_key="sk-batch", key_alias="batch-worker"),
UserAPIKeyAuth(api_key="sk-t", team_id="team-exempt"),
],
ids=["alias", "team"],
)
async def _pass_through_pre_call(key: UserAPIKeyAuth, body: dict) -> tuple[_GuardrailEndpoint, dict]:
endpoint = _GuardrailEndpoint()
data = {**body, "guardrail_to_apply": _make_guardrail(endpoint, **EXEMPT)}
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value
)
return endpoint, data
@pytest.mark.asyncio
@pytest.mark.parametrize("key", [PROD_KEY, BARE_KEY], ids=["key-with-identity", "key-without-identity"])
@pytest.mark.parametrize(
"buckets", [("metadata",), ("litellm_metadata",), ("metadata", "litellm_metadata")], ids=["md", "lmd", "both"]
)
async def test_pass_through_body_supplied_identity_does_not_exempt(key, buckets):
endpoint, _ = await _pass_through_pre_call(key, {"prompt": "hello", **{bucket: dict(FORGED) for bucket in buckets}})
assert [
(sent["request_data"].get("user_api_key_alias"), sent["request_data"].get("user_api_key_team_id"))
for sent in endpoint.received
] == [(key.key_alias, key.team_id)]
@pytest.mark.asyncio
@EXEMPT_KEYS
async def test_authenticated_exempt_key_is_skipped_on_pass_through(key):
endpoint, data = await _pass_through_pre_call(key, {"prompt": "hello"})
assert endpoint.received == []
assert [status for status, _ in _recorded_outcomes(data)] == ["not_run"]
async def _mcp_pre_call(key: UserAPIKeyAuth) -> tuple[_GuardrailEndpoint, dict]:
endpoint = _GuardrailEndpoint()
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
mcp_kwargs = {
"name": "search",
"arguments": {"query": "hello"},
"server_name": "docs",
"user_api_key_auth": key,
"user_api_key_user_id": key.user_id,
"user_api_key_team_id": key.team_id,
"user_api_key_end_user_id": None,
"user_api_key_hash": key.api_key,
"headers": {},
}
data = proxy_logging._convert_mcp_to_llm_format(
proxy_logging._create_mcp_request_object_from_kwargs(mcp_kwargs), mcp_kwargs
)
data["guardrail_to_apply"] = _make_guardrail(endpoint, event_hook=None, **EXEMPT)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.call_mcp_tool.value
)
return endpoint, data
@pytest.mark.asyncio
@EXEMPT_KEYS
async def test_authenticated_exempt_key_is_skipped_on_mcp_tool_call(key):
endpoint, data = await _mcp_pre_call(key)
assert endpoint.received == []
assert [status for status, _ in _recorded_outcomes(data)] == ["not_run"]
@pytest.mark.asyncio
async def test_other_key_is_scanned_on_mcp_tool_call():
endpoint, _ = await _mcp_pre_call(PROD_KEY)
assert [sent["request_data"]["user_api_key_alias"] for sent in endpoint.received] == ["prod-app"]