This commit is contained in:
Caduri 2026-09-30 10:30:03 +00:00 • committed by GitHub
commit 4cc22f6829
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
24 changed files with 1455 additions and 154 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -20,9 +20,13 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.identity_filter import (
IdentitySkipFilter,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
@ -170,6 +174,10 @@ def _structured_rows_to_write_back(
)
def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(**inputs)
class GenericGuardrailAPI(CustomGuardrail):
"""
Generic Guardrail API integration for LiteLLM.
@ -204,9 +212,14 @@ class GenericGuardrailAPI(CustomGuardrail):
streaming_end_of_stream_only: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
skip_if_key_alias_in: Sequence[str] | None = None,
skip_if_team_id_in: Sequence[str] | None = None,
async_handler: AsyncHTTPHandler | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.headers = headers or {}
self.extra_headers = extra_headers or []
@ -251,6 +264,8 @@ class GenericGuardrailAPI(CustomGuardrail):
"block_only" if streaming_transform_mode is None else streaming_transform_mode
)
self.identity_skip_filter: Final = IdentitySkipFilter.from_config(skip_if_key_alias_in, skip_if_team_id_in)
# Set supported event hooks
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -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,

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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) == {}

View file

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