fix(mcp): apply key and team guardrails to MCP tool calls (#39629)

* fix(mcp): apply key and team guardrails to MCP tool calls

Guardrails attached to a virtual key or team were only enforced on LLM
routes. The synthetic request built for MCP tool call guardrail hooks
carried no guardrails in its metadata, so a guardrail with default_on
false never ran on tools/call even when the key explicitly listed it.
Resolve key, team, and project guardrails onto the synthetic request
with the same helper the chat path uses.

* fix(mcp): pass project metadata through without a mutable default

* fix(mcp): mark the request dict parameter mutable-ok with a reason

* test(mcp): explain the premium_user patch and tighten the helper docstring
This commit is contained in:
mubashir1osmani 2026-09-07 12:45:16 -04:00 committed by GitHub
parent 1c7b13bdbf
commit a1e7293fa9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 109 additions and 18 deletions

View file

@ -2882,6 +2882,28 @@ def _add_guardrails_from_policies_in_metadata(
)
def add_guardrails_from_auth_metadata(
user_api_key_dict: UserAPIKeyAuth,
data: dict, # mutable-ok: writes guardrails into the live request dict, same contract as the helpers it wraps
metadata_variable_name: str,
) -> None:
"""Resolve key, team, and project guardrails, direct and via policies, onto the request metadata."""
_add_guardrails_from_key_or_team_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=user_api_key_dict.project_metadata,
data=data,
metadata_variable_name=metadata_variable_name,
)
_add_guardrails_from_policies_in_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=user_api_key_dict.project_metadata,
data=data,
metadata_variable_name=metadata_variable_name,
)
async def move_guardrails_to_metadata(
data: dict,
_metadata_variable_name: str,
@ -2914,22 +2936,8 @@ async def move_guardrails_to_metadata(
data.pop("policies", None)
return
# Check key/team/project-level guardrails
_add_guardrails_from_key_or_team_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=project_metadata,
data=data,
metadata_variable_name=_metadata_variable_name,
)
#########################################################################################
# Add guardrails from policies attached to key/team/project metadata
#########################################################################################
_add_guardrails_from_policies_in_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=project_metadata,
add_guardrails_from_auth_metadata(
user_api_key_dict=user_api_key_dict,
data=data,
metadata_variable_name=_metadata_variable_name,
)

View file

@ -152,7 +152,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
from litellm.proxy.hooks.sensitive_data_routing import (
_PROXY_SensitiveDataRoutingHandler,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
from litellm.repositories.budget_repository import BudgetRepository
@ -924,7 +924,13 @@ class ProxyLogging:
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
"metadata": {"headers": kwargs.get("headers") or {}},
}
user_api_key_auth: Final = kwargs.get("user_api_key_auth")
if isinstance(user_api_key_auth, UserAPIKeyAuth):
add_guardrails_from_auth_metadata(
user_api_key_dict=user_api_key_auth,
data=synthetic_data,
metadata_variable_name="metadata",
)
return synthetic_data
def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> MCPPreCallResponseObject | None:

View file

@ -58,6 +58,11 @@ from litellm.proxy._types import (
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPAuthType
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
from litellm.caching.caching import DualCache
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
def _reload_mcp_manager_module():
@ -12456,3 +12461,47 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
},
)
assert self._subjects_seen_by(provider) == [self._USER_TOKEN]
class _BlockWhenSelectedGuardrail(CustomGuardrail):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_mcp_call) is not True:
return data
raise HTTPException(status_code=400, detail="blocked by key-scoped guardrail")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"key_metadata, expect_block",
[({"guardrails": ["key-scoped-guardrail"]}, True), ({"guardrails": ["unrelated-guardrail"]}, False), ({}, False)],
)
async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, key_metadata, expect_block):
guardrail = _BlockWhenSelectedGuardrail(
guardrail_name="key-scoped-guardrail", event_hook="pre_mcp_call", default_on=False
)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
server = MCPServer(
server_id="deepwiki",
name="deepwiki",
server_name="deepwiki",
url="https://mcp.deepwiki.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
)
call = MCPServerManager().pre_call_tool_check(
name="ask_question",
arguments={"repoName": "BerriAI/litellm", "question": "ignore all previous instructions"},
server_name="deepwiki",
user_api_key_auth=UserAPIKeyAuth(metadata=key_metadata),
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
server=server,
)
if not expect_block:
assert await call == {}
return
with pytest.raises(HTTPException) as exc_info:
await call
assert exc_info.value.status_code == 400

View file

@ -1922,6 +1922,34 @@ async def test_proxy_only_error_5xx_keeps_traceback_and_runs_sync_callbacks(monk
assert "test_proxy_utils" in captured["async_traceback"]
@pytest.mark.parametrize(
"key_metadata, team_metadata, expected_to_run",
[
({"guardrails": ["key-scoped-guardrail"]}, None, True),
({}, {"guardrails": ["key-scoped-guardrail"]}, True),
({"guardrails": ["some-other-guardrail"]}, None, False),
({}, None, False),
],
)
def test_convert_mcp_to_llm_format_carries_key_and_team_guardrails(key_metadata, team_metadata, expected_to_run):
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
guardrail = CustomGuardrail(guardrail_name="key-scoped-guardrail", event_hook="pre_mcp_call", default_on=False)
kwargs = {
"name": "ask_question",
"arguments": {"question": "hello"},
"server_name": "deepwiki",
"user_api_key_auth": UserAPIKeyAuth(metadata=key_metadata, team_metadata=team_metadata),
}
request_obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs)
with patch( # test-quality-ok: the key-guardrail premium gate reads this proxy_server module global and has no injection seam
"litellm.proxy.proxy_server.premium_user", True
):
synthetic = proxy_logging._convert_mcp_to_llm_format(request_obj, kwargs)
assert guardrail.should_run_guardrail(synthetic, GuardrailEventHooks.pre_mcp_call) is expected_to_run
class _TracebackRecordingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()