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.
This commit is contained in:
karpanin 2026-09-16 02:53:06 +02:00
parent 226b1e1bb9
commit 5561a7d7a9
2 changed files with 151 additions and 6 deletions

View file

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

View file

@ -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="<MASKED>", raises=None):
super().__init__(guardrail_name="mcp-output-guardrail", event_hook=event_hook, default_on=True)
def __init__(self, event_hook, masked_text="<MASKED>", 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] == ["<MASKED>"]
@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] == ["<MASKED>"]
@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] == ["<MASKED>"]
@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