mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): source apply_guardrail identity from authenticated metadata
POST /guardrails/apply_guardrail built the guardrail's request_data from the request body alone, so the proxy-injected user_api_key_* fields never reached the guardrail. Key targeting saw no calling key and silently skipped scanning whenever apply_to_aliases was set, and a caller could name any alias to slip into an allowlist, dodge an exclusion, or misattribute a SIEM record The metadata common_processing_pre_call_logic produced is now merged over the body's, so identity comes from the key that authenticated while the caller's own fields, grounding documents included, still reach the guardrail Also tightens xecguard's nested key-context lift, which returned the first dict under litellm_params without checking for identity fields, so an empty metadata sibling could hide the identity that lived on litellm_metadata
This commit is contained in:
parent
b7af384dce
commit
1ac99623b7
4 changed files with 77 additions and 8 deletions
|
|
@ -2320,9 +2320,18 @@ async def apply_guardrail(
|
|||
if litellm_logging_obj is not None:
|
||||
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
|
||||
|
||||
# The proxy-injected metadata is merged last so it wins: the body's is
|
||||
# caller-controlled, and a guardrail reading it would let a caller name a
|
||||
# different virtual key than the one that authenticated.
|
||||
merged_metadata: Final = {
|
||||
name: value
|
||||
for source in (request.metadata, data.get("metadata"))
|
||||
if isinstance(source, dict)
|
||||
for name, value in source.items()
|
||||
}
|
||||
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": merged_metadata,
|
||||
}
|
||||
_input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type)
|
||||
guardrailed_inputs: Final = await active_guardrail.apply_guardrail(
|
||||
|
|
|
|||
|
|
@ -323,7 +323,7 @@ class XecGuardGuardrail(CustomGuardrail):
|
|||
if isinstance(nested, dict):
|
||||
for meta_key in ("metadata", "litellm_metadata"):
|
||||
md = nested.get(meta_key)
|
||||
if isinstance(md, dict):
|
||||
if isinstance(md, dict) and any(field in md for field in cls._KEY_IDENTITY_FIELDS):
|
||||
return {meta_key: md} # mutable-ok: lifts nested metadata to the readers' shape
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -2314,6 +2314,20 @@ class TestXecGuardLoggingHookKeyTargeting:
|
|||
)
|
||||
assert ctx["metadata"]["user_api_key_alias"] == "nested"
|
||||
|
||||
def test_key_context_nested_skips_metadata_without_identity(self):
|
||||
# logging path: an empty litellm_params.metadata must not shadow the
|
||||
# sibling key that actually carries the identity.
|
||||
gr = _extension_guardrail()
|
||||
ctx = gr._key_context(
|
||||
{
|
||||
"litellm_params": {
|
||||
"metadata": {},
|
||||
"litellm_metadata": {"user_api_key_alias": "sibling"},
|
||||
}
|
||||
}
|
||||
)
|
||||
assert gr._calling_key_identity(ctx) == ("sibling", None)
|
||||
|
||||
def test_key_context_returns_proxy_shape_untouched(self):
|
||||
# pre/during/post_call: reshaping to a single key would drop the other
|
||||
# location _calling_key_identity also reads.
|
||||
|
|
|
|||
|
|
@ -1497,7 +1497,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, proxy_metadata=None):
|
||||
mock_guardrail = mocker.Mock()
|
||||
mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result)
|
||||
|
||||
|
|
@ -1511,8 +1511,12 @@ def _patch_apply_guardrail_env(mocker, guardrail_result):
|
|||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
mock_processor = mocker.Mock()
|
||||
processed_data = {
|
||||
"guardrail_name": "test-guardrail",
|
||||
**({"metadata": proxy_metadata} if proxy_metadata is not None else {}),
|
||||
}
|
||||
mock_processor.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
|
||||
return_value=(processed_data, mock_logging_obj)
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
|
|
@ -1584,9 +1588,17 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
|
||||
"""Without metadata, request_data stays empty (backward-compatible)."""
|
||||
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
|
||||
async def test_apply_guardrail_forwards_proxy_identity_when_body_has_no_metadata(mocker):
|
||||
"""Guardrails that target specific virtual keys need the authenticated key
|
||||
even when the body carries no metadata of its own."""
|
||||
proxy_metadata = {
|
||||
"route": "/apply_guardrail",
|
||||
"user_api_key_alias": "prod-key",
|
||||
"user_api_key_hash": "hash-abc",
|
||||
}
|
||||
mock_guardrail = _patch_apply_guardrail_env(
|
||||
mocker, {"texts": ["ok"]}, proxy_metadata=proxy_metadata
|
||||
)
|
||||
|
||||
request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
|
||||
await apply_guardrail(
|
||||
|
|
@ -1597,11 +1609,45 @@ async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
|
|||
|
||||
mock_guardrail.apply_guardrail.assert_awaited_once_with(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data={},
|
||||
request_data={"metadata": proxy_metadata},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_proxy_identity_overrides_caller_metadata(mocker):
|
||||
"""A caller must not be able to name a virtual key other than the one that
|
||||
authenticated, while the body's non-identity fields still reach the guardrail."""
|
||||
mock_guardrail = _patch_apply_guardrail_env(
|
||||
mocker,
|
||||
{"texts": ["ok"]},
|
||||
proxy_metadata={
|
||||
"user_api_key_alias": "authenticated-key",
|
||||
"user_api_key_hash": "hash-real",
|
||||
},
|
||||
)
|
||||
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="test-guardrail",
|
||||
text="hello",
|
||||
metadata={
|
||||
"user_api_key_alias": "spoofed-key",
|
||||
"user_api_key_hash": "hash-spoofed",
|
||||
"forbidden_topics": ["tax"],
|
||||
},
|
||||
)
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
forwarded = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
|
||||
assert forwarded["user_api_key_alias"] == "authenticated-key"
|
||||
assert forwarded["user_api_key_hash"] == "hash-real"
|
||||
assert forwarded["forbidden_topics"] == ["tax"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker):
|
||||
"""Explicitly-sent empty messages/metadata must be forwarded, not dropped;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue