fix(guardrails): make guardrails see the authenticated identity on every path

Guardrails read user_api_key_alias, user_api_key_team_id, the key hash and the request route from request metadata, and the generic guardrail API forwards them to the vendor. The chat path strips caller copies of those fields and writes the real ones, but three other paths did not, so a caller could claim another key, team or request route (which call-type lookups key on)

/guardrails/apply_guardrail passed the body metadata straight to the guardrail. It now drops the fields the chat path treats as untrusted, plus the bare user_api_key and the caller's headers (a slightly wider strip than the chat path, on purpose). Then it adds the authenticated identity and the proxy's real request headers

Pass-through handed the raw body to pre_call_hook. It now runs the chat path's strip on both metadata buckets first, which also stops a body from switching off global guardrails, and drops the body's headers and proxy_server_request so it cannot pick the inbound headers a vendor sees. Those keys were already popped before the upstream send, so the forwarded body does not change. The pass-through guardrail text no longer includes litellm_metadata, which carried the key's identity dump into the scanned payload

The unified guardrail only filled litellm_metadata when it was missing, so a caller-supplied bucket won. It now overwrites the identity fields of an existing bucket and drops any user_api_key_token there, since the proxy never writes one. It leaves user_api_key_auth_metadata alone because on litellm_metadata routes the proxy has already merged team metadata into it

transform_user_api_key_dict_to_metadata used to dump every UserAPIKeyAuth field, including the raw token of non-sk keys, JWT claims, team membership, proxy config and org and project metadata with callback secrets. It now returns only the chat path's identity fields plus user_api_key_key_alias for existing readers, so MCP, pass-through, output-side handlers and realtime transcript guardrails all get that same identity allowlist and see the real user_api_key_alias. The generic guardrail uses user_api_key_token only from litellm_metadata and only when user_api_key_hash is missing, so a CLI session key's raw per-login token never reaches the vendor

The strip and the identity fields live in one place in litellm_pre_call_utils, and the chat path calls them too
This commit is contained in:
Caduri Katzav 2026-09-28 18:02:45 +03:00
parent 2c9b0e00ac
commit 1dbd73d512
18 changed files with 850 additions and 135 deletions

View file

@ -29,6 +29,7 @@ if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
from websockets.exceptions import ConnectionClosed
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
CLIENT_CONNECTION_CLASS = ClientConnection
@ -123,6 +124,12 @@ DefaultLoggedRealTimeEventTypes: Final = [
]
def _as_user_api_key_auth(user_api_key_dict: object) -> "UserAPIKeyAuth | None":
from litellm.proxy._types import UserAPIKeyAuth
return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None
class RealTimeStreaming:
def __init__(
self,
@ -831,6 +838,7 @@ class RealTimeStreaming:
typed user messages and tool outputs use ``pre_call``.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.guardrails import GuardrailEventHooks
if event_hooks is None:
@ -852,7 +860,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,17 @@ 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 {}
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,8 @@ 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"})
class PassThroughEndpointHandler(BaseTranslation):
"""
@ -80,7 +82,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

@ -292,9 +292,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",

View file

@ -1740,16 +1740,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

@ -155,16 +155,25 @@ 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."""
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
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):

View file

@ -493,6 +493,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
@ -520,7 +554,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(
@ -1642,6 +1676,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,
@ -2001,7 +2053,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
@ -2183,31 +2235,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
@ -2384,14 +2417,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

@ -100,6 +100,8 @@ 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
key_or_team_allows_client_message_redaction_opt_out,
strip_untrusted_caller_metadata,
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
@ -1120,6 +1122,18 @@ 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
),
)
# 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,

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,6 +1,7 @@
"""Tests for unified guardrail."""
import logging
from collections.abc import Callable
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Literal
@ -2395,3 +2396,228 @@ 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:
import json
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
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:
from litellm.proxy.utils import ProxyLogging
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."""
import copy
import json
from unittest.mock import MagicMock
from fastapi import Request
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request
_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:
import json
_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", **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"}']
@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

@ -1487,12 +1487,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 +1500,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 +1535,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 +1562,211 @@ 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."""
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
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."""
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
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 +1778,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

@ -7622,3 +7622,81 @@ 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.
"""
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
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"}
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"})
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
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"}}
]
assert upstream_bodies == [{"prompt": "hello"}]

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

@ -3550,3 +3550,134 @@ 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."""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
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."""
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail
from litellm.types.guardrails import GuardrailEventHooks
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."""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
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) == {}