From d7b6d3c904e15b238152899eebc56cb25e0ca7a3 Mon Sep 17 00:00:00 2001 From: albertbausili Date: Mon, 14 Sep 2026 09:03:15 +0200 Subject: [PATCH] fix(proxy): stop /apply_guardrail from forwarding caller-supplied key identity The endpoint passed request.metadata straight through as the guardrail's request_data metadata, so a caller could set user_api_key_alias or any other user_api_key_* field and have a guardrail attribute the call, or route policy, to another key. Every LLM route strips those keys and writes the authenticated identity instead The endpoint now does the same: caller fields keep their non-identity keys, and the user_api_key_* fields come from the proxy's own sanitized metadata for the authenticated key. A caller that sends no metadata still gets that identity forwarded, matching the shape guardrails see on every other route --- .../proxy/guardrails/guardrail_endpoints.py | 25 ++++++- .../guardrails/test_guardrail_endpoints.py | 66 ++++++++++++++++++- 2 files changed, 87 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index afb9997f2e6..74cf55c2f6e 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, Union, from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException, Request -from pydantic import BaseModel, ValidationError +from pydantic import BaseModel, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH @@ -2215,6 +2215,26 @@ async def test_custom_code_guardrail( ) +_GUARDRAIL_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object]) + + +def _metadata_fields(value: object) -> Mapping[str, object]: + try: + return _GUARDRAIL_METADATA_ADAPTER.validate_python(value) + except ValidationError: + return {} # mutable-ok: empty fallback for the request_data dict contract + + +def _guardrail_request_metadata(caller: object, proxy: object) -> Mapping[str, object]: + caller_fields: Final = tuple( + (key, value) for key, value in _metadata_fields(caller).items() if not key.startswith("user_api_key_") + ) + key_identity: Final = tuple( + (key, value) for key, value in _metadata_fields(proxy).items() if key.startswith("user_api_key_") + ) + return dict((*caller_fields, *key_identity)) # mutable-ok: request_data is the apply_guardrail dict contract + + def _resolve_guardrail_input_type(active_guardrail: CustomGuardrail, input_type: str) -> Literal["request", "response"]: """Return the effective input_type, auto-upgrading to 'response' for post_call guardrails.""" if input_type == "request": @@ -2377,9 +2397,10 @@ async def apply_guardrail( if litellm_logging_obj is not None: _patch_logging_obj_for_guardrail(litellm_logging_obj, request) + metadata: Final = _guardrail_request_metadata(request.metadata, data.get("metadata")) 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": metadata} if request.metadata is not None or metadata else {}), } _input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type) guardrailed_inputs: Final = await active_guardrail.apply_guardrail( diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 530f8ffd854..ed915b0121c 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -1612,7 +1612,7 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker): } -def _patch_apply_guardrail_env(mocker, guardrail_result): +def _patch_apply_guardrail_env(mocker, guardrail_result, processed_data=None): mock_guardrail = mocker.Mock() mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result) @@ -1627,7 +1627,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", @@ -1698,6 +1698,68 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker): ) +@pytest.mark.asyncio +async def test_apply_guardrail_replaces_caller_identity_with_the_authenticated_key(mocker): + """Identity fields come from the proxy's own sanitized metadata, never from + the caller, so a request cannot impersonate another key or probe its policy.""" + mock_guardrail = _patch_apply_guardrail_env( + mocker, + {"texts": ["ok"]}, + processed_data={ + "guardrail_name": "test-guardrail", + "metadata": { + "route": "/apply_guardrail", + "user_api_key_alias": "billing-app", + "user_api_key_user_id": "u-1", + }, + }, + ) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={ + "forbidden_topics": ["tax"], + "user_api_key_alias": "someone-else", + "user_api_key_user_email": "victim@example.com", + }, + ) + await apply_guardrail( + fastapi_request=mocker.Mock(), + request=request, + user_api_key_dict=UserAPIKeyAuth(key_alias="billing-app", user_id="u-1"), + ) + + forwarded = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert forwarded == { + "forbidden_topics": ["tax"], + "user_api_key_alias": "billing-app", + "user_api_key_user_id": "u-1", + } + + +@pytest.mark.asyncio +async def test_apply_guardrail_attaches_key_identity_when_caller_sends_no_metadata(mocker): + """A caller that sends no metadata still gets the authenticated identity + forwarded, the same shape every LLM route gives a guardrail.""" + mock_guardrail = _patch_apply_guardrail_env( + mocker, + {"texts": ["ok"]}, + processed_data={"guardrail_name": "test-guardrail", "metadata": {"user_api_key_alias": "billing-app"}}, + ) + + request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello") + await apply_guardrail( + fastapi_request=mocker.Mock(), + request=request, + user_api_key_dict=UserAPIKeyAuth(key_alias="billing-app"), + ) + + assert mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] == { + "metadata": {"user_api_key_alias": "billing-app"} + } + + @pytest.mark.asyncio async def test_apply_guardrail_omits_metadata_when_not_sent(mocker): """Without metadata, request_data stays empty (backward-compatible)."""