mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
8fe7b4f00f
commit
4b63b53fe4
9 changed files with 134 additions and 138 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue