From f028f32c9e599e2a2faf809c341c639d3aaa47ec Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 26 Jun 2026 19:53:27 -0700 Subject: [PATCH] fix(mcp): keep user token in authorization_code tools preview After to_server_spec maps oauth2 onto the v2 resolver, the interactive tools preview for an unsaved authorization_code server read the per-user token store, found nothing, and fail-closed with a 401, so the create/test tab could no longer list tools The preview now routes the just-authorized token (forwarded in oauth2_headers) through mcp_auth_header, so _create_mcp_client takes the per-request-override v1 path and uses it directly, matching v1's preview. Gated to the v2-mapped oauth2 case; M2M, delegate/passthrough, and token-exchange keep their existing preview path Adds tests: interactive oauth routes the forwarded token to mcp_auth_header, M2M and token-exchange do not --- .../mcp_server/rest_endpoints.py | 26 +++ .../mcp_server/test_rest_endpoints.py | 162 +++++++++++++++++- 2 files changed, 179 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7b4f1e13a52..e204ec7bd35 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -1039,6 +1039,32 @@ if MCP_AVAILABLE: static_headers=request.static_headers, ) + # Interactive authorization_code tools preview: the UI forwards the just-authorized token + # in oauth2_headers before any per-user credential is persisted. Route it through + # mcp_auth_header so _create_mcp_client's per-request-override deferral takes the v1 path + # (use the forwarded token directly - no resolver, no fail-closed challenge), exactly as + # v1's preview did. Gated to the v2-mapped oauth2 case (to_server_spec non-None == + # authorization_code): M2M (client_credentials) and delegate/passthrough return None from + # to_server_spec, and token-exchange is a different auth_type, so all three keep their + # existing v1 preview behavior untouched. + if ( + mcp_auth_header is None + and server_model.auth_type == MCPAuth.oauth2 + and oauth2_headers + ): + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 + to_server_spec, + ) + + forwarded_authorization = oauth2_headers.get("Authorization") + if forwarded_authorization and to_server_spec(server_model) is not None: + # The oauth2 client re-adds the "Bearer " scheme, so pass the bare token. + mcp_auth_header = ( + forwarded_authorization[7:] + if forwarded_authorization[:7].lower() == "bearer " + else forwarded_authorization + ) + client = await global_mcp_server_manager._create_mcp_client( server=server_model, mcp_auth_header=mcp_auth_header, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 27f9a311250..e92a54e313f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -270,6 +270,154 @@ class TestExecuteWithMcpClient: or "Authorization" not in captured["extra_headers"] ) + @pytest.mark.asyncio + async def test_interactive_oauth_routes_forwarded_token_to_mcp_auth_header( + self, monkeypatch + ): + """Interactive authorization_code preview (oauth2, no client credentials): the forwarded + just-authorized token must be routed to mcp_auth_header so _create_mcp_client takes the v1 + per-request-override path and uses it directly, instead of the v2 resolver fail-closing on + the not-yet-persisted credential. The "Bearer " scheme is stripped since the oauth2 client + re-adds it.""" + captured: dict = {} + + def fake_build_stdio_env(server, raw_headers): + return None + + async def fake_create_client(*args, **kwargs): + captured["mcp_auth_header"] = kwargs.get("mcp_auth_header") + return object() + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_build_stdio_env", + fake_build_stdio_env, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_create_mcp_client", + fake_create_client, + raising=False, + ) + + async def ok_operation(client): + return {"status": "ok"} + + payload = NewMCPServerRequest( + server_name="linear", + url="https://mcp.linear.app/mcp", + auth_type=MCPAuth.oauth2, + authorization_url="https://mcp.linear.app/authorize", + ) + + result = await rest_endpoints._execute_with_mcp_client( + payload, + ok_operation, + oauth2_headers={"Authorization": "Bearer forwarded-user-token"}, + ) + + assert result["status"] == "ok" + assert captured["mcp_auth_header"] == "forwarded-user-token" + + @pytest.mark.asyncio + async def test_m2m_does_not_route_forwarded_token_to_mcp_auth_header( + self, monkeypatch + ): + """M2M (client_credentials) must NOT route the forwarded header to mcp_auth_header - it stays + on the auto-fetch path. to_server_spec returns None for M2M, so the interactive shortcut is + skipped and mcp_auth_header remains None (the incoming header is dropped, as before).""" + captured: dict = {} + + def fake_build_stdio_env(server, raw_headers): + return None + + async def fake_create_client(*args, **kwargs): + captured["mcp_auth_header"] = kwargs.get("mcp_auth_header") + return object() + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_build_stdio_env", + fake_build_stdio_env, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_create_mcp_client", + fake_create_client, + raising=False, + ) + + async def ok_operation(client): + return {"status": "ok"} + + payload = NewMCPServerRequest( + server_name="m2m-server", + url="https://example.com", + auth_type=MCPAuth.oauth2, + token_url="https://auth.example.com/token", + credentials={"client_id": "my-id", "client_secret": "my-secret"}, + ) + + result = await rest_endpoints._execute_with_mcp_client( + payload, + ok_operation, + oauth2_headers={"Authorization": "Bearer sk-litellm-api-key"}, + ) + + assert result["status"] == "ok" + assert captured["mcp_auth_header"] is None + + @pytest.mark.asyncio + async def test_token_exchange_does_not_route_forwarded_token_to_mcp_auth_header( + self, monkeypatch + ): + """OBO / token-exchange (auth_type oauth2_token_exchange, not oauth2) must NOT route the + forwarded header to mcp_auth_header - resolve_mcp_auth performs the exchange on the v1 path. + The guard's auth_type == oauth2 check excludes it, so mcp_auth_header stays None (setting it + would short-circuit the exchange, since resolve_mcp_auth returns mcp_auth_header first).""" + captured: dict = {} + + def fake_build_stdio_env(server, raw_headers): + return None + + async def fake_create_client(*args, **kwargs): + captured["mcp_auth_header"] = kwargs.get("mcp_auth_header") + return object() + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_build_stdio_env", + fake_build_stdio_env, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_create_mcp_client", + fake_create_client, + raising=False, + ) + + async def ok_operation(client): + return {"status": "ok"} + + payload = NewMCPServerRequest( + server_name="obo-server", + url="https://example.com", + auth_type=MCPAuth.oauth2_token_exchange, + token_url="https://auth.example.com/token", + ) + + result = await rest_endpoints._execute_with_mcp_client( + payload, + ok_operation, + oauth2_headers={"Authorization": "Bearer subject-jwt"}, + ) + + assert result["status"] == "ok" + assert captured["mcp_auth_header"] is None + @pytest.mark.asyncio async def test_catches_exception_group(self, monkeypatch): """MCP SDK's anyio TaskGroup raises BaseExceptionGroup which does not @@ -1660,9 +1808,9 @@ class TestPreviewOpenAPITools: names = [t["name"] for t in result["tools"]] anthropic_re = re.compile(r"^[a-zA-Z0-9_-]{1,128}$") for name in names: - assert anthropic_re.match( - name - ), f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$" + assert anthropic_re.match(name), ( + f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$" + ) assert "actions_download-job-logs-for-workflow-run" in names assert "pulls_list-files" in names @@ -1719,9 +1867,7 @@ class TestPreviewOpenAPITools: registered_summary_to_name: dict = {} - def fake_create_tool_function( - path, method, operation, base_url - ): # noqa: ANN001 + def fake_create_tool_function(path, method, operation, base_url): # noqa: ANN001 def _f(): return None @@ -1734,9 +1880,7 @@ class TestPreviewOpenAPITools: ) class _StubRegistry: - def register_tool( - self, name, description, input_schema, handler - ): # noqa: ANN001 + def register_tool(self, name, description, input_schema, handler): # noqa: ANN001 registered_summary_to_name[description] = name monkeypatch.setattr(