diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 20efbe06ecc..ae05b9ad65b 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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( diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index 440389a79ae..8927856c855 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py index e3e4d236826..14338a58248 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py @@ -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. diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 45f5afef1bc..9160a9924fb 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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;