diff --git a/docs/assets/pr-2363/oauth-required-state.png b/docs/assets/pr-2363/oauth-required-state.png new file mode 100644 index 000000000..a969c3b9f Binary files /dev/null and b/docs/assets/pr-2363/oauth-required-state.png differ diff --git a/src/langbot/pkg/api/http/service/mcp.py b/src/langbot/pkg/api/http/service/mcp.py index 09bc185f7..67e8e14ae 100644 --- a/src/langbot/pkg/api/http/service/mcp.py +++ b/src/langbot/pkg/api/http/service/mcp.py @@ -446,15 +446,19 @@ class MCPService: persisted_session = runtime_mcp_session async def _refresh_and_report() -> None: - needs_start = persisted_session.status == MCPSessionStatus.ERROR or persisted_session.session is None - if needs_start: - await persisted_session.start() - else: - try: - await persisted_session.refresh() - except Exception: + try: + needs_start = ( + persisted_session.status == MCPSessionStatus.ERROR or persisted_session.session is None + ) + if needs_start: await persisted_session.start() - ctx.metadata['runtime_info'] = persisted_session.get_runtime_info_dict() + else: + try: + await persisted_session.refresh() + except Exception: + await persisted_session.start() + finally: + ctx.metadata['runtime_info'] = persisted_session.get_runtime_info_dict() coroutine = _refresh_and_report() else: @@ -471,8 +475,11 @@ class MCPService: async def _run_and_cleanup() -> None: try: await test_session.start() - ctx.metadata['runtime_info'] = test_session.get_runtime_info_dict() finally: + # start() raises for a failed connection. Preserve the + # terminal runtime state so the UI can render actionable + # failure phases such as OAuth-required. + ctx.metadata['runtime_info'] = test_session.get_runtime_info_dict() try: await test_session.shutdown() except Exception as exc: diff --git a/src/langbot/pkg/provider/tools/loaders/mcp.py b/src/langbot/pkg/provider/tools/loaders/mcp.py index 2084bd154..6cde880a3 100644 --- a/src/langbot/pkg/provider/tools/loaders/mcp.py +++ b/src/langbot/pkg/provider/tools/loaders/mcp.py @@ -1,5 +1,6 @@ from __future__ import annotations +import dataclasses import enum import json import math @@ -206,6 +207,13 @@ class MCPSessionStatus(enum.Enum): ERROR = 'error' +@dataclasses.dataclass(frozen=True) +class MCPOAuthChallenge: + """Bearer challenge metadata returned by an OAuth-protected MCP server.""" + + resource_metadata_url: str | None + + class _TransportReconnect(Exception): """Internal signal: the Box stdio WS transport dropped but the managed process is still alive. Triggers a lightweight transport reconnect that @@ -265,6 +273,7 @@ class RuntimeMCPSession: _ready_event: asyncio.Event error_message: str | None = None + _public_error_code: str = 'runtime_error' error_phase: MCPSessionErrorPhase | None = None @@ -510,6 +519,13 @@ class RuntimeMCPSession: await self._init_streamable_http_server() return except Exception as e: + if self._extract_oauth_challenge(e) is not None: + self.error_phase = MCPSessionErrorPhase.OAUTH_REQUIRED + self.ap.logger.info( + f'MCP server {self.server_name}: remote server requires OAuth authorization; ' + 'not falling back to SSE' + ) + raise if not self._should_fallback_to_sse(e): self.ap.logger.info( f'MCP server {self.server_name}: Streamable HTTP transport failed ' @@ -630,6 +646,7 @@ class RuntimeMCPSession: except Exception as e: self.status = MCPSessionStatus.ERROR self.error_message = str(e) + self._public_error_code = self._classify_public_error(e) self.ap.logger.error(f'Error in MCP session lifecycle {self.server_name}: {e}\n{traceback.format_exc()}') # Do NOT set _ready_event here — let _lifecycle_loop_with_retry # handle retries first. It will set the event when all retries @@ -752,6 +769,11 @@ class RuntimeMCPSession: except Exception as e: if self._shutdown_event.is_set(): return # Shutdown requested, don't retry + if self.error_phase == MCPSessionErrorPhase.OAUTH_REQUIRED: + self.retry_count = attempt + 1 + self.status = MCPSessionStatus.ERROR + self._ready_event.set() + return if self.error_phase == MCPSessionErrorPhase.BOX_UNAVAILABLE: box_service = getattr(self.ap, 'box_service', None) if box_service is not None and getattr(box_service, 'enabled', True): @@ -832,6 +854,39 @@ class RuntimeMCPSession: else: yield exc + @staticmethod + def _classify_public_error(exc: BaseException) -> str: + """Expose a safe category without transport URLs, headers, or arguments.""" + for leaf in RuntimeMCPSession._iter_exception_leaves(exc): + if isinstance(leaf, httpx.HTTPStatusError): + return f'http_{leaf.response.status_code}' + if isinstance(leaf, (httpx.TimeoutException, TimeoutError)): + return 'connection_timeout' + if isinstance(leaf, httpx.ConnectError): + return 'connection_unreachable' + return 'runtime_error' + + @staticmethod + def _extract_oauth_challenge(exc: BaseException) -> MCPOAuthChallenge | None: + """Extract an OAuth Bearer challenge from a remote MCP connection failure.""" + for leaf in RuntimeMCPSession._iter_exception_leaves(exc): + if not isinstance(leaf, httpx.HTTPStatusError) or leaf.response.status_code != 401: + continue + for header in leaf.response.headers.get_list('www-authenticate'): + bearer_match = re.search(r'(?:^|,)\s*Bearer(?:\s|,|$)', header, flags=re.IGNORECASE) + if bearer_match is None: + continue + metadata_match = re.search( + r'(?:^|,)\s*resource_metadata\s*=\s*(?:"([^"]+)"|([^,\s]+))', + header[bearer_match.end() :], + flags=re.IGNORECASE, + ) + if metadata_match is None: + continue + resource_metadata_url = metadata_match.group(1) or metadata_match.group(2) + return MCPOAuthChallenge(resource_metadata_url=resource_metadata_url) + return None + @staticmethod def _should_fallback_to_sse(exc: BaseException) -> bool: """Whether a Streamable HTTP failure matches legacy-SSE fallback. @@ -1374,7 +1429,7 @@ class RuntimeMCPSession: # environment values. Detailed diagnostics belong in AUDIT_VIEW # logs; resource-list responses expose only a stable status. 'error_message': 'MCP runtime failed' if self.error_message else None, - 'error_code': 'runtime_error' if self.error_message else None, + 'error_code': self._public_error_code if self.error_message else None, 'error_phase': self.error_phase.value if self.error_phase else None, 'retry_count': self.retry_count, 'tool_count': len(self.get_tools()), diff --git a/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py b/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py index 110fff431..48aa9a227 100644 --- a/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py +++ b/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py @@ -52,6 +52,7 @@ class MCPSessionErrorPhase(enum.Enum): MCP_INIT = 'mcp_init' RUNTIME = 'runtime' TOOL_CALL = 'tool_call' + OAUTH_REQUIRED = 'oauth_required' # Stdio MCP refused because Box is disabled in config or currently # unavailable. Not transient — retries would be pointless. The frontend # uses this phase to render a localized actionable message instead of diff --git a/tests/unit_tests/api/service/test_mcp_service.py b/tests/unit_tests/api/service/test_mcp_service.py index 31c531e98..5cddde1e4 100644 --- a/tests/unit_tests/api/service/test_mcp_service.py +++ b/tests/unit_tests/api/service/test_mcp_service.py @@ -1009,6 +1009,37 @@ class TestMCPServiceTestMCPServer: # Verify - returns task ID assert task_id == 123 + @pytest.mark.parametrize('refresh_first', [False, True]) + async def test_persisted_test_preserves_failure_details(self, refresh_first): + from langbot.pkg.provider.tools.loaders.mcp import MCPSessionStatus + + runtime_info = {'status': 'error', 'error_message': 'HTTP 403: access denied'} + session = SimpleNamespace( + status=MCPSessionStatus.CONNECTED if refresh_first else MCPSessionStatus.ERROR, + session=object(), + refresh=AsyncMock(side_effect=RuntimeError('refresh failed')), + start=AsyncMock(side_effect=RuntimeError('Connection failed, please check URL')), + get_runtime_info_dict=Mock(return_value=runtime_info), + ) + captured = {} + + def create_user_task(coroutine, **kwargs): + captured.update(coroutine=coroutine, context=kwargs['context']) + return SimpleNamespace(id=123) + + ap = SimpleNamespace( + tool_mgr=SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=session))), + task_mgr=SimpleNamespace(create_user_task=Mock(side_effect=create_user_task)), + ) + service = _service(ap) + service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'existing-server'})) + await service.test_mcp_server(_CONTEXT, 'existing-server', {}) + with pytest.raises(RuntimeError, match='Connection failed'): + await captured['coroutine'] + assert captured['context'].metadata['runtime_info'] == runtime_info + session.start.assert_awaited_once() + assert session.refresh.await_count == int(refresh_first) + async def test_test_mcp_server_not_found_raises(self): """Raises ValueError when server not found.""" # Setup @@ -1052,6 +1083,45 @@ class TestMCPServiceTestMCPServer: ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_called_once() assert task_id == 456 + async def test_transient_test_preserves_runtime_info_after_connection_failure(self): + runtime_info = { + 'status': 'error', + 'error_phase': 'oauth_required', + 'retry_count': 1, + } + mock_session = SimpleNamespace( + server_name='oauth-server', + start=AsyncMock(side_effect=RuntimeError('connection failed')), + get_runtime_info_dict=Mock(return_value=runtime_info), + shutdown=AsyncMock(), + ) + ap = SimpleNamespace( + tool_mgr=SimpleNamespace( + mcp_tool_loader=SimpleNamespace(load_mcp_server=AsyncMock(return_value=mock_session)) + ) + ) + captured: dict = {} + + def create_user_task(coroutine, **kwargs): + captured['coroutine'] = coroutine + captured['context'] = kwargs['context'] + return SimpleNamespace(id=457) + + ap.task_mgr = SimpleNamespace(create_user_task=Mock(side_effect=create_user_task)) + service = _service(ap) + + task_id = await service.test_mcp_server( + _CONTEXT, + '_', + {'name': 'OAuth server', 'mode': 'remote', 'enable': True, 'extra_args': {}}, + ) + + assert task_id == 457 + with pytest.raises(RuntimeError, match='connection failed'): + await captured['coroutine'] + assert captured['context'].metadata['runtime_info'] == runtime_info + mock_session.shutdown.assert_awaited_once_with() + async def test_rejected_transient_test_session_is_shut_down(self): ap = SimpleNamespace() mock_session = MagicMock() diff --git a/tests/unit_tests/provider/test_mcp_remote_transport.py b/tests/unit_tests/provider/test_mcp_remote_transport.py index 7abe8afea..91a142134 100644 --- a/tests/unit_tests/provider/test_mcp_remote_transport.py +++ b/tests/unit_tests/provider/test_mcp_remote_transport.py @@ -13,7 +13,15 @@ from aiohttp import web from mcp import types as mcp_types from langbot.pkg.api.http.context import ExecutionContext -from langbot.pkg.provider.tools.loaders.mcp import MCPToolCallTimeoutError, RuntimeMCPSession +from langbot.pkg.provider.tools.loaders.mcp import MCPSessionStatus, MCPToolCallTimeoutError, RuntimeMCPSession +from langbot.pkg.provider.tools.loaders.mcp_stdio import MCPSessionErrorPhase + + +TEST_EXECUTION_CONTEXT = ExecutionContext( + instance_uuid='instance-a', + workspace_uuid='workspace-a', + placement_generation=1, +) TEST_EXECUTION_CONTEXT = ExecutionContext( @@ -24,8 +32,9 @@ TEST_EXECUTION_CONTEXT = ExecutionContext( class _TransportProbe: - def __init__(self, streamable_status: int | None) -> None: + def __init__(self, streamable_status: int | None, streamable_headers: dict[str, str] | None = None) -> None: self.streamable_status = streamable_status + self.streamable_headers = streamable_headers or {} self.streamable_posts = 0 self.streamable_messages: list[str] = [] self.sse_gets = 0 @@ -93,7 +102,7 @@ class _TransportProbe: } ) return web.Response(status=202) - return web.Response(status=self.streamable_status) + return web.Response(status=self.streamable_status, headers=self.streamable_headers) self.sse_gets += 1 response = web.StreamResponse( @@ -136,8 +145,8 @@ class _TransportProbe: @asynccontextmanager -async def _transport_server(streamable_status: int | None): - probe = _TransportProbe(streamable_status) +async def _transport_server(streamable_status: int | None, streamable_headers: dict[str, str] | None = None): + probe = _TransportProbe(streamable_status, streamable_headers) application = web.Application() application.router.add_route('*', '/mcp', probe.handle_mcp_endpoint) application.router.add_post('/messages', probe.handle_sse_message) @@ -265,6 +274,45 @@ async def test_remote_transport_real_non_compatibility_error_does_not_fallback(s await _close_session(session) +def test_remote_transport_extracts_oauth_resource_metadata_from_bearer_challenge(): + request = httpx.Request('POST', 'https://mcp.example/mcp') + response = httpx.Response( + 401, + headers={ + 'WWW-Authenticate': ( + 'Basic realm="MCP", Bearer resource_metadata="https://mcp.example/.well-known/oauth-protected-resource"' + ) + }, + request=request, + ) + + with pytest.raises(httpx.HTTPStatusError) as exc_info: + response.raise_for_status() + + challenge = RuntimeMCPSession._extract_oauth_challenge(exc_info.value) + + assert challenge is not None + assert challenge.resource_metadata_url == 'https://mcp.example/.well-known/oauth-protected-resource' + + +@pytest.mark.asyncio +async def test_remote_transport_oauth_challenge_sets_non_retryable_authorization_state(): + headers = { + 'WWW-Authenticate': 'Bearer resource_metadata="https://mcp.example/.well-known/oauth-protected-resource"' + } + async with _transport_server(401, headers) as (probe, url): + session = _session(url) + + await session._lifecycle_loop_with_retry() + + assert session.status == MCPSessionStatus.ERROR + assert session.error_phase == MCPSessionErrorPhase.OAUTH_REQUIRED + assert session.retry_count == 1 + assert session._ready_event.is_set() + assert probe.streamable_posts == 1 + assert probe.sse_gets == 0 + + @pytest.mark.asyncio async def test_remote_transport_real_timeout_does_not_fallback(): async with _transport_server(None) as (probe, url): @@ -313,3 +361,25 @@ async def test_remote_transport_external_cancellation_is_not_converted_to_sse_fa finally: probe.release_streamable_request.set() await _close_session(session) + + +@pytest.mark.parametrize( + ('error', 'expected'), + [ + (httpx.ConnectError('secret host'), 'connection_unreachable'), + (httpx.ReadTimeout('secret URL'), 'connection_timeout'), + (TimeoutError('secret command'), 'connection_timeout'), + (RuntimeError('secret environment'), 'runtime_error'), + ( + httpx.HTTPStatusError( + 'secret response', + request=httpx.Request('POST', 'https://example.test/?token=secret'), + response=httpx.Response(403), + ), + 'http_403', + ), + ], +) +def test_public_error_category_does_not_expose_exception_details(error, expected): + grouped = ExceptionGroup('secret outer exception', [error]) + assert RuntimeMCPSession._classify_public_error(grouped) == expected diff --git a/web/playwright.config.ts b/web/playwright.config.ts index e15c6ef9e..90990a759 100644 --- a/web/playwright.config.ts +++ b/web/playwright.config.ts @@ -17,7 +17,7 @@ export default defineConfig({ }, ], webServer: { - command: 'pnpm exec vite --host 127.0.0.1 --port 4173', + command: 'corepack pnpm@8.9.2 exec vite --host 127.0.0.1 --port 4173', url: 'http://127.0.0.1:4173', reuseExistingServer: !process.env.CI, timeout: 120_000, diff --git a/web/src/app/home/mcp/components/mcp-form/MCPForm.tsx b/web/src/app/home/mcp/components/mcp-form/MCPForm.tsx index d1e4140a4..d7f71261a 100644 --- a/web/src/app/home/mcp/components/mcp-form/MCPForm.tsx +++ b/web/src/app/home/mcp/components/mcp-form/MCPForm.tsx @@ -8,7 +8,14 @@ import React, { } from 'react'; import { useTranslation } from 'react-i18next'; import type { TFunction } from 'i18next'; -import { Braces, Loader2, Trash2, Wrench, XCircle } from 'lucide-react'; +import { + Braces, + Loader2, + ShieldAlert, + Trash2, + Wrench, + XCircle, +} from 'lucide-react'; import { Resolver, useForm } from 'react-hook-form'; import { zodResolver } from '@hookform/resolvers/zod'; import { z } from 'zod'; @@ -101,7 +108,7 @@ function StatusDisplay({