diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 287913d8eb7..bdd3c983717 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4211,6 +4211,7 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, + connected_as: str | None = None, ) -> None: """Mint an exchange-backed server's upstream credential at the transport edge. @@ -4247,14 +4248,16 @@ class MCPServerManager: ) if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None: - raise_token_exchange_challenge(server, root_path=get_request_root_path()) + raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=connected_as) return resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) spec: Final = _to_server_spec_fail_closed(resolved_server) if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)): return if subject_token is None and isinstance(spec.config, TokenExchangeConfig): - raise_token_exchange_challenge(resolved_server, root_path=get_request_root_path()) + raise_token_exchange_challenge( + resolved_server, root_path=get_request_root_path(), connected_as=connected_as + ) match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec): case Ok(_): return @@ -4264,6 +4267,7 @@ class MCPServerManager: resolved_server, root_path=get_request_root_path(), claims=err.unauthorized.claims, + connected_as=connected_as, ) raise_public(err) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 470fc6c2308..988f77ea98b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1763,6 +1763,7 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, + connected_as=server_name, ) # Pass-through OAuth: when the admin has opted a server into diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py index 2a8c6479ae6..dc0257f2017 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py @@ -7,6 +7,7 @@ from .agent_365 import Agent365Guardrail if TYPE_CHECKING: from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger from litellm.types.guardrails import Guardrail, LitellmParams @@ -15,6 +16,7 @@ def initialize_guardrail( guardrail: "Guardrail", *, async_handler: "AsyncHTTPHandler | None" = None, + token_exchanger: "TokenExchanger | None" = None, ) -> Agent365Guardrail: import litellm from litellm.secret_managers.main import get_secret_str @@ -64,6 +66,7 @@ def initialize_guardrail( request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0, unreachable_fallback=litellm_params.unreachable_fallback, async_handler=async_handler, + token_exchanger=token_exchanger, event_hook=litellm_params.mode, default_on=litellm_params.default_on, ) 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 8194d7ab5c5..14d71f98900 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 @@ -2930,6 +2930,31 @@ class TestMCPServerManager: www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "" assert "resource_metadata" in www_authenticate + @pytest.mark.asyncio + async def test_preflight_rejected_subject_challenge_names_the_connected_segment(self): + """A subject rejected on ``/mcp/`` must point resource_metadata at that same + segment, the way the sign-in preflight does, so the client's discovery fetch resolves.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + return Error(CredError.of_unauthorized("subject token rejected by the IdP")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + server = self._token_exchange_server("te-preflight-segment") + + with pytest.raises(HTTPException) as exc_info: + await manager.preflight_token_exchange( + server=server, + oauth2_headers={"Authorization": "Bearer rejected-subject"}, + user_api_key_auth=None, + connected_as=server.server_id, + ) + headers = exc_info.value.headers or {} + www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "" + assert f"/.well-known/oauth-protected-resource/mcp/{server.server_id}" in www_authenticate, www_authenticate + @pytest.mark.asyncio async def test_preflight_token_exchange_maps_gateway_fault_to_public_status(self): """A gateway-fault CredError (e.g. invalid_client) must surface its public status (500) 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 ba90460d573..c57a3b42975 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 @@ -194,7 +194,16 @@ def _make_guardrail( def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchanger | None = None) -> Agent365Guardrail: - return _make_guardrail(handler, exchanger=exchanger) + return Agent365Guardrail( + guardrail_name="agent-365-guard", + tenant_id="tenant-abc", + client_id="client-xyz", + client_secret="secret-123", + async_handler=handler, + token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()), + event_hook="pre_mcp_call", + default_on=True, + ) def _server(**overrides: Any) -> MCPServer: @@ -320,11 +329,16 @@ class TestInitializeGuardrail: agent_id="yaml-agent", ) handler: Final = FakeHandler([_allow_response()]) + exchanger: Final = StubTokenExchanger(_obo_ok()) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler) + guardrail: Final = initialize_guardrail( + params, {"guardrail_name": "a365-stale"}, async_handler=handler, token_exchanger=exchanger + ) assert "ignoring api_base, resource_app_id, agent_id" in caplog.text - guardrail: Final = _make_guardrail(handler) await _run(guardrail, _mcp_data()) + _, server, config = exchanger.calls[0] + assert server.resource == AGENT_365_PROD_API_BASE + assert config.scopes == (f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}",) evaluate_call: Final = handler.calls[0] assert evaluate_call.url == EVALUATE_URL assert evaluate_call.json["agentId"] == "my-agent-key"