From 0e27a523e040a83ad3842b6f7fae16b101756412 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 14 Sep 2026 22:04:02 +0000 Subject: [PATCH] fix(mcp): probe Entra at connect, challenge only default-on Agent 365 guardrails, key listed tools by signer Connect-time sign-in challenge now asks Entra to exchange the presented assertion instead of only checking its compact-JWS shape, so expired, wrong-audience or forged bearers get the RFC 9728 challenge while gateway credential and provider failures still surface on the tool call. Only default_on guardrails the caller has not opted out of advertise or challenge, since anonymous metadata cannot see key-selected guardrails. Servers whose Authorization is minted per caller by MCPJWTSigner list tools per caller instead of sharing one cache slot Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 13 +++ .../proxy/_experimental/mcp_server/server.py | 14 +-- .../guardrail_hooks/agent_365/agent_365.py | 63 ++++++++---- .../mcp_server/test_mcp_server.py | 95 ++++++++++++++++++- .../mcp_server/test_mcp_server_manager.py | 32 +++++++ .../mcp_server/test_openapi_tool_auth.py | 2 + .../guardrail_hooks/test_agent_365.py | 22 ++++- 7 files changed, 212 insertions(+), 29 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 319735a8c08..00f19bf964e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4519,8 +4519,21 @@ class MCPServerManager: or self._references_per_user_env_var(server) or server.delegate_auth_to_upstream or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) + or self._signs_caller_identity_upstream(server) ) + @staticmethod + def _signs_caller_identity_upstream(server: MCPServer) -> bool: + """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may + tailor its catalog to the caller even though the server itself is configured as shared.""" + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server + get_mcp_jwt_signer, + ) + + if get_mcp_jwt_signer() is None: + return False + return not any(k.lower() == "authorization" for k in (server.static_headers or {})) + def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: """Key the listed-tool cache by every request input that can change the upstream catalog. diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 3b471033235..ca12d9e9937 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -82,8 +82,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.guardrails.guardrail_hooks.agent_365.agent_365 import ( - agent_365_authorization_servers, - agent_365_subject_token_present, + agent_365_sign_in_required, ) from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, @@ -4160,9 +4159,11 @@ if MCP_AVAILABLE: # then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the # header lost, so the discovery flow needs this pre-emptive challenge. Servers gated by an # Agent 365 guardrail (OBO to the evaluate API) get the same challenge, also when the only - # bearer is the LiteLLM key itself, which admits the caller but is not an exchangeable subject. - # Only on the server's own route: the per-server metadata ``resource`` must equal the URL the - # client connected to (RFC 9728 3.3), which aggregate ``/mcp`` and multi-server connects never do. + # bearer is the LiteLLM key itself, which admits the caller but is not an exchangeable subject, + # and when Entra refuses the presented assertion (expired, wrong audience), so the client + # signs in again instead of failing every tool call. Only on the server's own route: the + # per-server metadata ``resource`` must equal the URL the client connected to (RFC 9728 3.3), + # which aggregate ``/mcp`` and multi-server connects never do. granted_single_server = server is not None and await _key_granted_single_server( server, mcp_servers, user_api_key_auth, client_ip ) @@ -4171,8 +4172,7 @@ if MCP_AVAILABLE: or ( granted_single_server and tuple(_get_mcp_servers_in_path(get_route_relative_request_path(scope)) or ()) == (server_name,) - and not agent_365_subject_token_present(oauth2_headers) - and agent_365_authorization_servers(server, user_api_key_auth) + and await agent_365_sign_in_required(server, user_api_key_auth, oauth2_headers) ) ): from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index b134e1e7ff6..cb6ff9a5355 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -35,7 +35,6 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy.litellm_pre_call_utils import add_guardrails_from_auth_metadata from litellm.types.guardrails import GuardrailEventHooks from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -440,6 +439,18 @@ class Agent365Guardrail(CustomGuardrail): return call_id return str(uuid.uuid4()) + async def exchange_rejects_subject(self, assertion: str) -> bool: + """Whether Entra refuses ``assertion`` as the On-Behalf-Of subject (expired, wrong audience, bad + signature). Gateway credential rejections and endpoint failures answer ``False``: the caller cannot fix + those by signing in again, so the tool call reports them.""" + try: + await self._get_obo_token(assertion) + except Agent365TokenExchangeError as exc: + return exc.error_code not in _GATEWAY_OWNED_TOKEN_ERRORS + except (Agent365ThrottledError, Agent365MalformedResponseError, httpx.HTTPError, LitellmTimeout, TimeoutError): + return False + return False + async def _get_obo_token(self, assertion: str) -> str: cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest() now: Final = time.time() @@ -639,29 +650,32 @@ class Agent365Guardrail(CustomGuardrail): def _applies_to_caller(guardrail: Agent365Guardrail, user_api_key_auth: "UserAPIKeyAuth") -> bool: - probe: Final[dict[str, object]] = {"metadata": {}} # mutable-ok: filled in place by the key resolver - add_guardrails_from_auth_metadata( - user_api_key_dict=user_api_key_auth, data=probe, metadata_variable_name="metadata" - ) - return guardrail.should_run_guardrail(data=probe, event_type=GuardrailEventHooks.pre_mcp_call) + probe: Final[Mapping[str, object]] = { + "metadata": { + "user_api_key_metadata": user_api_key_auth.metadata, + "user_api_key_team_metadata": user_api_key_auth.team_metadata, + } + } + return guardrail.should_run_guardrail(data=dict(probe), event_type=GuardrailEventHooks.pre_mcp_call) def _applicable_guardrails( server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None" ) -> tuple[Agent365Guardrail, ...]: - """Agent 365 guardrails that gate ``server`` for this caller: the ``default_on`` ones for the anonymous - discovery fetch, otherwise those the caller's key, team, or policies select. Empty when the gateway - does not own sign-in for the server.""" + """Agent 365 guardrails whose sign-in the gateway advertises for ``server``: the ``default_on`` ones, minus + those the caller's key or team opted out of once the caller is known. A guardrail only a key or policy + selects still enforces at the tool call but never challenges, since the anonymous metadata fetch that + follows a challenge cannot see which key selected it and would advertise the wrong issuer.""" if server.auth_type == MCPAuth.oauth2 or not server.advertises_gateway_authorization_server: return () - registered: Final = tuple( + advertised: Final = tuple( callback for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(Agent365Guardrail) - if isinstance(callback, Agent365Guardrail) + if isinstance(callback, Agent365Guardrail) and callback.default_on ) if user_api_key_auth is None: - return tuple(g for g in registered if g.default_on) - return tuple(g for g in registered if _applies_to_caller(g, user_api_key_auth)) + return advertised + return tuple(g for g in advertised if _applies_to_caller(g, user_api_key_auth)) def entra_assertion(value: object) -> str | None: @@ -670,12 +684,29 @@ def entra_assertion(value: object) -> str | None: return value if isinstance(value, str) and value.count(".") == 2 else None -def agent_365_subject_token_present(oauth2_headers: Mapping[str, str] | None) -> bool: - """Whether the request's ``Authorization`` carries an Entra assertion the guardrail can exchange.""" +def _presented_assertion(oauth2_headers: Mapping[str, str] | None) -> str | None: authorization: Final = oauth2_headers.get("Authorization", "") if oauth2_headers else "" if not authorization.lower().startswith("bearer "): + return None + return entra_assertion(authorization[len("bearer ") :].strip()) + + +async def agent_365_sign_in_required( + server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", oauth2_headers: Mapping[str, str] | None +) -> bool: + """Whether the connect must answer with the RFC 9728 sign-in challenge: an Agent 365 guardrail gates + ``server`` for this caller and the request carries no Entra assertion, or one Entra will not exchange. + Decided at connect because a tool call's JSON-RPC error cannot carry ``WWW-Authenticate``.""" + guardrails: Final = _applicable_guardrails(server, user_api_key_auth) + if not guardrails: return False - return entra_assertion(authorization[len("bearer ") :].strip()) is not None + assertion: Final = _presented_assertion(oauth2_headers) + if assertion is None: + return True + for guardrail in guardrails: + if await guardrail.exchange_rejects_subject(assertion): + return True + return False def agent_365_authorization_servers(server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None") -> tuple[str, ...]: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 1817f101758..511433837c0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -6,6 +6,7 @@ from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -8915,6 +8916,15 @@ class TestAgent365ChallengeAtConnect: WWW-Authenticate header survives, instead of only inside the tools/call JSON-RPC error.""" GATEWAY_SCOPE = "api://gateway-app/access_as_user" + ENTRA_BEARER = {"Authorization": "Bearer eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1LTEifQ.c2ln"} + + @staticmethod + def _entra_response(status_code: int, body: dict[str, object]) -> httpx.Response: + return httpx.Response( + status_code=status_code, + json=body, + request=httpx.Request("POST", "https://login.microsoftonline.com/tenant-abc/oauth2/v2.0/token"), + ) def _server(self, scopes: list[str] | None) -> MCPServer: return MCPServer( @@ -8934,12 +8944,14 @@ class TestAgent365ChallengeAtConnect: import litellm from litellm.proxy.guardrails.guardrail_hooks.agent_365 import Agent365Guardrail + handler = AsyncMock() + handler.post.return_value = self._entra_response(200, {"access_token": "obo-token", "expires_in": 3599}) guardrail = Agent365Guardrail( guardrail_name="agent-365-guard", tenant_id="tenant-abc", client_id="client-xyz", client_secret="secret-123", - async_handler=AsyncMock(), + async_handler=handler, event_hook="pre_mcp_call", default_on=True, ) @@ -8958,6 +8970,7 @@ class TestAgent365ChallengeAtConnect: path: str = "/mcp/tools", mount_scope: dict[str, str] | None = None, granted: bool = True, + key_metadata: dict[str, object] | None = None, ) -> HTTPException | None: from litellm.proxy._experimental.mcp_server import server as server_module @@ -8983,7 +8996,9 @@ class TestAgent365ChallengeAtConnect: mcp_servers=["tools"], oauth2_headers=oauth2_headers, mcp_server_auth_headers=None, - user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + user_api_key_auth=UserAPIKeyAuth( + api_key="sk-litellm-virtual-key", user_id="u-1", metadata=key_metadata or {} + ), client_ip=None, ) except HTTPException as challenge: @@ -9048,8 +9063,80 @@ class TestAgent365ChallengeAtConnect: @pytest.mark.asyncio async def test_entra_assertion_present_connects(self, agent_365_guardrail): - bearer = {"Authorization": "Bearer eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1LTEifQ.c2ln"} - assert await self._connect(self._server([self.GATEWAY_SCOPE]), bearer) is None + assert await self._connect(self._server([self.GATEWAY_SCOPE]), self.ENTRA_BEARER) is None + + @pytest.mark.asyncio + async def test_assertion_entra_refuses_to_exchange_is_challenged(self, agent_365_guardrail): + """An expired, wrong-audience, or forged Entra token looks like a valid one. Only the OBO exchange + can tell, and its verdict must arrive at connect, where WWW-Authenticate reaches the client, + rather than inside every tools/call JSON-RPC error.""" + agent_365_guardrail.async_handler.post.return_value = self._entra_response( + 400, {"error": "invalid_grant", "error_description": "AADSTS700084: The refresh token was issued..."} + ) + + challenge = await self._connect(self._server([self.GATEWAY_SCOPE]), self.ENTRA_BEARER) + + assert challenge is not None and challenge.status_code == 401 + www_authenticate = (challenge.headers or {}).get("WWW-Authenticate", "") + assert 'error="invalid_token"' in www_authenticate + assert ( + 'resource_metadata="https://gw.example.com/.well-known/oauth-protected-resource/mcp/tools"' + in www_authenticate + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "entra_outcome", + [ + {"return_value": _entra_response(401, {"error": "invalid_client"})}, + {"return_value": _entra_response(503, {"error": "temporarily_unavailable"})}, + {"side_effect": httpx.ConnectError("dns")}, + ], + ids=["gateway-credentials-rejected", "entra-5xx", "entra-unreachable"], + ) + async def test_gateway_side_exchange_failures_do_not_send_the_caller_to_sign_in( + self, agent_365_guardrail, entra_outcome + ): + """Signing in again cannot fix the gateway's own client secret or an Entra outage, so those stay + with the tool call, which reports them as the guardrail's unavailable path.""" + agent_365_guardrail.async_handler.post = AsyncMock(**entra_outcome) + + assert await self._connect(self._server([self.GATEWAY_SCOPE]), self.ENTRA_BEARER) is None + + @pytest.mark.asyncio + async def test_key_selected_default_off_guardrail_never_challenges(self): + """The anonymous metadata fetch that follows a challenge cannot see which key selected the guardrail + and would advertise the gateway issuer, so a challenge here would send the client to the wrong IdP. + The guardrail still enforces at tools/call.""" + import litellm + from litellm.proxy.guardrails.guardrail_hooks.agent_365 import Agent365Guardrail + + guardrail = Agent365Guardrail( + guardrail_name="agent-365-guard", + tenant_id="tenant-abc", + client_id="client-xyz", + client_secret="secret-123", + async_handler=AsyncMock(), + event_hook="pre_mcp_call", + default_on=False, + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with patch( # test-quality-ok: key-selected guardrails read the proxy server premium global, no injection seam + "litellm.proxy.proxy_server.premium_user", True + ): + assert ( + await self._connect( + self._server([self.GATEWAY_SCOPE]), + None, + key_metadata={"guardrails": ["agent-365-guard"]}, + ) + is None + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) @pytest.mark.asyncio async def test_litellm_key_in_authorization_is_still_challenged(self, agent_365_guardrail): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f94d1dfa53b..bf670f62c4c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6827,6 +6827,38 @@ class TestMCPServerManager: listed = manager.get_listed_tool(server, "turn", other) assert listed is not None and listed.description == "everyone" + @pytest.mark.parametrize( + ("signer", "static_headers", "shared"), + [ + pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), + pytest.param(None, None, True, id="no-signer-stays-shared"), + ], + ) + def test_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared): + """MCPJWTSigner hands upstream a JWT naming the caller on an otherwise shared ``auth_type: none`` + server, so the upstream may tailor the catalog and the cache must not hand one caller another's.""" + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers + ) + alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="hashed-alice")) + bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", api_key="hashed-bob")) + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=signer, + ): + manager._create_prefixed_tools( + [MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice + ) + for_bob = manager.get_listed_tool(server, "srv-turn", bob) + + if shared: + assert for_bob is not None and for_bob.description == "alice view" + else: + assert for_bob is None + @pytest.mark.asyncio async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self): """Interleaved callers on a forwarded-header server: the hook must see the caller's own catalog.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index c534cf73219..67894d55e80 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -40,6 +40,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): fake_server.server_name = "openapi-petstore" fake_server.alias = None fake_server.short_prefix = None + fake_server.tool_name_to_description = None fake_tool = MagicMock() fake_tool.name = "list_pets" @@ -123,6 +124,7 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): fake_server.server_name = "openapi-petstore" fake_server.alias = None fake_server.short_prefix = None + fake_server.tool_name_to_description = None fake_tool = MagicMock() fake_tool.name = "delete_pet" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index a679c56f561..6466bd4549d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1094,7 +1094,9 @@ class TestAgent365AuthorizationServers: litellm.callbacks, twin, require_self=False ) - def test_key_selected_guardrail_challenges_only_that_key(self): + def test_key_selected_guardrail_never_advertises_sign_in(self): + """The challenge a guarded key would get and the anonymous metadata fetch that follows it must name + the same issuer. The anonymous fetch cannot see the key, so neither side advertises Entra.""" guardrail: Final = _make_guardrail(FakeHandler([])) guardrail.default_on = False litellm.logging_callback_manager.add_litellm_callback(guardrail) @@ -1108,10 +1110,26 @@ class TestAgent365AuthorizationServers: "litellm.proxy.proxy_server.premium_user", True ): assert agent_365_authorization_servers(server, plain_key) == () - assert agent_365_authorization_servers(server, guarded_key) == (ENTRA_ISSUER,) + assert agent_365_authorization_servers(server, guarded_key) == () assert agent_365_authorization_servers(server, None) == () assert agent_365_scopes_supported(_mcp_server(scopes=None), None) == () finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( litellm.callbacks, guardrail, require_self=False ) + + @pytest.mark.parametrize( + "metadata", + [{"opted_out_global_guardrails": ["agent-365-guard"]}, {"disable_global_guardrails": True}], + ids=["opted-out", "globals-disabled"], + ) + def test_key_opted_out_of_the_default_on_guardrail_is_not_challenged(self, registered_guardrail, metadata): + server: Final = _mcp_server(scopes=[GATEWAY_SCOPE]) + opted_out: Final = UserAPIKeyAuth(api_key="sk-out", user_id="u-3", metadata=metadata) + team_opted_out: Final = UserAPIKeyAuth(api_key="sk-team", user_id="u-4", team_metadata=metadata) + + assert agent_365_authorization_servers(server, opted_out) == () + assert agent_365_authorization_servers(server, team_opted_out) == () + assert agent_365_authorization_servers(server, UserAPIKeyAuth(api_key="sk-in", user_id="u-5")) == ( + ENTRA_ISSUER, + )