test(guardrails): assert whole guardrail payloads and cover the blocked-request log

The apply_guardrail tests compare the whole request_data again, with the
expected identity built from get_authenticated_identity_metadata and the key
hash pinned, so an extra leaked key or a wrong hash fails them. A new
pass-through test checks that a request blocked by a pre-call guardrail logs
the cleaned and redacted inbound headers. The realtime Gray Swan test now
injects an HTTP transport instead of overriding a private method, so the real
request path runs

_ensure_litellm_metadata is renamed to
_apply_authenticated_identity_to_litellm_metadata, and docstrings that only
repeated the code are gone. The apply_guardrail strip no longer lists
user_api_key, because the authenticated identity always overwrites it. The
identity helpers now return Mapping, since no caller mutates what they return
This commit is contained in:
Caduri Katzav 2026-09-30 17:42:36 +03:00
parent 8fe7b4f00f
commit 4b63b53fe4
9 changed files with 134 additions and 138 deletions

View file

@ -156,8 +156,7 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp
_PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_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."""
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)
@ -228,7 +227,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,
@ -274,7 +273,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,

View file

@ -500,7 +500,6 @@ def is_untrusted_caller_metadata_key(key: str) -> bool:
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
@ -512,17 +511,13 @@ def strip_untrusted_caller_metadata(
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."""
) -> 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 in _GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS)
if not (is_untrusted_caller_metadata_key(key) or key == "headers")
}
return {**caller_fields, **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)}
@ -1677,7 +1672,7 @@ class LiteLLMProxyRequestSetup:
return user_api_key_logged_metadata
@staticmethod
def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]:
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),
@ -1686,8 +1681,7 @@ class LiteLLMProxyRequestSetup:
}
@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."""
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),
@ -2053,7 +2047,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

View file

@ -1148,7 +1148,7 @@ async def pass_through_request(
general_settings_view,
)
# Guardrails forward these to vendors as the inbound request headers; all are popped before the upstream send.
# 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")):

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,33 +569,23 @@ 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_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,
)
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}
_ensure_litellm_metadata(data, user_auth)
_apply_authenticated_identity_to_litellm_metadata(data, user_auth)
assert data["litellm_metadata"] is bucket
assert bucket["existing"] == "value"

View file

@ -2434,8 +2434,6 @@ def _cli_session_key(route: str) -> UserAPIKeyAuth:
class TestGuardrailsSeeAuthenticatedIdentity:
"""A request body cannot make a guardrail vendor see another key's identity, and the real one reaches it."""
@staticmethod
def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail:
def vendor(request: httpx.Request) -> httpx.Response:
@ -2505,7 +2503,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
@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 = {
@ -2528,8 +2525,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
async def test_chat_path_request_keeps_proxy_metadata_and_sends_stable_hash(
self, monkeypatch, route: str, make_key: Callable[[str], UserAPIKeyAuth]
) -> None:
"""After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata
in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token."""
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
request = MagicMock(spec=Request)
request.url = MagicMock()
@ -2599,8 +2594,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
@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")
@ -2616,4 +2609,4 @@ class TestGuardrailsSeeAuthenticatedIdentity:
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
assert vendor_payloads[0]["request_data"].get("user_api_key_hash") is None, "a kept token becomes the hash"

View file

@ -1,5 +1,6 @@
import json
import time
from collections.abc import Mapping
from datetime import datetime
from typing import Dict, List, Optional
from unittest.mock import AsyncMock
@ -36,6 +37,7 @@ 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 (
@ -1490,6 +1492,10 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker):
}
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)
@ -1532,18 +1538,14 @@ 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()
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"]
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={"metadata": {**_identity(caller), "forbidden_topics": ["tax"]}},
input_type="request",
)
@pytest.mark.asyncio
@ -1559,20 +1561,18 @@ 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)
request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]
assert request_data["messages"] == messages
assert request_data["metadata"]["forbidden_topics"] == ["tax"]
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={"messages": messages, "metadata": {**_identity(caller), "forbidden_topics": ["tax"]}},
input_type="request",
)
@pytest.mark.asyncio
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",
@ -1595,18 +1595,17 @@ async def test_apply_guardrail_authenticated_identity_overrides_client_metadata(
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"]
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):
"""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")
@ -1628,20 +1627,13 @@ async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set
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 == {**_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"
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,
@ -1654,16 +1646,15 @@ async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied
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())
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["headers"] == real_headers
assert metadata == {**_identity(caller), "headers": real_headers}
@pytest.mark.asyncio
async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker):
"""End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when
the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash."""
vendor_payloads = []
def vendor(request: httpx.Request) -> httpx.Response:
@ -1692,7 +1683,6 @@ async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identit
@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(
@ -1712,7 +1702,6 @@ async def test_apply_guardrail_request_route_comes_from_the_key(mocker):
@pytest.mark.asyncio
async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker):
"""A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key."""
raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
vendor_payloads = []
@ -1739,27 +1728,21 @@ async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocke
@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"]})
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=UserAPIKeyAuth(key_alias="known-caller", team_id="known-team"),
)
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
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"
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 must be forwarded, not dropped, and empty
metadata still carries the authenticated identity."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(
@ -1768,15 +1751,12 @@ 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(key_alias="known-caller"),
)
caller = UserAPIKeyAuth(key_alias="known-caller")
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
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"
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]}, request_data={"messages": [], "metadata": _identity(caller)}, input_type="request"
)
@pytest.mark.asyncio

View file

@ -7632,11 +7632,6 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
@pytest.mark.asyncio
async def test_pass_through_request_strips_caller_identity_before_guardrail_hooks():
"""
Regression: a pass-through body skips add_litellm_data_to_request, so forged user_api_key_* fields, guardrail
control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body
is unchanged because these keys never reach it.
"""
upstream_bodies = []
def transport_handler(upstream_request: httpx.Request) -> httpx.Response:
@ -7731,10 +7726,48 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook
@pytest.mark.asyncio
async def test_pass_through_post_call_guardrails_receive_real_inbound_headers():
"""Post-call guardrails run on a copy of the body the litellm-param pop already stripped, so without an explicit
re-attach an operator's extra_headers allowlist forwarded nothing on the response side."""
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"})
@ -7786,7 +7819,7 @@ async def test_pass_through_post_call_guardrails_receive_real_inbound_headers():
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"}
)

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,
@ -61,11 +69,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()
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
mock_guardrail.apply_guardrail.assert_called_once_with(
inputs={"texts": ["Test text with PII"]},
request_data={"metadata": _identity_with_hash(user_api_key_dict)},
input_type="request",
)
@pytest.mark.asyncio
@ -197,8 +205,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()
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
mock_guardrail.apply_guardrail.assert_called_once_with(
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,6 +17,7 @@ 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
@ -3556,7 +3558,6 @@ async def test_provider_bytes_are_sent_raw_after_pacing():
@pytest.mark.asyncio
async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch):
"""Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it."""
received_request_data = []
class IdentityRecordingGuardrail(CustomGuardrail):
@ -3590,26 +3591,20 @@ async def test_realtime_transcript_guardrail_receives_authenticated_identity(mon
@pytest.mark.asyncio
async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch):
"""Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets."""
vendor_payloads: list[dict[str, object]] = []
class RecordingGraySwan(GraySwanGuardrail):
async def _call_grayswan_api(self, payload):
vendor_payloads.append(payload)
return {"violation": 0.0, "violated_rules": []}
def vendor(request: httpx.Request) -> httpx.Response:
vendor_payloads.append(json.loads(request.content))
return httpx.Response(200, json={"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,
)
],
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",
@ -3643,7 +3638,6 @@ async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pyte
async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value(
monkeypatch: pytest.MonkeyPatch, sdk_value: object
):
"""Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one."""
received_request_data: list[dict[str, object]] = []
class IdentityRecordingGuardrail(CustomGuardrail):