mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
976d1466c7
commit
d7b6d3c904
2 changed files with 87 additions and 4 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue