From 5561a7d7a9966ab136b7b3b51f92ae20bca3325b Mon Sep 17 00:00:00 2001 From: karpanin Date: Wed, 16 Sep 2026 02:53:06 +0200 Subject: [PATCH] fix(guardrails): run post_mcp_call guardrails attached by key, team or policy Every MCP call site hands post_mcp_call_hook the request's Logging.model_call_details, which nests the request metadata under litellm_params. should_run_guardrail only reads the top level, so the guardrail list a key, team or policy resolved onto the request was invisible there: a post_mcp_call guardrail only ever ran when it was default_on, and opted_out_global_guardrails never took one off. Lift the metadata buckets out of litellm_params before gating, so MCP tool results get scanned and masked by exactly the guardrails the request resolved, the same way the inbound pre_mcp_call hook already does. --- litellm/proxy/utils.py | 26 +++- tests/test_litellm/proxy/test_proxy_utils.py | 131 ++++++++++++++++++- 2 files changed, 151 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 479bd0a55af..893425f93ba 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -936,6 +936,26 @@ def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, return MappingProxyType({key: value for key, value in _entries if value is not None}) +def _mcp_guardrail_gate_data(request_data: Mapping[str, object]) -> Mapping[str, object]: + """``request_data`` reshaped so guardrail gating can read the request's metadata. + + Every MCP call site hands ``post_mcp_call_hook`` the request's + ``Logging.model_call_details``, which nests the request metadata under + ``litellm_params``. ``CustomGuardrail.should_run_guardrail`` only reads the + top level, so without this lift the guardrail list a key, team or policy + resolved onto the request is invisible and only ``default_on`` guardrails + ever run, and ``opted_out_global_guardrails`` never takes one off. + """ + params: Final = request_data.get("litellm_params") + if not isinstance(params, Mapping): + return request_data + nested: Final = cast("Mapping[str, object]", params) # cast-ok: litellm_params is always a str-keyed dict + lifted: Final = MappingProxyType({key: nested[key] for key in ("metadata", "litellm_metadata") if key in nested}) + if not lifted: + return request_data + return MappingProxyType({**request_data, **lifted}) + + @dataclass(frozen=True) class _CallbackCapabilities: """Cached per-hook capability flags derived from ``litellm.callbacks``. @@ -3345,15 +3365,13 @@ class ProxyLogging: verbose_proxy_logger.debug("MCP guardrail translation handler unavailable; skipping post_mcp_call hook") return response + gate_data: Final = _mcp_guardrail_gate_data(request_data) for callback in caps.resolved_callbacks: if not isinstance(callback, CustomGuardrail): continue if "apply_guardrail" not in type(callback).__dict__ or callback.use_native_lifecycle_hooks: continue - if ( - callback.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_mcp_call) - is not True - ): + if callback.should_run_guardrail(data=gate_data, event_type=GuardrailEventHooks.post_mcp_call) is not True: continue response = await self._run_guardrail_with_metrics( callback, diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 94ccc2762c5..d107b90040b 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -1523,8 +1523,8 @@ class TestSendEmailStartTls: class _RecordingMCPGuardrail(CustomGuardrail): """Unified guardrail that masks every text it is handed.""" - def __init__(self, event_hook, masked_text="", raises=None): - super().__init__(guardrail_name="mcp-output-guardrail", event_hook=event_hook, default_on=True) + def __init__(self, event_hook, masked_text="", raises=None, default_on=True): + super().__init__(guardrail_name="mcp-output-guardrail", event_hook=event_hook, default_on=default_on) self.masked_text = masked_text self.raises = raises self.call_count = 0 @@ -1656,6 +1656,133 @@ async def test_post_mcp_call_hook_propagates_guardrail_block(restore_callbacks): ) +def _mcp_logging_payload(metadata, metadata_key="metadata"): + """The request payload every MCP call site hands ``post_mcp_call_hook``. + + ``Logging.model_call_details`` nests the request metadata under + ``litellm_params``; see ``litellm.utils.function_setup``. The bucket is + named ``litellm_metadata`` on the routes that reserve ``metadata`` for the + provider, which is how MCP tools invoked from /responses arrive here. + """ + return { + "model": "MCP: echo", + "call_type": "call_mcp_tool", + "litellm_params": {"api_base": "", metadata_key: metadata}, + } + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_runs_guardrail_requested_through_request_metadata(restore_callbacks): + """A guardrail attached by key, team or policy (not default_on) must mask MCP tool output. + + The MCP paths pass the logging payload, where the resolved guardrail list + sits under litellm_params.metadata, so gating on the top level alone left + masking working only for default_on guardrails. + """ + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call, default_on=False) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=_mcp_logging_payload({"guardrails": [guardrail.guardrail_name]}), + user_api_key_dict=None, + ) + + assert guardrail.call_count == 1 + assert [item.text for item in returned.content] == [""] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_skips_guardrail_absent_from_request_metadata(restore_callbacks): + """A non-default_on guardrail the request never asked for must leave tool output alone.""" + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call, default_on=False) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=_mcp_logging_payload({"guardrails": ["some-other-guardrail"]}), + user_api_key_dict=None, + ) + + assert guardrail.call_count == 0 + assert [item.text for item in returned.content] == ["jane@example.com"] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_runs_guardrail_requested_through_litellm_metadata(restore_callbacks): + """The same lift must work for routes whose request metadata bucket is litellm_metadata.""" + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call, default_on=False) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=_mcp_logging_payload( + {"guardrails": [guardrail.guardrail_name]}, metadata_key="litellm_metadata" + ), + user_api_key_dict=None, + ) + + assert guardrail.call_count == 1 + assert [item.text for item in returned.content] == [""] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_honors_opt_out_in_request_metadata(restore_callbacks): + """A key that opted out of a global guardrail must not have MCP tool output scanned by it.""" + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=_mcp_logging_payload( + {"user_api_key_metadata": {"opted_out_global_guardrails": [guardrail.guardrail_name]}} + ), + user_api_key_dict=None, + ) + + assert guardrail.call_count == 0 + assert [item.text for item in returned.content] == ["jane@example.com"] + + +@pytest.mark.asyncio +async def test_post_mcp_call_hook_runs_default_on_guardrail_when_payload_has_no_metadata_bucket(restore_callbacks): + """A logging payload carrying no metadata bucket must leave the Default On path working. + + There is nothing to lift here, and that no-op must not change the verdict. + """ + from mcp.types import CallToolResult, TextContent + + guardrail = _RecordingMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call) + litellm.callbacks = [guardrail] + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False) + + returned = await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data={"model": "MCP: echo", "litellm_params": {"api_base": ""}}, + user_api_key_dict=None, + ) + + assert guardrail.call_count == 1 + assert [item.text for item in returned.content] == [""] + + @pytest.mark.asyncio async def test_prisma_health_check_failure_names_itself_at_operator_visible_level(caplog): """A failing DB health check has to name the check that failed, at a level