mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
226b1e1bb9
commit
5561a7d7a9
2 changed files with 151 additions and 6 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue