diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index f18acb7d88d..9529ede4b97 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -76,6 +76,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.guardrails.guardrail_hooks.agent_365.agent_365 import agent_365_authorization_servers from litellm.types.mcp import MCPAuth, MCPCredentials from litellm.types.mcp_server.mcp_server_manager import MCPServer, MCPTokenEndpointAuthMethod @@ -2400,6 +2401,14 @@ async def _build_oauth_protected_resource_response( if obo_response is not None: return obo_response + agent_365_issuers: Final = agent_365_authorization_servers(mcp_server, None) if mcp_server else () + if mcp_server is not None and agent_365_issuers: + return { + "authorization_servers": list(agent_365_issuers), + "resource": resource_url, + "scopes_supported": list(mcp_server.scopes or ()), + } + if explicitly_named and mcp_server is not None and mcp_server.advertises_gateway_authorization_server: return { "authorization_servers": [f"{request_base_url}/mcp"], diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d259fa32cd9..24527a380d4 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -79,6 +79,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils +from litellm.proxy.guardrails.guardrail_hooks.agent_365.agent_365 import agent_365_authorization_servers from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, get_chain_id_from_headers, @@ -4101,8 +4102,16 @@ if MCP_AVAILABLE: # (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata # so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM # 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. - if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers: + # 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. + if ( + server + and not oauth2_headers + and ( + server.auth_type == MCPAuth.oauth2_token_exchange + or agent_365_authorization_servers(server, user_api_key_auth) + ) + ): from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph raise_token_exchange_challenge, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index a7a65bfe466..ab84777dbeb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7090,6 +7090,77 @@ async def test_build_oauth_protected_resource_response_obo_end_to_end(): global_mcp_server_manager.registry.clear() +@pytest.fixture +def agent_365_guardrail(): + 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=True, + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + yield guardrail + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + +async def _agent_365_gated_prm(scopes): + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _build_oauth_protected_resource_response, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry["tools"] = MCPServer( + server_id="tools", + name="tools", + server_name="tools", + alias="tools", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + scopes=scopes, + ) + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + try: + return await _build_oauth_protected_resource_response( + request=mock_request, mcp_server_name="tools", use_standard_pattern=True + ) + finally: + global_mcp_server_manager.registry.clear() + + +@pytest.mark.asyncio +async def test_agent_365_gated_server_prm_names_the_entra_tenant(agent_365_guardrail): + response = await _agent_365_gated_prm(scopes=["api://gateway-app/access_as_user"]) + assert response == { + "authorization_servers": ["https://login.microsoftonline.com/tenant-abc/v2.0"], + "resource": "https://litellm.example.com/mcp/tools", + "scopes_supported": ["api://gateway-app/access_as_user"], + } + + +@pytest.mark.asyncio +async def test_agent_365_prm_falls_back_to_gateway_without_scopes(agent_365_guardrail): + response = await _agent_365_gated_prm(scopes=None) + assert response["authorization_servers"] == ["https://litellm.example.com/mcp"] + + def _token_request(headers): """A real Starlette request with case-insensitive headers (matches production).""" from starlette.requests import Request 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 12805e355f9..69e932007aa 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 @@ -8739,6 +8739,90 @@ class TestSingleServerPreflightReachesIdJag: preflight.assert_not_awaited() +class TestAgent365ChallengeAtConnect: + """A missing Entra bearer on an Agent 365 gated server is challenged at connect (RFC 9728), where the + WWW-Authenticate header survives, instead of only inside the tools/call JSON-RPC error.""" + + GATEWAY_SCOPE = "api://gateway-app/access_as_user" + + def _server(self, scopes: list[str] | None) -> MCPServer: + return MCPServer( + server_id="id-tools", + name="tools", + alias="tools", + server_name="tools", + url="https://tools.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + scopes=scopes, + mcp_info={"server_name": "tools"}, + ) + + @pytest.fixture + def agent_365_guardrail(self): + 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=True, + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + yield guardrail + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + async def _connect(self, server: MCPServer, oauth2_headers: dict[str, str] | None) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + with ( + patch.object( # test-quality-ok: route wiring must use the manager's configured server + server_module.global_mcp_server_manager, "get_mcp_server_by_name", return_value=server + ), + patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer + server_module, "_get_allowed_mcp_servers", AsyncMock(return_value=[]) + ), + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/tools", "headers": []}, + 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"), + client_ip=None, + ) + + @pytest.mark.asyncio + async def test_no_bearer_gets_the_discovery_challenge(self, agent_365_guardrail): + with pytest.raises(HTTPException) as exc: + await self._connect(self._server([self.GATEWAY_SCOPE]), None) + + assert exc.value.status_code == 401 + www_authenticate = (exc.value.headers or {}).get("WWW-Authenticate", "") + assert 'error="invalid_token"' in www_authenticate + assert 'resource_metadata="/.well-known/oauth-protected-resource/mcp/tools"' in www_authenticate + + @pytest.mark.asyncio + async def test_bearer_present_connects(self, agent_365_guardrail): + await self._connect(self._server([self.GATEWAY_SCOPE]), {"Authorization": "Bearer entra-user-token"}) + + @pytest.mark.asyncio + async def test_server_without_advertised_scopes_is_not_challenged(self, agent_365_guardrail): + await self._connect(self._server(None), None) + + @pytest.mark.asyncio + async def test_no_registered_guardrail_means_no_challenge(self): + await self._connect(self._server([self.GATEWAY_SCOPE]), None) + + def _make_obo_server(alias: str) -> MCPServer: return MCPServer( server_id=f"id-{alias}",