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
This commit is contained in:
albertbausili 2026-09-14 09:03:15 +02:00
parent 976d1466c7
commit d7b6d3c904
2 changed files with 87 additions and 4 deletions

View file

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

View file

@ -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)."""