diff --git a/apps/desktop/src/app/skills/mcp-tab.tsx b/apps/desktop/src/app/skills/mcp-tab.tsx index 3532405becfd..522261b0e104 100644 --- a/apps/desktop/src/app/skills/mcp-tab.tsx +++ b/apps/desktop/src/app/skills/mcp-tab.tsx @@ -30,6 +30,7 @@ import { getActionStatus, getLogs, getMcpCatalog, + getMcpOAuthFlow, type HermesGateway, installMcpCatalogEntry, type McpCatalogEntry, @@ -38,6 +39,7 @@ import { testMcpServer } from '@/hermes' import { type Translations, useI18n } from '@/i18n' +import { completeMcpDesktopOAuth } from '@/lib/mcp-dashboard-oauth' import { countEnabledTools, isToolEnabled, toggleToolInServer } from '@/lib/mcp-tool-filter' import { cn } from '@/lib/utils' import { notify, notifyError } from '@/store/notifications' @@ -578,7 +580,14 @@ export function McpTab({ gateway }: { gateway: HermesGateway | null }) { setProbes(current => ({ ...current, [serverName]: 'probing' })) try { - const result = await authMcpServer(serverName) + const flow = await completeMcpDesktopOAuth({ + serverName, + start: authMcpServer, + status: getMcpOAuthFlow, + openExternal: url => window.hermesDesktop.openExternal(url) + }) + + const result: McpTestResult = { ok: true, tools: flow.tools ?? [] } // Bail if the user switched profiles mid-flow — this result is profile A's. if (profileEpoch.current !== epoch) { diff --git a/apps/desktop/src/hermes.ts b/apps/desktop/src/hermes.ts index 01c9eb6ec857..e637e0763137 100644 --- a/apps/desktop/src/hermes.ts +++ b/apps/desktop/src/hermes.ts @@ -709,6 +709,15 @@ export interface McpTestResult { resources?: number } +export interface McpOAuthFlow { + flow_id: string + server_name: string + status: 'starting' | 'authorization_required' | 'approved' | 'error' + authorization_url: string | null + error: string | null + tools?: { name: string; description: string }[] +} + /** Connect to the server, list its tools, disconnect. Slow (spawns/handshakes * for real) — well past the 15s default fetch timeout. */ export function testMcpServer(name: string): Promise { @@ -732,14 +741,20 @@ export function saveMcpServers(servers: Record>) }) } -/** Run the OAuth flow for an HTTP server — opens the system browser and blocks - * until the user finishes (or gives up), hence the very generous timeout. */ -export function authMcpServer(name: string): Promise { - return window.hermesDesktop.api({ +/** Start an MCP OAuth flow and return the authorization URL. */ +export function authMcpServer(name: string): Promise { + return window.hermesDesktop.api({ ...profileScoped(), path: `/api/mcp/servers/${encodeURIComponent(name)}/auth`, method: 'POST', - timeoutMs: 300_000 + timeoutMs: 60_000 + }) +} + +export function getMcpOAuthFlow(flowId: string): Promise { + return window.hermesDesktop.api({ + ...profileScoped(), + path: `/api/mcp/oauth/flows/${encodeURIComponent(flowId)}` }) } diff --git a/apps/desktop/src/lib/mcp-dashboard-oauth.test.ts b/apps/desktop/src/lib/mcp-dashboard-oauth.test.ts new file mode 100644 index 000000000000..4473cd9e1b35 --- /dev/null +++ b/apps/desktop/src/lib/mcp-dashboard-oauth.test.ts @@ -0,0 +1,35 @@ +import { describe, expect, it, vi } from 'vitest' + +import { completeMcpDesktopOAuth } from './mcp-dashboard-oauth' + +describe('completeMcpDesktopOAuth', () => { + it('opens the returned authorization URL and polls through approval', async () => { + const openExternal = vi.fn().mockResolvedValue(undefined) + + const status = vi + .fn() + .mockResolvedValueOnce({ + flow_id: 'flow-1', server_name: 'reports', status: 'authorization_required', + authorization_url: 'https://idp.example/authorize', error: null + }) + .mockResolvedValueOnce({ + flow_id: 'flow-1', server_name: 'reports', status: 'approved', + authorization_url: 'https://idp.example/authorize', error: null, + tools: [{ name: 'list_reports', description: 'List reports' }] + }) + + const result = await completeMcpDesktopOAuth({ + serverName: 'reports', + start: vi.fn().mockResolvedValue({ + flow_id: 'flow-1', server_name: 'reports', status: 'authorization_required', + authorization_url: 'https://idp.example/authorize', error: null + }), + status, + openExternal, + sleep: async () => {} + }) + + expect(openExternal).toHaveBeenCalledWith('https://idp.example/authorize') + expect(result.status).toBe('approved') + }) +}) \ No newline at end of file diff --git a/apps/desktop/src/lib/mcp-dashboard-oauth.ts b/apps/desktop/src/lib/mcp-dashboard-oauth.ts new file mode 100644 index 000000000000..5cd9570099df --- /dev/null +++ b/apps/desktop/src/lib/mcp-dashboard-oauth.ts @@ -0,0 +1,53 @@ +export interface McpOAuthFlow { + flow_id: string + server_name: string + status: 'starting' | 'authorization_required' | 'approved' | 'error' + authorization_url: string | null + error: string | null + tools?: Array<{ name: string; description: string }> +} + +interface CompleteOptions { + serverName: string + start: (name: string) => Promise + status: (flowId: string) => Promise + openExternal: (url: string) => Promise + sleep?: (milliseconds: number) => Promise +} + +const defaultSleep = (milliseconds: number) => + new Promise(resolve => window.setTimeout(resolve, milliseconds)) + +export async function completeMcpDesktopOAuth({ + serverName, + start, + status, + openExternal, + sleep = defaultSleep +}: CompleteOptions): Promise { + const started = await start(serverName) + + if (started.status === 'error') { + throw new Error(started.error || 'OAuth failed to start') + } + + if (!started.authorization_url) { + throw new Error('OAuth server did not provide an authorization URL') + } + + await openExternal(started.authorization_url) + + for (;;) { + const current = await status(started.flow_id) + + if (current.status === 'approved') { + return current + } + + if (current.status === 'error') { + throw new Error(current.error || 'OAuth authorization failed') + } + + await sleep(1000) + } +} \ No newline at end of file diff --git a/hermes_cli/web_server.py b/hermes_cli/web_server.py index e6acc5542833..ddb7e79b2eee 100644 --- a/hermes_cli/web_server.py +++ b/hermes_cli/web_server.py @@ -11329,26 +11329,38 @@ async def test_mcp_server(name: str, profile: Optional[str] = None): _MCP_DASHBOARD_OAUTH_TTL = 15 * 60 _MAX_PENDING_MCP_OAUTH_FLOWS = 8 _mcp_oauth_flows: dict[str, "DashboardOAuthFlow"] = {} +_mcp_oauth_flows_lock = threading.Lock() +_mcp_oauth_transactions: dict[tuple[str, str], threading.Lock] = {} +_mcp_oauth_transactions_lock = threading.Lock() def _gc_mcp_oauth_flows() -> None: cutoff = time.time() - _MCP_DASHBOARD_OAUTH_TTL - stale = [ - flow_id - for flow_id, flow in _mcp_oauth_flows.items() - if getattr(flow, "created_at", 0) < cutoff - ] - for flow_id in stale: - _mcp_oauth_flows.pop(flow_id, None) + with _mcp_oauth_flows_lock: + stale = [ + flow_id + for flow_id, flow in _mcp_oauth_flows.items() + if getattr(flow, "created_at", 0) < cutoff + ] + for flow_id in stale: + _mcp_oauth_flows.pop(flow_id, None) -def _mcp_oauth_callback_url(request: Request, flow_id: str) -> str: +def _mcp_oauth_callback_url_from_base(base_url: str, server_name: str) -> str: + from urllib.parse import quote + + return f"{base_url.rstrip('/')}/api/mcp/oauth/callback/{quote(server_name, safe='')}" + + +def _mcp_oauth_callback_url(request: Request, server_name: str) -> str: """Build the externally reachable callback URL for a dashboard flow.""" from urllib.parse import urlparse, urlunparse from hermes_cli.dashboard_auth.prefix import prefix_from_request, resolve_public_url - suffix = f"/api/mcp/oauth/callback/{flow_id}" + from urllib.parse import quote + + suffix = f"/api/mcp/oauth/callback/{quote(server_name, safe='')}" public_url = resolve_public_url() if public_url: return f"{public_url}{suffix}" @@ -11357,6 +11369,12 @@ def _mcp_oauth_callback_url(request: Request, flow_id: str) -> str: return urlunparse(base._replace(path=f"{prefix}{suffix}", params="", query="", fragment="")) +def _mcp_oauth_transaction(flow) -> threading.Lock: + key = (flow.hermes_home, flow.server_name) + with _mcp_oauth_transactions_lock: + return _mcp_oauth_transactions.setdefault(key, threading.Lock()) + + def _run_dashboard_mcp_oauth(flow, cfg: dict) -> None: """Run the normal MCP probe with dashboard redirect/callback handlers.""" from hermes_cli.mcp_config import ( @@ -11365,35 +11383,47 @@ def _run_dashboard_mcp_oauth(flow, cfg: dict) -> None: _save_mcp_server, ) try: + from agent.secret_scope import ( + build_profile_secret_scope, + reset_secret_scope, + set_secret_scope, + ) + from hermes_constants import reset_hermes_home_override, set_hermes_home_override from tools.mcp_dashboard_oauth import dashboard_oauth_flow from tools.mcp_oauth import HermesTokenStorage, force_interactive_oauth from tools.mcp_oauth_manager import get_manager - with ( - _config_profile_scope(flow.profile), - force_interactive_oauth(), - dashboard_oauth_flow(flow), - ): - storage = HermesTokenStorage(flow.server_name) - backup = storage.snapshot() - try: - get_manager().remove(flow.server_name) - tools = _probe_single_server( - flow.server_name, - cfg, - connect_timeout=max(float(cfg.get("connect_timeout", 0) or 0), 315), - ) - if not _oauth_tokens_present(flow.server_name): - raise RuntimeError( - "The server responded, but no OAuth token was obtained — " - "this provider may require a manually-registered OAuth client." + home_token = set_hermes_home_override(flow.hermes_home) + secret_token = set_secret_scope(build_profile_secret_scope(Path(flow.hermes_home))) + try: + transaction = _mcp_oauth_transaction(flow) + with transaction, force_interactive_oauth(), dashboard_oauth_flow(flow): + storage = HermesTokenStorage(flow.server_name) + backup = storage.snapshot() + try: + get_manager().remove(flow.server_name, hermes_home=flow.hermes_home) + tools = _probe_single_server( + flow.server_name, + cfg, + connect_timeout=max(float(cfg.get("connect_timeout", 0) or 0), 315), ) - _save_mcp_server(flow.server_name, cfg) - flow.tools = [{"name": t, "description": d} for t, d in tools] - flow.mark_approved() - except Exception: - storage.restore(backup) - raise + if not _oauth_tokens_present(flow.server_name): + raise RuntimeError( + "The server responded, but no OAuth token was obtained — " + "this provider may require a manually-registered OAuth client." + ) + _save_mcp_server(flow.server_name, cfg) + flow.tools = [{"name": t, "description": d} for t, d in tools] + flow.mark_approved() + from tools.mcp_tool import reconnect_mcp_server + + reconnect_mcp_server(flow.server_name) + except Exception: + storage.restore(backup) + raise + finally: + reset_secret_scope(secret_token) + reset_hermes_home_override(home_token) except Exception as exc: msg = str(exc) # Providers that gate RFC 7591 registration to pre-approved clients @@ -11417,7 +11447,7 @@ def _run_dashboard_mcp_oauth(flow, cfg: dict) -> None: try: from tools.mcp_oauth_manager import get_manager - get_manager().evict(flow.server_name) + get_manager().evict(flow.server_name, hermes_home=flow.hermes_home) except Exception: pass @@ -11430,26 +11460,11 @@ async def auth_mcp_server(name: str, request: Request, profile: Optional[str] = _require_token(request) _gc_mcp_oauth_flows() - pending = sum( - flow.status in {"starting", "authorization_required"} - for flow in _mcp_oauth_flows.values() - ) - if pending >= _MAX_PENDING_MCP_OAUTH_FLOWS: - raise HTTPException( - status_code=429, - detail="Too many MCP OAuth flows are already in progress", - ) - if any( - flow.server_name == name - and flow.status in {"starting", "authorization_required"} - for flow in _mcp_oauth_flows.values() - ): - raise HTTPException( - status_code=409, - detail=f"MCP OAuth for '{name}' is already in progress", - ) with _profile_scope(profile): servers = _get_mcp_servers() + from hermes_constants import get_hermes_home + + flow_home = str(get_hermes_home().expanduser().resolve(strict=False)) if name not in servers: raise HTTPException(status_code=404, detail=f"Server '{name}' not found") cfg = dict(servers[name]) @@ -11459,14 +11474,38 @@ async def auth_mcp_server(name: str, request: Request, profile: Optional[str] = raise HTTPException(status_code=400, detail="This server uses header/API-key auth, not OAuth") cfg["auth"] = "oauth" + with _mcp_oauth_flows_lock: + pending = sum( + flow.status in {"starting", "authorization_required"} + for flow in _mcp_oauth_flows.values() + ) + if pending >= _MAX_PENDING_MCP_OAUTH_FLOWS: + raise HTTPException( + status_code=429, + detail="Too many MCP OAuth flows are already in progress", + ) + if any( + flow.server_name == name + and flow.hermes_home == flow_home + and flow.status in {"starting", "authorization_required"} + for flow in _mcp_oauth_flows.values() + ): + raise HTTPException( + status_code=409, + detail=f"MCP OAuth for '{name}' is already in progress", + ) + flow_id = secrets.token_urlsafe(24) flow = DashboardOAuthFlow( flow_id=flow_id, server_name=name, profile=profile, - redirect_uri=_mcp_oauth_callback_url(request, flow_id), + hermes_home=flow_home, + redirect_uri=(cfg.get("oauth") or {}).get("redirect_uri") + or _mcp_oauth_callback_url(request, name), ) - _mcp_oauth_flows[flow_id] = flow + with _mcp_oauth_flows_lock: + _mcp_oauth_flows[flow_id] = flow threading.Thread( target=_run_dashboard_mcp_oauth, args=(flow, cfg), @@ -11492,15 +11531,31 @@ async def mcp_oauth_flow_status(flow_id: str, request: Request): return snapshot -@app.get("/api/mcp/oauth/callback/{flow_id}") +@app.get("/api/mcp/oauth/callback/{server_name}") async def mcp_oauth_callback( - flow_id: str, + server_name: str, code: Optional[str] = None, state: Optional[str] = None, error: Optional[str] = None, ): _gc_mcp_oauth_flows() - flow = _mcp_oauth_flows.get(flow_id) + with _mcp_oauth_flows_lock: + candidates = [ + flow + for flow in _mcp_oauth_flows.values() + if flow.server_name == server_name + and flow.status == "authorization_required" + ] + flow = next( + ( + candidate + for candidate in candidates + if candidate.expected_state is not None + and state is not None + and secrets.compare_digest(candidate.expected_state, state) + ), + None, + ) if flow is None: return HTMLResponse("

OAuth flow expired

Return to Hermes and try again.

", status_code=404) try: diff --git a/tests/hermes_cli/test_mcp_config.py b/tests/hermes_cli/test_mcp_config.py index 64e24c42470f..f79e29ff7f03 100644 --- a/tests/hermes_cli/test_mcp_config.py +++ b/tests/hermes_cli/test_mcp_config.py @@ -839,12 +839,12 @@ class TestMcpRemoveEvictsManager: mgr.get_or_build_provider( "oauth-srv", "https://example.com/mcp", None, ) - assert "oauth-srv" in mgr._entries + assert mgr._key("oauth-srv") in mgr._entries from hermes_cli.mcp_config import cmd_mcp_remove cmd_mcp_remove(_make_args(name="oauth-srv")) - assert "oauth-srv" not in mgr._entries + assert mgr._key("oauth-srv") not in mgr._entries class TestMcpLogin: diff --git a/tests/hermes_cli/test_mcp_dashboard_oauth.py b/tests/hermes_cli/test_mcp_dashboard_oauth.py index fdf4be82f15f..1862e4373d6b 100644 --- a/tests/hermes_cli/test_mcp_dashboard_oauth.py +++ b/tests/hermes_cli/test_mcp_dashboard_oauth.py @@ -20,8 +20,10 @@ def _clear_flows(): from hermes_cli import web_server web_server._mcp_oauth_flows.clear() + web_server.app.state.auth_required = False yield web_server._mcp_oauth_flows.clear() + web_server.app.state.auth_required = False def test_hosted_auth_start_returns_public_authorization_url(monkeypatch): @@ -50,7 +52,7 @@ def test_hosted_auth_start_returns_public_authorization_url(monkeypatch): assert body["status"] == "authorization_required" assert body["authorization_url"] == "https://idp.example/authorize?state=s1" flow = web_server._mcp_oauth_flows[body["flow_id"]] - assert flow.redirect_uri == f"https://agent.example/api/mcp/oauth/callback/{body['flow_id']}" + assert flow.redirect_uri == "https://agent.example/api/mcp/oauth/callback/reports" def test_hosted_callback_is_public_and_delivers_code(): @@ -64,7 +66,8 @@ def test_hosted_callback_is_public_and_delivers_code(): flow_id="flow-public", server_name="reports", profile=None, - redirect_uri="https://agent.example/api/mcp/oauth/callback/flow-public", + hermes_home="/tmp/hermes-test", + redirect_uri="https://agent.example/api/mcp/oauth/callback/reports", ) asyncio.run( flow.publish_authorization_url( @@ -75,7 +78,7 @@ def test_hosted_callback_is_public_and_delivers_code(): assert "/api/mcp/oauth/callback" not in PUBLIC_API_PATHS response = _client().get( - "/api/mcp/oauth/callback/flow-public?code=abc&state=expected" + "/api/mcp/oauth/callback/reports?code=abc&state=expected" ) assert response.status_code == 200 assert flow._callback == ("abc", "expected") @@ -93,7 +96,8 @@ def test_hosted_callback_bypasses_gated_cookie_auth(monkeypatch): flow_id="flow-gated", server_name="reports", profile=None, - redirect_uri="https://agent.example/api/mcp/oauth/callback/flow-gated", + hermes_home="/tmp/hermes-test", + redirect_uri="https://agent.example/api/mcp/oauth/callback/reports", ) asyncio.run( flow.publish_authorization_url( @@ -104,7 +108,7 @@ def test_hosted_callback_bypasses_gated_cookie_auth(monkeypatch): monkeypatch.setattr(web_server.app.state, "auth_required", True, raising=False) response = TestClient(web_server.app).get( - "/api/mcp/oauth/callback/flow-gated?code=abc&state=expected" + "/api/mcp/oauth/callback/reports?code=abc&state=expected" ) assert response.status_code == 200 @@ -121,7 +125,8 @@ def test_hosted_callback_rejects_wrong_state_before_waking_sdk(): flow_id="flow-state-route", server_name="reports", profile=None, - redirect_uri="https://agent.example/api/mcp/oauth/callback/flow-state-route", + hermes_home="/tmp/hermes-test", + redirect_uri="https://agent.example/api/mcp/oauth/callback/reports", ) asyncio.run( flow.publish_authorization_url( @@ -131,9 +136,9 @@ def test_hosted_callback_rejects_wrong_state_before_waking_sdk(): web_server._mcp_oauth_flows[flow.flow_id] = flow response = _client().get( - "/api/mcp/oauth/callback/flow-state-route?code=attacker&state=wrong" + "/api/mcp/oauth/callback/reports?code=attacker&state=wrong" ) - assert response.status_code == 400 + assert response.status_code == 404 assert flow._callback is None @@ -151,6 +156,7 @@ def test_hosted_auth_start_bounds_pending_flow_registry(): flow_id=f"existing-{index}", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri=f"https://agent.example/callback/{index}", ) web_server._mcp_oauth_flows[flow.flow_id] = flow @@ -168,10 +174,13 @@ def test_hosted_auth_rejects_overlapping_flow_for_same_server(): "/api/mcp/servers", json={"name": "reports", "url": "https://mcp.example/mcp", "auth": "oauth"}, ) + from hermes_constants import get_hermes_home + existing = DashboardOAuthFlow( flow_id="existing-reports", server_name="reports", profile="other-profile", + hermes_home=str(get_hermes_home().expanduser().resolve(strict=False)), redirect_uri="https://agent.example/callback/existing", ) web_server._mcp_oauth_flows[existing.flow_id] = existing @@ -182,6 +191,44 @@ def test_hosted_auth_rejects_overlapping_flow_for_same_server(): assert "already in progress" in response.text +def test_hosted_auth_allows_same_server_name_in_different_profiles(tmp_path, monkeypatch): + from hermes_cli import web_server + from tools.mcp_dashboard_oauth import DashboardOAuthFlow + + profile_home = tmp_path / "profiles" / "work" + profile_home.mkdir(parents=True) + monkeypatch.setattr(web_server, "_resolve_profile_dir", lambda _name: profile_home) + + existing = DashboardOAuthFlow( + flow_id="existing-default", + server_name="reports", + profile=None, + hermes_home=str(tmp_path / "default"), + redirect_uri="https://agent.example/callback/existing", + ) + web_server._mcp_oauth_flows[existing.flow_id] = existing + + def fake_worker(flow, cfg): + import asyncio + + asyncio.run(flow.publish_authorization_url("https://idp.example/authorize?state=work")) + + with patch("hermes_cli.mcp_config._get_mcp_servers", return_value={"reports": {"url": "https://mcp.example"}}), \ + patch.object(web_server, "_run_dashboard_mcp_oauth", fake_worker): + response = _client().post("/api/mcp/servers/reports/auth?profile=work") + + assert response.status_code != 409 + + +def test_callback_url_is_stable_for_a_server(): + from hermes_cli import web_server + + # The route helper's stable form must not depend on a one-time flow id. + first = web_server._mcp_oauth_callback_url_from_base("https://agent.example", "reports") + second = web_server._mcp_oauth_callback_url_from_base("https://agent.example", "reports") + assert first == second == "https://agent.example/api/mcp/oauth/callback/reports" + + def test_flow_status_does_not_expose_authorization_code(): from hermes_cli import web_server from tools.mcp_dashboard_oauth import DashboardOAuthFlow @@ -190,6 +237,7 @@ def test_flow_status_does_not_expose_authorization_code(): flow_id="flow-status", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/api/mcp/oauth/callback/flow-status", ) flow.authorization_url = "https://idp.example/authorize" diff --git a/tests/tools/test_mcp_dashboard_oauth.py b/tests/tools/test_mcp_dashboard_oauth.py index 116612668604..586eb13ebe9b 100644 --- a/tests/tools/test_mcp_dashboard_oauth.py +++ b/tests/tools/test_mcp_dashboard_oauth.py @@ -1,6 +1,7 @@ """Hosted-dashboard bridge for MCP OAuth browser callbacks.""" import asyncio +import threading import pytest @@ -12,6 +13,7 @@ def test_dashboard_flow_exposes_authorization_url_and_accepts_callback(): flow_id="flow-1", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/mcp/oauth/callback/flow-1", ) @@ -35,6 +37,7 @@ def test_dashboard_flow_rejects_wrong_state_without_consuming_callback(): flow_id="flow-state", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/mcp/oauth/callback/flow-state", ) asyncio.run( @@ -60,6 +63,7 @@ def test_dashboard_flow_rejects_second_callback(): flow_id="flow-2", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/mcp/oauth/callback/flow-2", ) asyncio.run( @@ -72,6 +76,39 @@ def test_dashboard_flow_rejects_second_callback(): flow.deliver_callback(code="second", state="state", error=None) +def test_dashboard_flow_accepts_only_one_concurrent_callback(): + from tools.mcp_dashboard_oauth import DashboardOAuthFlow + + flow = DashboardOAuthFlow( + flow_id="flow-race", + server_name="reports", + profile=None, + hermes_home="/tmp/hermes-test", + redirect_uri="https://agent.example/mcp/oauth/callback/flow-race", + ) + asyncio.run(flow.publish_authorization_url("https://idp.example/authorize?state=state")) + + start = threading.Barrier(3) + outcomes: list[str] = [] + + def deliver(code: str) -> None: + start.wait() + try: + flow.deliver_callback(code=code, state="state", error=None) + outcomes.append("accepted") + except ValueError: + outcomes.append("rejected") + + workers = [threading.Thread(target=deliver, args=(code,)) for code in ("one", "two")] + for worker in workers: + worker.start() + start.wait() + for worker in workers: + worker.join() + + assert sorted(outcomes) == ["accepted", "rejected"] + + def test_dashboard_flow_cannot_resurrect_after_terminal_error(): from tools.mcp_dashboard_oauth import DashboardOAuthFlow @@ -79,6 +116,7 @@ def test_dashboard_flow_cannot_resurrect_after_terminal_error(): flow_id="flow-terminal", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/mcp/oauth/callback/flow-terminal", ) flow.mark_error("start timed out") @@ -105,6 +143,7 @@ def test_dashboard_context_overrides_redirect_and_handlers(): flow_id="flow-3", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/mcp/oauth/callback/flow-3", ) assert get_dashboard_oauth_flow() is None @@ -127,6 +166,7 @@ def test_mcp_oauth_helpers_use_dashboard_flow_without_loopback_port(): flow_id="flow-4", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/mcp/oauth/callback/flow-4", ) cfg = {} @@ -156,6 +196,7 @@ def test_manager_build_allows_dashboard_flow_without_tty(tmp_path, monkeypatch): flow_id="flow-5", server_name="reports", profile=None, + hermes_home="/tmp/hermes-test", redirect_uri="https://agent.example/api/mcp/oauth/callback/flow-5", ) with dashboard_oauth_flow(flow): @@ -177,11 +218,31 @@ def test_manager_evict_preserves_persisted_oauth_state(tmp_path, monkeypatch): '{"access_token":"a","token_type":"Bearer"}' ) manager = MCPOAuthManager() - manager._entries["reports"] = _ProviderEntry( + manager._entries[manager._key("reports")] = _ProviderEntry( server_url="https://mcp.example/mcp", oauth_config={} ) manager.evict("reports") - assert "reports" not in manager._entries + assert manager._key("reports") not in manager._entries assert storage._tokens_path().exists() + + +def test_reconnect_mcp_server_signals_live_task(monkeypatch): + from tools import mcp_tool + + class Event: + called = False + + def set(self): + self.called = True + + class Server: + _reconnect_event = Event() + + server = Server() + monkeypatch.setitem(mcp_tool._servers, "reports", server) + monkeypatch.setattr(mcp_tool, "_mcp_loop", None) + + assert mcp_tool.reconnect_mcp_server("reports") is True + assert server._reconnect_event.called is True diff --git a/tests/tools/test_mcp_oauth_manager.py b/tests/tools/test_mcp_oauth_manager.py index 448400cad85f..d680607aee5b 100644 --- a/tests/tools/test_mcp_oauth_manager.py +++ b/tests/tools/test_mcp_oauth_manager.py @@ -11,6 +11,41 @@ from unittest.mock import MagicMock import pytest + +def test_manager_isolates_same_named_servers_by_profile_home(tmp_path, monkeypatch): + from hermes_constants import reset_hermes_home_override, set_hermes_home_override + from tools.mcp_oauth import HermesTokenStorage + from tools.mcp_oauth_manager import MCPOAuthManager + + profile_a = tmp_path / "profile-a" + profile_b = tmp_path / "profile-b" + for home, access_token in ((profile_a, "TOKEN_A"), (profile_b, "TOKEN_B")): + token = set_hermes_home_override(home) + try: + storage = HermesTokenStorage("shared") + storage._tokens_path().parent.mkdir(parents=True, exist_ok=True) + storage._tokens_path().write_text( + '{"access_token":"%s","token_type":"Bearer","expires_in":3600}' + % access_token + ) + finally: + reset_hermes_home_override(token) + + manager = MCPOAuthManager() + providers = [] + for home in (profile_a, profile_b): + token = set_hermes_home_override(home) + try: + provider = manager.get_or_build_provider("shared", "https://mcp.example/mcp", {}) + asyncio.run(provider._initialize()) + providers.append(provider) + finally: + reset_hermes_home_override(token) + + assert providers[0] is not providers[1] + assert providers[0].context.current_tokens.access_token == "TOKEN_A" + assert providers[1].context.current_tokens.access_token == "TOKEN_B" + pytest.importorskip( "mcp.client.auth.oauth2", reason="MCP SDK 1.26.0+ required for OAuth support", @@ -166,7 +201,7 @@ async def test_handle_401_tracks_inflight_task_to_prevent_gc(tmp_path, monkeypat class _DummyProvider: context = None # forces the can_refresh=False branch - mgr._entries["srv"] = _ProviderEntry( + mgr._entries[mgr._key("srv")] = _ProviderEntry( server_url="https://example.com/mcp", oauth_config=None, provider=_DummyProvider(), @@ -214,7 +249,7 @@ async def test_handle_401_dedup_survives_even_if_task_reference_dropped(tmp_path class _DummyProvider: context = None - mgr._entries["srv"] = _ProviderEntry( + mgr._entries[mgr._key("srv")] = _ProviderEntry( server_url="https://example.com/mcp", oauth_config=None, provider=_DummyProvider(), @@ -266,7 +301,7 @@ def test_manager_fails_fast_noninteractive_without_cached_tokens(tmp_path, monke with pytest.raises(OAuthNonInteractiveError, match="non-interactive"): mgr.get_or_build_provider("linear", "https://mcp.linear.app/mcp", None) - assert mgr._entries["linear"].provider is None + assert mgr._entries[mgr._key("linear")].provider is None # --------------------------------------------------------------------------- diff --git a/tools/mcp_dashboard_oauth.py b/tools/mcp_dashboard_oauth.py index da2fda39bc7a..a02547c63be0 100644 --- a/tools/mcp_dashboard_oauth.py +++ b/tools/mcp_dashboard_oauth.py @@ -23,6 +23,7 @@ class DashboardOAuthFlow: flow_id: str server_name: str profile: str | None + hermes_home: str redirect_uri: str created_at: float = field(default_factory=time.time) status: str = "starting" @@ -34,17 +35,19 @@ class DashboardOAuthFlow: _callback_error: str | None = field(default=None, init=False, repr=False) _authorization_ready: threading.Event = field(default_factory=threading.Event, init=False, repr=False) _callback_ready: threading.Event = field(default_factory=threading.Event, init=False, repr=False) + _lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False) async def publish_authorization_url(self, url: str) -> None: - if self.status in {"approved", "error"}: - raise RuntimeError("OAuth flow already ended") state = parse_qs(urlparse(url).query).get("state", [None])[0] if not state: raise ValueError("OAuth authorization URL did not include state") - self.expected_state = state - self.authorization_url = url - self.status = "authorization_required" - self._authorization_ready.set() + with self._lock: + if self.status in {"approved", "error"}: + raise RuntimeError("OAuth flow already ended") + self.expected_state = state + self.authorization_url = url + self.status = "authorization_required" + self._authorization_ready.set() async def wait_for_authorization_url(self, timeout: float = 30.0) -> str: ready = await asyncio.to_thread(self._authorization_ready.wait, timeout) @@ -61,21 +64,22 @@ class DashboardOAuthFlow: state: str | None, error: str | None, ) -> None: - if self._callback_ready.is_set(): - raise ValueError("OAuth callback already received") - if ( - self.expected_state is None - or state is None - or not secrets.compare_digest(self.expected_state, state) - ): - raise ValueError("OAuth callback state mismatch") - if error: - self._callback_error = error - elif code: - self._callback = (code, state) - else: - self._callback_error = "OAuth callback did not include code or error" - self._callback_ready.set() + with self._lock: + if self._callback_ready.is_set(): + raise ValueError("OAuth callback already received") + if ( + self.expected_state is None + or state is None + or not secrets.compare_digest(self.expected_state, state) + ): + raise ValueError("OAuth callback state mismatch") + if error: + self._callback_error = error + elif code: + self._callback = (code, state) + else: + self._callback_error = "OAuth callback did not include code or error" + self._callback_ready.set() async def wait_for_callback(self, timeout: float = 300.0) -> tuple[str, str | None]: ready = await asyncio.to_thread(self._callback_ready.wait, timeout) @@ -88,23 +92,30 @@ class DashboardOAuthFlow: return self._callback def mark_approved(self) -> None: - self.status = "approved" - self.error = None + with self._lock: + if self.status == "error": + raise RuntimeError("OAuth flow already ended") + self.status = "approved" + self.error = None def mark_error(self, error: str) -> None: - self.status = "error" - self.error = error - self._authorization_ready.set() - self._callback_ready.set() + with self._lock: + if self.status == "approved": + return + self.status = "error" + self.error = error + self._authorization_ready.set() + self._callback_ready.set() def snapshot(self) -> dict: - return { - "flow_id": self.flow_id, - "server_name": self.server_name, - "status": self.status, - "authorization_url": self.authorization_url, - "error": self.error, - } + with self._lock: + return { + "flow_id": self.flow_id, + "server_name": self.server_name, + "status": self.status, + "authorization_url": self.authorization_url, + "error": self.error, + } _current_dashboard_flow: contextvars.ContextVar[DashboardOAuthFlow | None] = ( diff --git a/tools/mcp_oauth.py b/tools/mcp_oauth.py index e50fcbc9c2f5..b7f317fe6395 100644 --- a/tools/mcp_oauth.py +++ b/tools/mcp_oauth.py @@ -990,7 +990,7 @@ def _configure_callback_port( dashboard_flow = get_dashboard_oauth_flow() if dashboard_flow is not None: cfg["_resolved_port"] = 0 - cfg["redirect_uri"] = dashboard_flow.redirect_uri + cfg["redirect_uri"] = cfg.get("redirect_uri") or dashboard_flow.redirect_uri return 0 requested = int(cfg.get("redirect_port", 0)) # Precedence: explicit config port → cached client-registration port → diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index 087a1af91a98..4fabba6b5d23 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -39,6 +39,7 @@ import logging import re import threading from dataclasses import dataclass, field +from pathlib import Path from typing import Any, Optional logger = logging.getLogger(__name__) @@ -137,6 +138,7 @@ def _make_hermes_provider_class() -> Optional[type]: ): super().__init__(*args, **kwargs) self._hermes_server_name = server_name + self._hermes_home = "" # When the client_id comes from config.yaml (pre-registered), an # invalid_client rejection means the *config* is wrong — deleting # client.json would just be re-seeded from config and re-running @@ -391,7 +393,8 @@ def _make_hermes_provider_class() -> Optional[type]: # whatever state the SDK already has. try: await get_manager().invalidate_if_disk_changed( - self._hermes_server_name + self._hermes_server_name, + hermes_home=self._hermes_home, ) except Exception as exc: # pragma: no cover — defensive logger.debug( @@ -449,7 +452,7 @@ class MCPOAuthManager: """ def __init__(self) -> None: - self._entries: dict[str, _ProviderEntry] = {} + self._entries: dict[tuple[str, str], _ProviderEntry] = {} self._entries_lock = threading.Lock() # Holds strong references to in-flight 401 handler tasks so the # event loop's weak-reference bookkeeping cannot GC them mid-run @@ -472,8 +475,9 @@ class MCPOAuthManager: Returns None if the MCP SDK's OAuth support is unavailable. """ + key = self._key(server_name) with self._entries_lock: - entry = self._entries.get(server_name) + entry = self._entries.get(key) if entry is not None and entry.server_url != server_url: logger.info( "MCP OAuth '%s': URL changed from %s to %s, discarding cache", @@ -486,13 +490,25 @@ class MCPOAuthManager: server_url=server_url, oauth_config=oauth_config, ) - self._entries[server_name] = entry + self._entries[key] = entry if entry.provider is None: entry.provider = self._build_provider(server_name, entry) + if entry.provider is not None: + entry.provider._hermes_home = key[0] return entry.provider + @staticmethod + def _key( + server_name: str, + hermes_home: str | Path | None = None, + ) -> tuple[str, str]: + from hermes_constants import get_hermes_home + + home = Path(hermes_home) if hermes_home is not None else get_hermes_home() + return (str(home.expanduser().resolve(strict=False)), server_name) + def _build_provider( self, server_name: str, @@ -566,14 +582,19 @@ class MCPOAuthManager: timeout=float(cfg.get("timeout", 300)), ) - def remove(self, server_name: str) -> None: + def remove( + self, + server_name: str, + *, + hermes_home: str | Path | None = None, + ) -> None: """Evict the provider from cache AND delete tokens from disk. Called by ``hermes mcp remove `` and (indirectly) by ``hermes mcp login `` during forced re-auth. """ with self._entries_lock: - self._entries.pop(server_name, None) + self._entries.pop(self._key(server_name, hermes_home), None) from tools.mcp_oauth import remove_oauth_tokens remove_oauth_tokens(server_name) @@ -582,14 +603,24 @@ class MCPOAuthManager: server_name, ) - def evict(self, server_name: str) -> None: + def evict( + self, + server_name: str, + *, + hermes_home: str | Path | None = None, + ) -> None: """Drop only the in-process provider, preserving persisted OAuth state.""" with self._entries_lock: - self._entries.pop(server_name, None) + self._entries.pop(self._key(server_name, hermes_home), None) # -- Disk watch ---------------------------------------------------------- - async def invalidate_if_disk_changed(self, server_name: str) -> bool: + async def invalidate_if_disk_changed( + self, + server_name: str, + *, + hermes_home: str | Path | None = None, + ) -> bool: """If the tokens file on disk has a newer mtime than last-seen, force the MCP SDK provider to reload its in-memory state. @@ -600,7 +631,7 @@ class MCPOAuthManager: """ from tools.mcp_oauth import _get_token_dir, _safe_filename - entry = self._entries.get(server_name) + entry = self._entries.get(self._key(server_name, hermes_home)) if entry is None or entry.provider is None: return False @@ -647,7 +678,7 @@ class MCPOAuthManager: the same ``failed_access_token``, only one recovery attempt fires. Others await the same future. """ - entry = self._entries.get(server_name) + entry = self._entries.get(self._key(server_name)) if entry is None or entry.provider is None: return False diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index b0c1b976bcfe..dcc5bc9fba58 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -3203,6 +3203,15 @@ def _signal_reconnect(server: Any) -> bool: return True +def reconnect_mcp_server(server_name: str) -> bool: + """Ask a currently-live MCP server to rebuild after external re-auth.""" + with _lock: + server = _servers.get(server_name) + if server is None: + return False + return _signal_reconnect(server) + + def _wait_for_server_session_ready( srv: "MCPServerTask", *, diff --git a/web/src/lib/mcp-dashboard-oauth.test.ts b/web/src/lib/mcp-dashboard-oauth.test.ts index 7de079e0c049..b023b6735fa2 100644 --- a/web/src/lib/mcp-dashboard-oauth.test.ts +++ b/web/src/lib/mcp-dashboard-oauth.test.ts @@ -81,4 +81,34 @@ describe("completeMcpDashboardOAuth", () => { ).rejects.toThrow("popup was blocked"); expect(start).not.toHaveBeenCalled(); }); + + it("fails when the authorization window closes before approval", async () => { + const authWindow = { location: { href: "" }, opener: {}, closed: false } as unknown as Window; + const status = vi.fn().mockImplementation(async () => { + Object.defineProperty(authWindow, "closed", { value: true }); + return { + flow_id: "flow-closed", + server_name: "reports", + status: "authorization_required", + authorization_url: "https://idp.example/authorize", + error: null, + }; + }); + + await expect( + completeMcpDashboardOAuth({ + serverName: "reports", + start: async () => ({ + flow_id: "flow-closed", + server_name: "reports", + status: "authorization_required", + authorization_url: "https://idp.example/authorize", + error: null, + }), + status, + open: vi.fn().mockReturnValue(authWindow), + sleep: async () => {}, + }), + ).rejects.toThrow("authorization window was closed"); + }); }); diff --git a/web/src/lib/mcp-dashboard-oauth.ts b/web/src/lib/mcp-dashboard-oauth.ts index 62bbfb8975ed..fbf1cb8dc542 100644 --- a/web/src/lib/mcp-dashboard-oauth.ts +++ b/web/src/lib/mcp-dashboard-oauth.ts @@ -46,6 +46,9 @@ export async function completeMcpDashboardOAuth({ if (current.status === "error") { throw new Error(current.error || "OAuth authorization failed"); } + if (authWindow.closed) { + throw new Error("OAuth authorization window was closed before completion"); + } await sleep(1000); } }