diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d2fbb26bb02..d1b962d1226 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -12,7 +12,9 @@ import litellm from litellm._logging import redact_internal_details_from_client_message, verbose_logger from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend +from litellm.proxy._types import UserAPIKeyAuth from litellm.types.llms.openai import ( OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, @@ -123,6 +125,10 @@ DefaultLoggedRealTimeEventTypes: Final = [ ] +def _as_user_api_key_auth(user_api_key_dict: object) -> UserAPIKeyAuth | None: + return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None + + class RealTimeStreaming: def __init__( self, @@ -852,7 +858,12 @@ class RealTimeStreaming: try: await callback.apply_guardrail( inputs={"texts": [transcript], "images": []}, - request_data={"user_api_key_dict": self.user_api_key_dict}, + request_data={ + "user_api_key_dict": self.user_api_key_dict, + "litellm_metadata": BaseTranslation.transform_user_api_key_dict_to_metadata( + _as_user_api_key_auth(self.user_api_key_dict) + ), + }, input_type="request", ) except Exception as e: diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 89ad67f0485..074d1b275ef 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -81,43 +81,18 @@ class BaseTranslation(ABC): @staticmethod def transform_user_api_key_dict_to_metadata( - user_api_key_dict: Any | None, + user_api_key_dict: "UserAPIKeyAuth | None", ) -> dict[str, object]: - """ - Transform user_api_key_dict to a metadata dict with prefixed keys. - - Converts keys like 'user_id' to 'user_api_key_user_id' to clearly indicate - the source of the metadata. - - Args: - user_api_key_dict: UserAPIKeyAuth object or dict with user information - - Returns: - Dict with keys prefixed with 'user_api_key_' - """ + """The authenticated key's identity as prefixed metadata, an allowlist safe to hand to guardrail vendors.""" if user_api_key_dict is None: return {} + # Lazy: `import litellm` loads this module before litellm.Router exists, and litellm_pre_call_utils imports it + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup - # Convert to dict if it's a Pydantic object - user_dict = user_api_key_dict.model_dump() if hasattr(user_api_key_dict, "model_dump") else user_api_key_dict - - if not isinstance(user_dict, dict): - return {} - - # Transform keys to be prefixed with 'user_api_key_' - transformed: Final[dict[str, object]] = {} - for key, value in user_dict.items(): - # Skip None values and internal fields - if value is None or key.startswith("_"): - continue - - # If key already has the prefix, use as-is, otherwise add prefix - if key.startswith("user_api_key_"): - transformed[key] = value - else: - transformed[f"user_api_key_{key}"] = value - - return transformed + return { + **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict), + "user_api_key_key_alias": user_api_key_dict.key_alias, + } @staticmethod def merge_user_api_key_metadata_into_request( diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index 1f295a6e656..6d8f13e1044 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -20,6 +20,10 @@ if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging +_PROXY_OWNED_PAYLOAD_KEYS: Final = frozenset( + {"metadata", "litellm_metadata", "litellm_logging_obj", "proxy_server_request"} +) + class PassThroughEndpointHandler(BaseTranslation): """ @@ -80,7 +84,7 @@ class PassThroughEndpointHandler(BaseTranslation): from litellm.litellm_core_utils.safe_json_dumps import safe_dumps payload_to_check: Final = { - k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj") + k: v for k, v in data.items() if not k.startswith("_") and k not in _PROXY_OWNED_PAYLOAD_KEYS } verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail") return safe_dumps(payload_to_check) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 6053ab26726..18d62142503 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -34,6 +34,7 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( ) from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router +from litellm.proxy.litellm_pre_call_utils import caller_metadata_with_authenticated_identity from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import GuardrailsRepository @@ -2404,9 +2405,14 @@ async def apply_guardrail( if litellm_logging_obj is not None: _patch_logging_obj_for_guardrail(litellm_logging_obj, request) + processed_metadata: Final = data.get("metadata") + inbound_headers: Final = processed_metadata.get("headers") if isinstance(processed_metadata, dict) else None request_data: Final[dict] = { **({"messages": request.messages} if request.messages is not None else {}), - **({"metadata": request.metadata} if request.metadata is not None else {}), + "metadata": { + **caller_metadata_with_authenticated_identity(request.metadata, user_api_key_dict), + **({"headers": inbound_headers} if inbound_headers is not None else {}), + }, } _input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type) guardrailed_inputs: Final = await active_guardrail.apply_guardrail( 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 3d1a173635e..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())) @@ -292,9 +307,8 @@ class GenericGuardrailAPI(CustomGuardrail): if value is not None: result_metadata[field_name] = value - # handle user_api_key_token = user_api_key_hash - if metadata_dict.get("user_api_key_token") is not None: - result_metadata["user_api_key_hash"] = metadata_dict.get("user_api_key_token") + if litellm_metadata.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata: + result_metadata["user_api_key_hash"] = litellm_metadata["user_api_key_token"] verbose_proxy_logger.debug( "Generic Guardrail API: Extracted user metadata: %s", @@ -432,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 @@ -442,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/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index e4822195bec..cea6eb17aac 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -1798,16 +1798,6 @@ class PanwPrismaAirsHandler(CustomGuardrail): call_id, _mcp_tool, ) - elif not request_data and logging_obj is None and input_type == "request": - # Direct /apply_guardrail endpoint — empty request_data, no - # logging_obj. Existing behavior: synthesize UUID. - call_id = str(uuid.uuid4()) - request_data["litellm_call_id"] = call_id - verbose_proxy_logger.warning( - "PANW Prisma AIRS: litellm_call_id missing from empty " - "request_data, synthesized %s (direct /apply_guardrail?)", - call_id, - ) else: call_id = str(uuid.uuid4()) request_data["litellm_call_id"] = call_id diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 37c1829def4..095cdc0efa0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -20,7 +20,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation, StreamingScanKey from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import ( MCP_GUARDRAIL_CALL_TYPES, @@ -35,10 +37,6 @@ if TYPE_CHECKING: # Imported lazily at runtime (inside the streaming hook) to avoid a # module-level cyclic import with litellm.integrations.custom_guardrail. from litellm.integrations.custom_guardrail import ModifyResponseException - from litellm.llms.base_llm.guardrail_translation.base_translation import ( - BaseTranslation, - StreamingScanKey, - ) # Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message) @@ -57,7 +55,7 @@ class _EndpointTranslation(Protocol): def process_output_streaming_response(self) -> "Callable[..., Awaitable[object]]": ... @property - def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ... + def get_streaming_scan_key(self) -> Callable[[Sequence[object]], StreamingScanKey | None]: ... @property def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ... @@ -72,7 +70,7 @@ def _as_endpoint_translation(translation: _EndpointTranslation) -> _EndpointTran def resolve_endpoint_translation( user_api_key_dict: UserAPIKeyAuth, first_response_item: object | None -) -> "tuple[str, BaseTranslation] | None": +) -> tuple[str, BaseTranslation] | None: """ Resolve the endpoint guardrail translation for a streamed response: the request route wins, falling back to inferring the call type from the first @@ -109,7 +107,7 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]: return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0) -def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool: +def _is_redundant_scan(scan_key: StreamingScanKey | None, last_scan_key: StreamingScanKey | None) -> bool: if scan_key is None: return False return scan_key == last_scan_key or scan_key.has_nothing_to_scan @@ -156,16 +154,20 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp } -def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: - """Populate data['litellm_metadata'] from user_api_key_dict if absent.""" - if "litellm_metadata" not in data: - from litellm.llms.base_llm.guardrail_translation.base_translation import ( - BaseTranslation, - ) +_PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata"}) - user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict) - if user_metadata: - data["litellm_metadata"] = user_metadata + +def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: + """Overwrite the identity fields of data['litellm_metadata'] from the authenticated key, in place.""" + existing: Final = data.get("litellm_metadata") + if isinstance(existing, dict): + identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict) + existing.update({key: value for key, value in identity.items() if key not in _PROXY_ENRICHED_IDENTITY_FIELDS}) + existing.pop("user_api_key_token", None) + return + user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict) + if user_metadata: + data["litellm_metadata"] = user_metadata class UnifiedLLMGuardrails(CustomLogger): @@ -406,7 +408,7 @@ class UnifiedLLMGuardrails(CustomLogger): @staticmethod def _resolve_transform_call_type( user_api_key_dict: UserAPIKeyAuth, - mappings: Mapping[CallTypes, type["BaseTranslation"]], + mappings: Mapping[CallTypes, type[BaseTranslation]], ) -> str | None: """Resolve the call type for the incremental_diff path, or None if the route is unresolvable / unsupported. @@ -669,7 +671,7 @@ class UnifiedLLMGuardrails(CustomLogger): call_type: str, sampling_rate: int, end_of_stream_only: bool, - mappings: Mapping[CallTypes, type["BaseTranslation"]], + mappings: Mapping[CallTypes, type[BaseTranslation]], ) -> AsyncGenerator[object, None]: """Emit guardrail text transformations as new deltas on the stream. diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d48451de6b1..5e7bb9b2c44 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -496,6 +496,40 @@ def _strip_untrusted_request_header_controls( headers.pop(header_name, None) +def is_untrusted_caller_metadata_key(key: str) -> bool: + return key.startswith("user_api_key_") or key in _UNTRUSTED_METADATA_CONTROL_FIELDS + + +def strip_untrusted_caller_metadata( + data: MutableMapping[str, object], *, allow_client_message_redaction_opt_out: bool +) -> None: + """Remove, in place, the proxy-owned slots a caller put in either metadata bucket of a request body.""" + for user_meta in (data.get("metadata"), data.get("litellm_metadata")): + if not isinstance(user_meta, dict): + continue + _strip_untrusted_request_header_controls( + user_meta.get("headers"), + allow_client_message_redaction_opt_out=allow_client_message_redaction_opt_out, + ) + for untrusted_key in tuple(key for key in user_meta if is_untrusted_caller_metadata_key(key)): + user_meta.pop(untrusted_key, None) + + +_GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS: Final = frozenset({"user_api_key", "headers"}) + + +def caller_metadata_with_authenticated_identity( + caller_metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth +) -> dict[str, object]: + """Caller metadata minus proxy-owned slots, bare user_api_key and headers, with the key's identity on top.""" + caller_fields: Final = { + key: value + for key, value in (caller_metadata or {}).items() + if not (is_untrusted_caller_metadata_key(key) or key in _GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS) + } + return {**caller_fields, **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)} + + def _is_false_like(value: object) -> bool: if isinstance(value, bool): return value is False @@ -523,7 +557,7 @@ def _key_or_team_allows_client_mock_response( ) -def _key_or_team_allows_client_message_redaction_opt_out( +def key_or_team_allows_client_message_redaction_opt_out( user_api_key_dict: UserAPIKeyAuth, ) -> bool: return _key_or_team_metadata_flag_is_true( @@ -1645,6 +1679,24 @@ class LiteLLMProxyRequestSetup: ) return user_api_key_logged_metadata + @staticmethod + def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: + return { + "user_api_key_metadata": strip_callback_config(user_api_key_dict.metadata), + "user_api_key_team_metadata": strip_callback_config(user_api_key_dict.team_metadata), + "user_api_key_object_permission_id": user_api_key_dict.object_permission_id, + "user_api_key_team_object_permission_id": user_api_key_dict.team_object_permission_id, + } + + @staticmethod + def get_authenticated_identity_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: + """Identity fields derived from the authenticated key alone, for paths that skip the chat-path build.""" + return { + **LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict), + "user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict), + **LiteLLMProxyRequestSetup.get_key_scoped_metadata(user_api_key_dict), + } + @staticmethod def add_user_api_key_auth_to_request_metadata( data: dict, @@ -2006,7 +2058,7 @@ async def add_litellm_data_to_request( # These keys are injected by the proxy itself below — user-supplied values # must not be trusted. _allow_client_mock_response: Final = _key_or_team_allows_client_mock_response(user_api_key_dict) - _allow_client_message_redaction_opt_out = _key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict) + _allow_client_message_redaction_opt_out = key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict) for _internal_key in _UNTRUSTED_ROOT_CONTROL_FIELDS: if _allow_client_mock_response and _internal_key in _CLIENT_MOCK_CONTROL_FIELDS: continue @@ -2188,31 +2240,12 @@ async def add_litellm_data_to_request( # profile_id) don't see attacker-injected admin slots preserved in # the deepcopy. - # Strip internal pipeline state and admin-injection slots from user input. # Runs AFTER the string-to-dict parse above so JSON-string metadata (sent # via multipart/form-data or extra_body) cannot smuggle admin fields past # the isinstance(dict) guard. - # - # The proxy populates a family of ``user_api_key_*`` fields below - # (user_api_key_metadata, user_api_key_user_id, user_api_key_alias, - # user_api_key_spend, user_api_key_team_metadata, …) into - # data[_metadata_variable_name]. Because the proxy only writes to ONE of - # the two metadata dicts, a caller pre-populating any of these keys on - # the OTHER metadata dict would have their forged values surface in - # guardrails, spend tracking, audit logs, and identity resolution. Strip - # by prefix so new ``user_api_key_*`` fields added in the future are - # covered without per-key maintenance. - for _meta_key in ("metadata", "litellm_metadata"): - _user_meta = data.get(_meta_key) - if isinstance(_user_meta, dict): - _strip_untrusted_request_header_controls( - _user_meta.get("headers"), - allow_client_message_redaction_opt_out=(_allow_client_message_redaction_opt_out), - ) - for _k in [ - k for k in _user_meta if k.startswith("user_api_key_") or k in _UNTRUSTED_METADATA_CONTROL_FIELDS - ]: - _user_meta.pop(_k, None) + strip_untrusted_caller_metadata( + data, allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out + ) # Strip pricing overrides AFTER the litellm_metadata string-to-dict parse # above, for the same reason as the user_api_key_* strip — JSON-string @@ -2389,14 +2422,7 @@ async def add_litellm_data_to_request( data[_metadata_variable_name]["user_api_key_user_model_max_budget"] = user_model_budget # rebind-ok: out-param data[_metadata_variable_name].update(carried_budget_metadata(user_api_key_dict)) - data[_metadata_variable_name]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata) - data[_metadata_variable_name]["user_api_key_team_metadata"] = strip_callback_config(user_api_key_dict.team_metadata) - data[_metadata_variable_name]["user_api_key_object_permission_id"] = getattr( - user_api_key_dict, "object_permission_id", None - ) - data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = getattr( - user_api_key_dict, "team_object_permission_id", None - ) + data[_metadata_variable_name].update(LiteLLMProxyRequestSetup.get_key_scoped_metadata(user_api_key_dict)) data[_metadata_variable_name]["headers"] = _logging_safe_headers data[_metadata_variable_name]["endpoint"] = str(request.url) # Carry the proxy-receive instant via metadata (like `endpoint`) so the diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..7bd1bb7871f 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -26,6 +26,7 @@ from fastapi import ( status, ) from fastapi.responses import StreamingResponse +from starlette.datastructures import Headers from starlette.datastructures import UploadFile as StarletteUploadFile from starlette.websockets import WebSocketState from websockets.asyncio.client import connect @@ -65,6 +66,7 @@ from litellm.llms.base_llm.managed_resources.utils import ( ) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.passthrough import BasePassthroughUtils +from litellm.proxy._experimental.mcp_server.utils import upstream_credential_headers from litellm.proxy._types import ( ConfigFieldInfo, ConfigFieldUpdate, @@ -100,6 +102,10 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + clean_headers, + key_or_team_allows_client_message_redaction_opt_out, + redact_credential_headers, + strip_untrusted_caller_metadata, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -1022,6 +1028,14 @@ from litellm.passthrough.timeout_utils import ( ) +def _guardrail_request_headers(headers: Headers, litellm_key_header_name: str | None) -> Mapping[str, str]: + cleaned: Final = clean_headers(headers, litellm_key_header_name=litellm_key_header_name) + mcp_credential_headers: Final = upstream_credential_headers(cleaned) + return redact_credential_headers( + {name: value for name, value in cleaned.items() if name.lower() not in mcp_credential_headers} + ) + + async def pass_through_request( request: Request, target: str, @@ -1120,6 +1134,26 @@ async def pass_through_request( _parsed_body = {} else: _parsed_body = await _read_request_body(request) + strip_untrusted_caller_metadata( + _parsed_body, + allow_client_message_redaction_opt_out=key_or_team_allows_client_message_redaction_opt_out( + user_api_key_dict + ), + ) + # Lazy: proxy_server imports this module + from litellm.proxy.proxy_server import ( + general_settings as proxy_general_settings, + ) + from litellm.proxy.proxy_server import ( + general_settings_view, + ) + + # Guardrails forward these to vendors as the inbound request headers; all are popped before the upstream send. + _parsed_body.pop("proxy_server_request", None) + _parsed_body.pop("headers", None) + for _caller_bucket in (_parsed_body.get("metadata"), _parsed_body.get("litellm_metadata")): + if isinstance(_caller_bucket, dict): + _caller_bucket.pop("headers", None) verbose_proxy_logger.debug( "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", url, @@ -1174,6 +1208,10 @@ async def pass_through_request( if _parsed_body is None: _parsed_body = {} _parsed_body["litellm_logging_obj"] = logging_obj + guardrail_headers: Final = _guardrail_request_headers( + request.headers, litellm_key_header_name=proxy_general_settings.get("litellm_key_header_name") + ) + _parsed_body["proxy_server_request"] = {"headers": guardrail_headers} ### CALL HOOKS ### - modify incoming data / reject request before calling the model _parsed_body = await proxy_logging_obj.pre_call_hook( @@ -1222,13 +1260,6 @@ async def pass_through_request( # provider IDs before forwarding upstream. Gated by feature flag and # enterprise managed-files hook. Runs after pre_call_hook so # guardrails have already seen the managed IDs. - from litellm.proxy.proxy_server import ( - general_settings as proxy_general_settings, - ) - from litellm.proxy.proxy_server import ( - general_settings_view, - ) - _managed_id_provider: Final = resolve_passthrough_managed_id_provider(custom_llm_provider) if proxy_general_settings.get("passthrough_managed_object_ids", False) and _managed_id_provider is not None: @@ -1636,6 +1667,7 @@ async def pass_through_request( **existing_metadata, "guardrails": guardrails_to_run, } + hook_data["proxy_server_request"] = {"headers": guardrail_headers} post_call_guardrail_data = hook_data response_body = await proxy_logging_obj.post_call_success_hook( data=hook_data, 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/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index a5e79f84ef1..3f77a24b707 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -416,6 +416,22 @@ class TestMetadataExtraction: assert request_metadata["user_api_key_hash"] == "hashed-token-value" assert request_metadata["user_api_key_user_id"] == "test-user" + @pytest.mark.parametrize( + "request_data, expected_hash", + [ + pytest.param({"metadata": {"user_api_key_token": "caller-token"}}, None, id="caller-bucket-token-ignored"), + pytest.param( + {"litellm_metadata": {"user_api_key_token": "proxy-token", "user_api_key_hash": "logged-key"}}, + "logged-key", + id="hash-wins-over-token", + ), + ], + ) + def test_token_fallback_only_from_litellm_metadata_and_only_without_hash( + self, generic_guardrail, request_data, expected_hash + ): + assert generic_guardrail._extract_user_api_key_metadata(request_data).get("user_api_key_hash") == expected_hash + @pytest.mark.asyncio async def test_metadata_extraction_empty_when_no_metadata(self, generic_guardrail): """Test metadata extraction returns empty dict when no metadata available""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index 53af7f36a5f..cce13844201 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -582,15 +582,20 @@ def test_ensure_litellm_metadata_populates_from_user_api_key_dict() -> None: assert data["litellm_metadata"]["user_api_key_team_id"] == "t1" -def test_ensure_litellm_metadata_noop_when_already_present() -> None: - """Verify _ensure_litellm_metadata does not overwrite existing litellm_metadata.""" +def test_ensure_litellm_metadata_overrides_caller_identity_in_existing_bucket() -> None: + """An existing litellm_metadata keeps its other keys, but its identity comes from the authenticated key.""" from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( _ensure_litellm_metadata, ) - user_auth = UserAPIKeyAuth(user_id="should-not-appear") - data: dict = {"litellm_metadata": {"existing": "value"}} + user_auth = UserAPIKeyAuth(user_id="auth-user", key_alias="auth-alias", team_id="auth-team") + bucket: dict = {"existing": "value", "user_api_key_alias": "batch-worker", "user_api_key_team_id": "team-exempt"} + data: dict = {"litellm_metadata": bucket} _ensure_litellm_metadata(data, user_auth) - assert data["litellm_metadata"] == {"existing": "value"} + assert data["litellm_metadata"] is bucket + assert bucket["existing"] == "value" + assert bucket["user_api_key_alias"] == "auth-alias" + assert bucket["user_api_key_team_id"] == "auth-team" + assert bucket["user_api_key_user_id"] == "auth-user" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index c90f88ec110..afebcbdf8a6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1,10 +1,16 @@ """Tests for unified guardrail.""" +import copy +import json import logging +from collections.abc import Callable from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal +from unittest.mock import MagicMock +import httpx import pytest +from fastapi import Request import litellm from litellm.caching import DualCache @@ -22,6 +28,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( openai_messages_without_tool, ) from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.mistral.ocr.guardrail_translation.handler import OCRHandler from litellm.llms.openai.chat.guardrail_translation.handler import ( OpenAIChatCompletionsHandler, @@ -33,12 +40,15 @@ from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import MCPGuardrailTranslationHandler, ) from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import ( unified_guardrail as unified_module, ) from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request +from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices @@ -2395,3 +2405,215 @@ class TestTranslationMappingsAreReadLive: assert not [ name for name, value in vars(unified_module).items() if isinstance(value, dict) and CallTypes.aocr in value ] + + +_RAW_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" + + +def _sk_key(route: str) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-real-caller-key", + key_alias="prod-app", + team_id="team-prod", + metadata={"key_label": "k1"}, + team_metadata={"phoenix_project_name": "team-proj", "priority": "high"}, + request_route=route, + ) + + +def _cli_session_key(route: str) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + token=_RAW_CLI_SESSION_TOKEN, + key_alias="cli-session-alice", + user_id="alice", + is_session_token=True, + team_id="team-prod", + team_metadata={"phoenix_project_name": "team-proj", "priority": "high"}, + request_route=route, + ) + + +class TestGuardrailsSeeAuthenticatedIdentity: + """A request body cannot make a guardrail vendor see another key's identity, and the real one reaches it.""" + + @staticmethod + def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail: + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"action": "NONE"}) + + guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic") + guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + return guardrail + + @pytest.mark.asyncio + @pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"]) + async def test_pass_through_body_cannot_forge_identity(self, monkeypatch, bucket: str) -> None: + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + vendor_payloads: list[dict[str, object]] = [] + key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") + data = { + "guardrail_to_apply": self._generic_guardrail(vendor_payloads), + "prompt": "hello", + bucket: { + "user_api_key_alias": "batch-worker", + "user_api_key_team_id": "team-exempt", + "user_api_key_token": "forged-hash", + }, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert len(vendor_payloads) == 1 + identity = vendor_payloads[0]["request_data"] + assert identity["user_api_key_alias"] == "prod-app" + assert identity["user_api_key_team_id"] == "team-prod" + assert identity["user_api_key_hash"] == key.api_key + + @pytest.mark.asyncio + async def test_mcp_tool_call_reaches_vendor_with_key_alias(self) -> None: + vendor_payloads: list[dict[str, object]] = [] + key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") + 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"] = self._generic_guardrail(vendor_payloads) + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.call_mcp_tool.value + ) + + assert len(vendor_payloads) == 1 + identity = vendor_payloads[0]["request_data"] + assert identity["user_api_key_alias"] == "prod-app" + assert identity["user_api_key_team_id"] == "team-prod" + assert identity["user_api_key_hash"] == key.api_key + + @pytest.mark.asyncio + async def test_pass_through_body_cannot_forge_request_route(self, monkeypatch) -> None: + """Guardrails and call-type lookups key on user_api_key_request_route, so it must be the key's own route.""" + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + key = UserAPIKeyAuth(api_key="sk-real-caller-key", request_route="/openai/v1/chat/completions") + data = { + "guardrail_to_apply": RecordingGuardrail(), + "prompt": "hello", + "litellm_metadata": {"user_api_key_request_route": "/v1/embeddings"}, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert data["litellm_metadata"]["user_api_key_request_route"] == "/openai/v1/chat/completions" + + @pytest.mark.asyncio + @pytest.mark.parametrize("route", ["/v1/messages", "/v1/responses", "/v1/chat/completions"]) + @pytest.mark.parametrize( + "make_key", [pytest.param(_sk_key, id="sk-key"), pytest.param(_cli_session_key, id="cli-session-key")] + ) + async def test_chat_path_request_keeps_proxy_metadata_and_sends_stable_hash( + self, monkeypatch, route: str, make_key: Callable[[str], UserAPIKeyAuth] + ) -> None: + """After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata + in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token.""" + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + request = MagicMock(spec=Request) + request.url = MagicMock() + request.url.path = route + request.url.__str__.return_value = "http://localhost" + route + request.method = "POST" + request.query_params = {} + request.headers = {"Content-Type": "application/json"} + request.client = MagicMock() + request.client.host = "127.0.0.1" + request.state = MagicMock() + key = make_key(route) + data = await add_litellm_data_to_request( + data={"model": "m", "messages": [{"role": "user", "content": "hi"}]}, + request=request, + user_api_key_dict=key, + proxy_config=MagicMock(), + general_settings={}, + version="v", + ) + proxy_bucket = data.get("litellm_metadata") + unshared = ("litellm_parent_otel_span", "user_api_key_auth") + bucket_before = ( + copy.deepcopy({k: v for k, v in proxy_bucket.items() if k not in unshared}) if proxy_bucket else None + ) + vendor_payloads: list[dict[str, object]] = [] + data["guardrail_to_apply"] = self._generic_guardrail(vendor_payloads) + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.acompletion.value + ) + + assert len(vendor_payloads) == 1 + assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == LiteLLMProxyRequestSetup.get_logged_api_key( + key + ) + assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0]) + if bucket_before is not None: + assert data["litellm_metadata"] is proxy_bucket + assert {k: v for k, v in proxy_bucket.items() if k in bucket_before} == bucket_before + assert bucket_before["user_api_key_auth_metadata"]["priority"] == "high" + assert "user_api_key_token" not in proxy_bucket + + @pytest.mark.asyncio + @pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata", None]) + async def test_pass_through_cli_session_key_sends_stable_hash(self, monkeypatch, bucket: str | None) -> None: + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + vendor_payloads: list[dict[str, object]] = [] + key = _cli_session_key("/anthropic/v1/messages") + forged_bucket = {bucket: {"user_api_key_token": "forged-hash"}} if bucket else {} + data = { + "guardrail_to_apply": self._generic_guardrail(vendor_payloads), + "prompt": "hello", + "proxy_server_request": {"headers": {"x-tenant": "tenant-real"}}, + **forged_bucket, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert len(vendor_payloads) == 1 + assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice" + assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0]) + assert vendor_payloads[0]["texts"] == ['{"prompt": "hello"}'] + assert vendor_payloads[0]["request_headers"] == {"x-tenant": "[present]"} + + @pytest.mark.asyncio + async def test_token_only_key_drops_forged_token_already_in_proxy_bucket(self, monkeypatch) -> None: + """A key with no api_key logs no hash, so a user_api_key_token left in litellm_metadata would become the + vendor's hash if the hook kept it.""" + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + vendor_payloads: list[dict[str, object]] = [] + key = UserAPIKeyAuth(token="abc123hashed", key_alias="prod-app") + data = { + "guardrail_to_apply": self._generic_guardrail(vendor_payloads), + "prompt": "hello", + "litellm_metadata": {"user_api_key_token": "forged-hash"}, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert len(vendor_payloads) == 1 + assert "user_api_key_token" not in data["litellm_metadata"] + assert vendor_payloads[0]["request_data"].get("user_api_key_hash") is None diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 508736fb78e..0d6b9b9d13f 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -4,11 +4,13 @@ from datetime import datetime from typing import Dict, List, Optional from unittest.mock import AsyncMock +import httpx import pytest from fastapi import HTTPException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_endpoints import ( CreateGuardrailRequest, @@ -33,6 +35,7 @@ from litellm.proxy.guardrails.guardrail_endpoints import ( from litellm.proxy.guardrails.guardrail_endpoints import ( test_custom_code_guardrail as run_custom_code_test_endpoint, ) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) from litellm.proxy.guardrails.guardrail_registry import ( @@ -1487,12 +1490,12 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker): } -def _patch_apply_guardrail_env(mocker, guardrail_result): +def _patch_apply_guardrail_env(mocker, guardrail_result, processed_data=None, guardrail=None): mock_guardrail = mocker.Mock() mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result) mock_registry = mocker.Mock() - mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail + mock_registry.get_initialized_guardrail_callback.return_value = guardrail or mock_guardrail mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry) mock_logging_obj = mocker.Mock() @@ -1500,7 +1503,7 @@ def _patch_apply_guardrail_env(mocker, guardrail_result): mock_logging_obj.model_call_details = {} mock_processor = mocker.Mock() mock_processor.common_processing_pre_call_logic = AsyncMock( - return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj) + return_value=(processed_data or {"guardrail_name": "test-guardrail"}, mock_logging_obj) ) mocker.patch( "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", @@ -1535,11 +1538,12 @@ async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker): user_api_key_dict=UserAPIKeyAuth(), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["What are tax loopholes?"]}, - request_data={"metadata": {"forbidden_topics": ["tax"]}}, - input_type="request", - ) + mock_guardrail.apply_guardrail.assert_awaited_once() + call = mock_guardrail.apply_guardrail.await_args.kwargs + assert call["inputs"] == {"texts": ["What are tax loopholes?"]} + assert call["input_type"] == "request" + assert "messages" not in call["request_data"] + assert call["request_data"]["metadata"]["forbidden_topics"] == ["tax"] @pytest.mark.asyncio @@ -1561,39 +1565,201 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker): user_api_key_dict=UserAPIKeyAuth(), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["What are tax loopholes?"]}, - request_data={ - "messages": messages, - "metadata": {"forbidden_topics": ["tax"]}, - }, - input_type="request", - ) + request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] + assert request_data["messages"] == messages + assert request_data["metadata"]["forbidden_topics"] == ["tax"] @pytest.mark.asyncio -async def test_apply_guardrail_omits_metadata_when_not_sent(mocker): - """Without metadata, request_data stays empty (backward-compatible).""" +async def test_apply_guardrail_authenticated_identity_overrides_client_metadata(mocker): + """A caller must not be able to claim another key's or team's identity in the body metadata.""" + mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + caller = UserAPIKeyAuth( + api_key="sk-real-caller-key", + key_alias="real-caller", + team_id="real-team", + user_id="real-user", + ) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={ + "user_api_key_alias": "exempt-batch-worker", + "user_api_key_team_id": "exempt-team", + "user_api_key_user_id": "someone-else", + "user_api_key_hash": "forged-hash", + "forbidden_topics": ["tax"], + }, + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["user_api_key_alias"] == "real-caller" + assert metadata["user_api_key_team_id"] == "real-team" + assert metadata["user_api_key_user_id"] == "real-user" + assert metadata["user_api_key_hash"] == caller.api_key + assert metadata["user_api_key_hash"] != "forged-hash" + assert metadata["forbidden_topics"] == ["tax"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set(mocker): + """Proxy-owned slots in the body never reach the guardrail, including user_api_key_token, which guardrails + map onto the key hash, and control fields the chat path also strips.""" + mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + caller = UserAPIKeyAuth(metadata={"zguard_policy_id": "strict"}, object_permission_id="perm-real") + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={ + "user_api_key_alias": "exempt-batch-worker", + "user_api_key_team_id": "exempt-team", + "user_api_key_token": "forged-hash", + "user_api_key_metadata": {"zguard_policy_id": "permissive"}, + "user_api_key_object_permission_id": "perm-forged", + "user_api_key": "forged-key", + "applied_guardrails": ["already-ran"], + "headers": {"x-end-user": "someone-else"}, + "trace_label": "nightly", + }, + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["user_api_key_alias"] is None + assert metadata["user_api_key_team_id"] is None + assert "user_api_key_token" not in metadata + assert metadata["user_api_key_metadata"] == {"zguard_policy_id": "strict"} + assert metadata["user_api_key_object_permission_id"] == "perm-real" + assert metadata["user_api_key"] is None + assert "applied_guardrails" not in metadata + assert "headers" not in metadata + assert metadata["trace_label"] == "nightly" + + +@pytest.mark.asyncio +async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied_ones(mocker): + """Guardrails forward metadata headers to vendors, so they must be the proxy's view of the request.""" + real_headers = {"user-agent": "real-client/1.0"} + mock_guardrail = _patch_apply_guardrail_env( + mocker, + {"texts": ["ok"]}, + processed_data={"guardrail_name": "test-guardrail", "metadata": {"headers": real_headers}}, + ) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={"headers": {"user-agent": "forged/1.0", "x-end-user": "someone-else"}}, + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=UserAPIKeyAuth()) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["headers"] == real_headers + + +@pytest.mark.asyncio +async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker): + """End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when + the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash.""" + vendor_payloads = [] + + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"action": "NONE"}) + + generic_guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic") + generic_guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + _patch_apply_guardrail_env(mocker, {"texts": ["unused"]}, guardrail=generic_guardrail) + caller = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="real-caller", team_id="real-team") + + request = ApplyGuardrailRequest( + guardrail_name="generic", + text="hello", + metadata={"user_api_key_alias": "exempt-batch-worker", "user_api_key_token": "forged-hash"}, + ) + response = await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + assert response.response_text == "hello" + assert len(vendor_payloads) == 1 + identity = vendor_payloads[0]["request_data"] + assert identity["user_api_key_alias"] == "real-caller" + assert identity["user_api_key_team_id"] == "real-team" + assert identity["user_api_key_hash"] == caller.api_key + + +@pytest.mark.asyncio +async def test_apply_guardrail_request_route_comes_from_the_key(mocker): + """Guardrails pick call-type behavior from user_api_key_request_route, so the body cannot choose it.""" + mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={"user_api_key_request_route": "/v1/embeddings"}, + ) + await apply_guardrail( + fastapi_request=mocker.Mock(), + request=request, + user_api_key_dict=UserAPIKeyAuth(request_route="/guardrails/apply_guardrail"), + ) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["user_api_key_request_route"] == "/guardrails/apply_guardrail" + + +@pytest.mark.asyncio +async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker): + """A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key.""" + raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" + vendor_payloads = [] + + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"action": "NONE"}) + + generic_guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic") + generic_guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + _patch_apply_guardrail_env(mocker, {"texts": ["unused"]}, guardrail=generic_guardrail) + caller = UserAPIKeyAuth( + token=raw_session_token, key_alias="cli-session-alice", user_id="alice", is_session_token=True + ) + + request = ApplyGuardrailRequest( + guardrail_name="generic", text="hello", metadata={"user_api_key_token": raw_session_token} + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + assert len(vendor_payloads) == 1 + assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice" + assert raw_session_token not in json.dumps(vendor_payloads[0]) + + +@pytest.mark.asyncio +async def test_apply_guardrail_carries_authenticated_identity_when_no_metadata_sent(mocker): + """request_data always carries the authenticated identity, even when the body has no metadata.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello") await apply_guardrail( fastapi_request=mocker.Mock(), request=request, - user_api_key_dict=UserAPIKeyAuth(), + user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller", team_id="known-team"), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["hello"]}, - request_data={}, - input_type="request", - ) + call = mock_guardrail.apply_guardrail.await_args.kwargs + assert call["inputs"] == {"texts": ["hello"]} + assert "messages" not in call["request_data"] + assert call["request_data"]["metadata"]["user_api_key_alias"] == "known-caller" + assert call["request_data"]["metadata"]["user_api_key_team_id"] == "known-team" @pytest.mark.asyncio async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker): - """Explicitly-sent empty messages/metadata must be forwarded, not dropped; - only omitted fields stay out of request_data.""" + """Explicitly-sent empty messages must be forwarded, not dropped, and empty + metadata still carries the authenticated identity.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) request = ApplyGuardrailRequest( @@ -1605,14 +1771,12 @@ async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(moc await apply_guardrail( fastapi_request=mocker.Mock(), request=request, - user_api_key_dict=UserAPIKeyAuth(), + user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller"), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["hello"]}, - request_data={"messages": [], "metadata": {}}, - input_type="request", - ) + request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] + assert request_data["messages"] == [] + assert request_data["metadata"]["user_api_key_alias"] == "known-caller" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..1b397ea3286 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -26,7 +26,12 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import ( + _extract_inbound_headers, +) +from litellm.proxy.litellm_pre_call_utils import _REDACTED_HEADER_VALUE from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, @@ -48,6 +53,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.types import utils as types_utils +from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -7622,3 +7628,166 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() metadata = kwargs["litellm_params"]["metadata"] assert metadata["user_api_key"] == "cli-session-alice" assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice" + + +@pytest.mark.asyncio +async def test_pass_through_request_strips_caller_identity_before_guardrail_hooks(): + """ + Regression: a pass-through body skips add_litellm_data_to_request, so forged user_api_key_* fields, guardrail + control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body + is unchanged because these keys never reach it. + """ + upstream_bodies = [] + + def transport_handler(upstream_request: httpx.Request) -> httpx.Response: + upstream_bodies.append(json.loads(upstream_request.content)) + return httpx.Response(200, json={"ok": True}) + + real_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.PassThroughEndpoint, + params={"timeout": resolve_pass_through_request_timeout(None)}, + ) + cache_dict = litellm.in_memory_llm_clients_cache.cache_dict + cache_key = next(key for key, cached in cache_dict.items() if cached is real_handler) + cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))) + + hook_data = [] + + def record_hook_data(user_api_key_dict, data, call_type): + hook_data.append({key: value for key, value in data.items() if key != "litellm_logging_obj"}) + return data + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=record_hook_data) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + forged = { + "user_api_key_alias": "batch-worker", + "user_api_key_team_id": "team-exempt", + "user_api_key_token": "forged-hash", + "user_api_key_request_route": "/v1/embeddings", + "disable_global_guardrails": True, + "headers": {"x-authenticated-user": "admin@corp"}, + "trace_label": "nightly", + } + forged_headers = {"x-authenticated-user": "admin@corp", "x-litellm-end-user-id": "victim", "x-tenant": "forged"} + body = { + "prompt": "hello", + "metadata": forged, + "litellm_metadata": forged, + "headers": forged_headers, + "proxy_server_request": {"headers": forged_headers}, + } + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.headers = Headers( + { + "content-type": "application/json", + "x-tenant": "tenant-real", + "authorization": "Bearer sk-real-caller-key", + "x-my-key": "sk-custom-header-key", + "x-mcp-github-authorization": "Bearer mcp-upstream-token", + "cookie": "session=secret", + } + ) + mock_request.query_params = QueryParams({}) + mock_request.body = AsyncMock(return_value=json.dumps(body).encode()) + + try: + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging + ), # test-quality-ok: read at call time + patch("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "x-my-key"}), + ): + response = await pass_through_request( + request=mock_request, + target="https://upstream.test/v1/generate", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"), + ) + finally: + cache_dict[cache_key] = real_handler + + assert response.status_code == 200 + assert hook_data == [ + { + "prompt": "hello", + "metadata": {"trace_label": "nightly"}, + "litellm_metadata": {"trace_label": "nightly"}, + "proxy_server_request": { + "headers": { + "content-type": "application/json", + "x-tenant": "tenant-real", + "cookie": _REDACTED_HEADER_VALUE, + } + }, + } + ] + vendor_headers = _extract_inbound_headers(request_data=hook_data[0], logging_obj=None, extra_allowlist={"x-tenant"}) + assert vendor_headers is not None and vendor_headers["x-tenant"] == "tenant-real" + assert upstream_bodies == [{"prompt": "hello"}] + + +@pytest.mark.asyncio +async def test_pass_through_post_call_guardrails_receive_real_inbound_headers(): + """Post-call guardrails run on a copy of the body the litellm-param pop already stripped, so without an explicit + re-attach an operator's extra_headers allowlist forwarded nothing on the response side.""" + + def transport_handler(upstream_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"completion": "hi"}) + + real_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.PassThroughEndpoint, + params={"timeout": resolve_pass_through_request_timeout(None)}, + ) + cache_dict = litellm.in_memory_llm_clients_cache.cache_dict + cache_key = next(key for key, cached in cache_dict.items() if cached is real_handler) + cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))) + + post_call_data = [] + + def record_post_call(data, user_api_key_dict, response): + post_call_data.append(data) + return response + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data) + mock_proxy_logging.post_call_success_hook = AsyncMock(side_effect=record_post_call) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.headers = Headers( + {"content-type": "application/json", "x-tenant": "tenant-real", "authorization": "Bearer sk-real-caller-key"} + ) + mock_request.query_params = QueryParams({}) + mock_request.body = AsyncMock( + return_value=json.dumps({"prompt": "hello", "headers": {"x-tenant": "forged"}}).encode() + ) + + try: + with patch( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging + ): # test-quality-ok: read at call time + response = await pass_through_request( + request=mock_request, + target="https://upstream.test/v1/generate", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"), + guardrails_config=["gg"], + ) + finally: + cache_dict[cache_key] = real_handler + + assert response.status_code == 200 + assert len(post_call_data) == 1, "the post-call guardrail hook did not run" + assert post_call_data[0]["proxy_server_request"] == { + "headers": {"content-type": "application/json", "x-tenant": "tenant-real"} + } + vendor_headers = _extract_inbound_headers( + request_data=post_call_data[0], logging_obj=None, extra_allowlist={"x-tenant"} + ) + assert vendor_headers is not None and vendor_headers["x-tenant"] == "tenant-real" diff --git a/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py index 4f44a4adeed..bab9958dc5d 100644 --- a/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py +++ b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py @@ -61,11 +61,11 @@ async def test_apply_guardrail_endpoint_returns_correct_response( assert response.response_text == "Redacted text: [REDACTED] and [REDACTED]" # Verify the guardrail was called with correct parameters - mock_guardrail.apply_guardrail.assert_called_once_with( - inputs={"texts": ["Test text with PII"]}, - request_data={}, - input_type="request", - ) + mock_guardrail.apply_guardrail.assert_called_once() + call = mock_guardrail.apply_guardrail.call_args.kwargs + assert call["inputs"] == {"texts": ["Test text with PII"]} + assert call["input_type"] == "request" + assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key @pytest.mark.asyncio @@ -197,6 +197,8 @@ async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_loggi assert response.response_text == "Processed text" # Verify the guardrail was called with correct parameters - mock_guardrail.apply_guardrail.assert_called_once_with( - inputs={"texts": ["Test text"]}, request_data={}, input_type="request" - ) + mock_guardrail.apply_guardrail.assert_called_once() + call = mock_guardrail.apply_guardrail.call_args.kwargs + assert call["inputs"] == {"texts": ["Test text"]} + assert call["input_type"] == "request" + assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 7e6d4d24905..a5713b355af 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -17,6 +17,8 @@ from litellm.litellm_core_utils.realtime_streaming import ( client_sent_openai_beta_realtime_header, ) from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail from litellm.types.guardrails import GuardrailEventHooks @@ -3550,3 +3552,120 @@ async def test_provider_bytes_are_sent_raw_after_pacing(): assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}'] provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01") + + +@pytest.mark.asyncio +async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch): + """Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it.""" + received_request_data = [] + + class IdentityRecordingGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + received_request_data.append(request_data) + return inputs + + monkeypatch.setattr( + litellm, + "callbacks", + [ + IdentityRecordingGuardrail( + guardrail_name="identity_recorder", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + ], + ) + key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=key) + + blocked = await streaming.run_realtime_guardrails("hello there") + + assert blocked is False + assert len(received_request_data) == 1 + identity = received_request_data[0]["litellm_metadata"] + assert identity["user_api_key_alias"] == "prod-app" + assert identity["user_api_key_team_id"] == "team-prod" + assert identity["user_api_key_hash"] == key.api_key + + +@pytest.mark.asyncio +async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch): + """Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets.""" + vendor_payloads: list[dict[str, object]] = [] + + class RecordingGraySwan(GraySwanGuardrail): + async def _call_grayswan_api(self, payload): + vendor_payloads.append(payload) + return {"violation": 0.0, "violated_rules": []} + + monkeypatch.setattr( + litellm, + "callbacks", + [ + RecordingGraySwan( + guardrail_name="grayswan", + api_key="test-key", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + ], + ) + key = UserAPIKeyAuth( + api_key="sk-real-caller-key", + key_alias="prod-app", + team_id="team-prod", + organization_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-ORG"}}]}, + jwt_claims={"email": "alice@corp.example", "name": "Alice Smith"}, + team_member={"user_id": "alice", "user_email": "alice@corp.example", "role": "admin"}, + ) + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=key) + + await streaming.run_realtime_guardrails("hello there", event_hooks=[GuardrailEventHooks.pre_call]) + + assert len(vendor_payloads) == 1 + vendor_metadata = vendor_payloads[0]["litellm_metadata"] + assert vendor_metadata["user_api_key_alias"] == "prod-app" + assert vendor_metadata["user_api_key_team_id"] == "team-prod" + serialized = json.dumps(vendor_metadata) + for leaked in ("SECRET-ORG", "Alice Smith", "alice@corp.example"): + assert leaked not in serialized + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "sdk_value", + [ + pytest.param({"key_alias": "forged", "team_id": "team-exempt"}, id="dict"), + pytest.param({"spend": "not-a-number"}, id="malformed-dict"), + pytest.param("sk-raw-string", id="string"), + ], +) +async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value( + monkeypatch: pytest.MonkeyPatch, sdk_value: object +): + """Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one.""" + received_request_data: list[dict[str, object]] = [] + + class IdentityRecordingGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + received_request_data.append(request_data) + return inputs + + monkeypatch.setattr( + litellm, + "callbacks", + [ + IdentityRecordingGuardrail( + guardrail_name="identity_recorder", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + ], + ) + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=sdk_value) + + blocked = await streaming.run_realtime_guardrails("hello there") + + assert blocked is False + assert len(received_request_data) == 1 + assert not [key for key in received_request_data[0]["litellm_metadata"] if key.startswith("user_api_key")] diff --git a/tests/unit/llms/base_llm/guardrail_translation/__init__.py b/tests/unit/llms/base_llm/guardrail_translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py b/tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py new file mode 100644 index 00000000000..06745c5a15f --- /dev/null +++ b/tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py @@ -0,0 +1,53 @@ +import json +from typing import Final + +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.proxy._types import UserAPIKeyAuth + +RAW_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" + + +def _fully_populated_session_key() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + token=RAW_SESSION_TOKEN, + is_session_token=True, + key_alias="cli-session-alice", + user_id="alice", + team_id="team-prod", + org_id="org-1", + metadata={"logging": [{"callback_name": "langfuse", "callback_vars": {"langfuse_secret_key": "SECRET-KEY"}}]}, + team_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-TEAM"}}]}, + organization_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-ORG"}}]}, + project_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-PROJECT"}}]}, + jwt_claims={"sub": "alice", "email": "alice@corp.example", "name": "Alice Smith"}, + team_member={"user_id": "alice", "user_email": "alice@corp.example", "role": "admin"}, + config={"internal": "proxy-config"}, + ) + + +def test_transform_emits_only_identity_never_credentials_or_callback_secrets(): + metadata = BaseTranslation.transform_user_api_key_dict_to_metadata(_fully_populated_session_key()) + serialized = json.dumps(metadata, default=str) + + assert metadata["user_api_key_alias"] == "cli-session-alice" + assert metadata["user_api_key_key_alias"] == "cli-session-alice" + assert metadata["user_api_key_hash"] == "cli-session-alice" + assert metadata["user_api_key_team_id"] == "team-prod" + assert RAW_SESSION_TOKEN not in serialized + assert "callback_vars" not in serialized + assert "SECRET" not in serialized + assert "Alice Smith" not in serialized + assert "proxy-config" not in serialized + for dropped in ( + "user_api_key_token", + "user_api_key_jwt_claims", + "user_api_key_team_member", + "user_api_key_organization_metadata", + "user_api_key_project_metadata", + "user_api_key_config", + ): + assert dropped not in metadata + + +def test_transform_of_no_key_is_empty(): + assert BaseTranslation.transform_user_api_key_dict_to_metadata(None) == {} 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"]