From 779041bf26ec4ec4ca1413bcadabcb32e137aa29 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Tue, 29 Sep 2026 12:00:37 +0300 Subject: [PATCH] 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 --- .../generic_guardrail_api/__init__.py | 2 + .../generic_guardrail_api/config_parsing.py | 15 + .../generic_guardrail_api.py | 32 +- .../generic_guardrail_api/identity_filter.py | 40 ++ .../guardrail_hooks/generic_guardrail_api.py | 26 ++ tests/unit/proxy/guardrails/__init__.py | 0 .../guardrails/guardrail_hooks/__init__.py | 0 .../generic_guardrail_api/__init__.py | 0 .../test_identity_filter.py | 397 ++++++++++++++++++ 9 files changed, 509 insertions(+), 3 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/identity_filter.py create mode 100644 tests/unit/proxy/guardrails/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/__init__.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_identity_filter.py 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 e3511d46544..f0b6d5d28d7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py new file mode 100644 index 00000000000..10934c8e244 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py @@ -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 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 320682c89b7..1155aa26829 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 @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/identity_filter.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/identity_filter.py new file mode 100644 index 00000000000..bdbfecaa07b --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/identity_filter.py @@ -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 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..b33bea87c25 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -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], diff --git a/tests/unit/proxy/guardrails/__init__.py b/tests/unit/proxy/guardrails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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_identity_filter.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_identity_filter.py new file mode 100644 index 00000000000..1ea74661667 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_identity_filter.py @@ -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"]