This commit is contained in:
Caduri 2026-09-30 14:42:51 +00:00 • committed by GitHub
commit 71b5ca194a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
18 changed files with 954 additions and 163 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

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

@ -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,19 @@ 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 _apply_authenticated_identity_to_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
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):
@ -227,7 +228,7 @@ class UnifiedLLMGuardrails(CustomLogger):
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
_ensure_litellm_metadata(data, user_api_key_dict)
_apply_authenticated_identity_to_litellm_metadata(data, user_api_key_dict)
data = await endpoint_translation.process_input_messages(
data=data,
@ -273,7 +274,7 @@ class UnifiedLLMGuardrails(CustomLogger):
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
_ensure_litellm_metadata(data, user_api_key_dict)
_apply_authenticated_identity_to_litellm_metadata(data, user_api_key_dict)
return await endpoint_translation.process_input_messages(
data=data,
@ -406,7 +407,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 +670,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,35 @@ 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:
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)
def caller_metadata_with_authenticated_identity(
caller_metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth
) -> Mapping[str, object]:
caller_fields: Final = {
key: value
for key, value in (caller_metadata or {}).items()
if not (is_untrusted_caller_metadata_key(key) or key == "headers")
}
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 +552,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 +1674,23 @@ class LiteLLMProxyRequestSetup:
)
return user_api_key_logged_metadata
@staticmethod
def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> Mapping[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) -> Mapping[str, object]:
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 +2052,9 @@ 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: Final = 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 +2236,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 +2418,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,
)
# Only the proxy's own view of the inbound headers may reach guardrail vendors
_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

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

@ -10,6 +10,9 @@ from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import (
GraySwanGuardrail,
GraySwanGuardrailAPIError,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
_apply_authenticated_identity_to_litellm_metadata,
)
from litellm.types.guardrails import GuardrailEventHooks
@ -566,31 +569,26 @@ def test_prepare_payload_includes_litellm_metadata(
assert payload["litellm_metadata"]["user_api_key_team_id"] == "team-456"
def test_ensure_litellm_metadata_populates_from_user_api_key_dict() -> None:
"""Verify _ensure_litellm_metadata populates litellm_metadata."""
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
_ensure_litellm_metadata,
)
def test_missing_litellm_metadata_is_populated_from_user_api_key_dict() -> None:
user_auth = UserAPIKeyAuth(user_id="u1", team_id="t1", api_key="sk-test-hashed")
data: dict = {}
_ensure_litellm_metadata(data, user_auth)
_apply_authenticated_identity_to_litellm_metadata(data, user_auth)
assert "litellm_metadata" in data
assert data["litellm_metadata"]["user_api_key_user_id"] == "u1"
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."""
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
_ensure_litellm_metadata,
)
def test_existing_litellm_metadata_keeps_its_keys_but_takes_the_authenticated_identity() -> None:
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}
user_auth = UserAPIKeyAuth(user_id="should-not-appear")
data: dict = {"litellm_metadata": {"existing": "value"}}
_apply_authenticated_identity_to_litellm_metadata(data, user_auth)
_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,208 @@ 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:
@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:
_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:
_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:
_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, "a kept token becomes the hash"

View file

@ -1,14 +1,17 @@
import json
import time
from collections.abc import Mapping
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 +36,8 @@ 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
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
from litellm.proxy.guardrails.guardrail_registry import (
@ -1487,12 +1492,16 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker):
}
def _patch_apply_guardrail_env(mocker, guardrail_result):
def _identity(caller: UserAPIKeyAuth) -> Mapping[str, object]:
return LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(caller)
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 +1509,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",
@ -1529,15 +1538,12 @@ async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker):
text="What are tax loopholes?",
metadata={"forbidden_topics": ["tax"]},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
caller = UserAPIKeyAuth()
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={"metadata": {"forbidden_topics": ["tax"]}},
request_data={"metadata": {**_identity(caller), "forbidden_topics": ["tax"]}},
input_type="request",
)
@ -1555,45 +1561,188 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker):
messages=messages,
metadata={"forbidden_topics": ["tax"]},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
caller = UserAPIKeyAuth()
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={
"messages": messages,
"metadata": {"forbidden_topics": ["tax"]},
},
request_data={"messages": messages, "metadata": {**_identity(caller), "forbidden_topics": ["tax"]}},
input_type="request",
)
@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):
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 == {**_identity(caller), "forbidden_topics": ["tax"]}
assert (
metadata["user_api_key_alias"],
metadata["user_api_key_team_id"],
metadata["user_api_key_user_id"],
metadata["user_api_key_hash"],
) == ("real-caller", "real-team", "real-user", caller.api_key)
@pytest.mark.asyncio
async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set(mocker):
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 == {**_identity(caller), "trace_label": "nightly"}, "proxy-owned slots must not reach it"
assert metadata["user_api_key_metadata"] == {"zguard_policy_id": "strict"}
assert metadata["user_api_key_object_permission_id"] == "perm-real"
@pytest.mark.asyncio
async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied_ones(mocker):
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"}},
)
caller = UserAPIKeyAuth()
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 == {**_identity(caller), "headers": real_headers}
@pytest.mark.asyncio
async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker):
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):
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
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(),
user_api_key_dict=UserAPIKeyAuth(request_route="/guardrails/apply_guardrail"),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]},
request_data={},
input_type="request",
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):
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):
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
caller = UserAPIKeyAuth(key_alias="known-caller", team_id="known-team")
request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]}, request_data={"metadata": _identity(caller)}, input_type="request"
)
metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
assert (metadata["user_api_key_alias"], metadata["user_api_key_team_id"]) == ("known-caller", "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."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(
@ -1602,16 +1751,11 @@ async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(moc
messages=[],
metadata={},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
caller = UserAPIKeyAuth(key_alias="known-caller")
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]},
request_data={"messages": [], "metadata": {}},
input_type="request",
inputs={"texts": ["hello"]}, request_data={"messages": [], "metadata": _identity(caller)}, input_type="request"
)

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,199 @@ 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():
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_pre_call_block_logs_cleaned_inbound_headers():
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=HTTPException(status_code=400, detail="blocked"))
mock_proxy_logging.post_call_failure_hook = AsyncMock()
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",
"cookie": "session=secret",
}
)
mock_request.query_params = QueryParams({})
forged_headers = {"x-tenant": "forged"}
mock_request.body = AsyncMock(
return_value=json.dumps(
{"prompt": "hello", "headers": forged_headers, "proxy_server_request": {"headers": forged_headers}}
).encode()
)
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), # test-quality-ok: read at call time
pytest.raises(ProxyException),
):
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"),
)
mock_proxy_logging.post_call_failure_hook.assert_awaited_once()
logged_request = mock_proxy_logging.post_call_failure_hook.await_args.kwargs["request_data"]
assert logged_request["proxy_server_request"] == {
"headers": {"content-type": "application/json", "x-tenant": "tenant-real", "cookie": _REDACTED_HEADER_VALUE}
}
@pytest.mark.asyncio
async def test_pass_through_post_call_guardrails_receive_real_inbound_headers():
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"}
}, "the post-call body copy is already stripped, so the headers must be re-attached"
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

@ -11,9 +11,17 @@ from fastapi import HTTPException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.guardrails import ApplyGuardrailRequest, ApplyGuardrailResponse
def _identity_with_hash(caller: UserAPIKeyAuth) -> dict[str, object]:
return {
**LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(caller),
"user_api_key_hash": caller.api_key,
}
@pytest.mark.asyncio
async def test_apply_guardrail_endpoint_returns_correct_response(
mock_proxy_logging_ctx,
@ -63,7 +71,7 @@ async def test_apply_guardrail_endpoint_returns_correct_response(
# Verify the guardrail was called with correct parameters
mock_guardrail.apply_guardrail.assert_called_once_with(
inputs={"texts": ["Test text with PII"]},
request_data={},
request_data={"metadata": _identity_with_hash(user_api_key_dict)},
input_type="request",
)
@ -198,5 +206,7 @@ async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_loggi
# 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"
inputs={"texts": ["Test text"]},
request_data={"metadata": _identity_with_hash(user_api_key_dict)},
input_type="request",
)

View file

@ -5,6 +5,7 @@ from dataclasses import dataclass
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from websockets.exceptions import ConnectionClosed
from websockets.frames import Close
@ -16,7 +17,10 @@ from litellm.litellm_core_utils.realtime_streaming import (
RealTimeStreaming,
client_sent_openai_beta_realtime_header,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
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 +3554,112 @@ 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):
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):
vendor_payloads: list[dict[str, object]] = []
def vendor(request: httpx.Request) -> httpx.Response:
vendor_payloads.append(json.loads(request.content))
return httpx.Response(200, json={"violation": 0.0, "violated_rules": []})
grayswan = GraySwanGuardrail(
guardrail_name="grayswan",
api_key="test-key",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
grayswan.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor))
monkeypatch.setattr(litellm, "callbacks", [grayswan])
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
):
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) == {}