From ceaff9946f16f1be610630982b451ceca26d20e2 Mon Sep 17 00:00:00 2001 From: Mouhand-Kaddo Date: Tue, 8 Sep 2026 10:05:07 +0400 Subject: [PATCH 1/2] feat: OAuth 2.1 for streamable-HTTP MCP servers (/mcp auth) Interactive login via the SDK's OAuthClientProvider: loopback callback + browser for /mcp auth, cached-credentials-only for startup and reconnects. Credentials live per endpoint under /mcp-auth (0600/atomic). Stored tokens are stamped with their acquisition time so an expired access token is refreshed silently on restart (the SDK restores stored tokens but not their absolute expiry; seeding it lets its own refresh path run). auth_required is distinct from failed; AS-controlled error text is control-char-stripped and capped; logout drops session + credentials. The OAuth-E2E-then-plain-HTTP test interference was sse-starlette's process-global AppStatus.should_exit: its shutdown watcher captures the uvicorn server from the signal table and, when a test fixture's teardown sets should_exit, flips the global flag and kills every later SSE response at birth. Fixed per-test in tests/conftest.py and pinned by an ordered regression pair in tests/test_mcp.py. --- docs/mcp.md | 33 +- src/lecode/agent/tools/base.py | 4 + src/lecode/config/models.py | 4 + src/lecode/extras/mcp_auth.py | 382 ++++++++++++++++++ src/lecode/extras/mcp_client.py | 241 +++++++++++- src/lecode/slash/handlers.py | 53 ++- src/lecode/tui/loading.py | 4 +- tests/conftest.py | 22 ++ tests/test_mcp.py | 67 +++- tests/test_mcp_auth.py | 658 ++++++++++++++++++++++++++++++++ 10 files changed, 1433 insertions(+), 35 deletions(-) create mode 100644 src/lecode/extras/mcp_auth.py create mode 100644 tests/test_mcp_auth.py diff --git a/docs/mcp.md b/docs/mcp.md index 4daa3a6..ad87276 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -36,8 +36,34 @@ headers = { Authorization = "Bearer …" } A `[mcp.servers]` entry named `exa` or `context7` replaces the auto-configured definition. -OAuth for HTTP servers is intentionally not wired up (it needs an interactive -browser flow); pass bearer tokens via `headers` instead. +## OAuth 2.1 + +HTTP servers that advertise OAuth (like GlitchTip) authenticate interactively: + +```toml +[mcp.servers.glitchtip] +transport = "http" +url = "https://your-glitchtip.example.com/mcp" +auth = "oauth" +``` + +- On startup lecode connects with **cached credentials only** — an expired + access token is refreshed silently from its refresh token; a missing or + rejected refresh shows `authentication required`. Startup never blocks on + a browser. +- `/mcp auth glitchtip` runs the interactive login: it opens your browser, + the server walks you through approval, and the redirect lands back on a + loopback port lecode serves (`http://127.0.0.1:/callback`). +- `/mcp logout glitchtip` drops the session and the persisted credentials. +- Credentials (access + refresh tokens, client registration) are stored per + endpoint under `~/.config/lecode/mcp-auth/` (0600 files, 0700 directory, + plaintext JSON). `LECODE_CONFIG_DIR` moves them. + +The protocol itself — resource/server metadata discovery, dynamic client +registration, PKCE, token exchange and refresh - is handled by the SDK's +OAuth client; lecode supplies storage, the browser step, and the loopback +callback. `auth = "oauth"` conflicts with a static `Authorization` header +(that header path is the alternative for servers without OAuth). ## Permissions @@ -68,6 +94,9 @@ Rule targets for MCP tools are the canonical `mcp::` name. ``` /mcp per-server state: connected (n tools) / failed / disabled + / authentication required /mcp tools list one server's tools /mcp reconnect drop and re-establish a server session +/mcp auth interactive OAuth login (opens the browser) +/mcp logout drop a server session and its stored credentials ``` diff --git a/src/lecode/agent/tools/base.py b/src/lecode/agent/tools/base.py index d5754b9..6b5c780 100644 --- a/src/lecode/agent/tools/base.py +++ b/src/lecode/agent/tools/base.py @@ -83,6 +83,10 @@ def __init__(self, tools: list[Tool] | None = None) -> None: def register(self, tool: Tool) -> None: self._tools[tool.name] = tool + def unregister(self, name: str) -> None: + """Drop one tool by exact name (no-op when unknown).""" + self._tools.pop(name, None) + def get(self, name: str) -> Tool | None: return self._tools.get(name) diff --git a/src/lecode/config/models.py b/src/lecode/config/models.py index 1ccc491..83bd09f 100644 --- a/src/lecode/config/models.py +++ b/src/lecode/config/models.py @@ -165,6 +165,10 @@ class McpServerConfig(BaseModel): # http url: str | None = None headers: dict[str, str] = Field(default_factory=dict) + #: ``"oauth"`` enables the SDK's OAuth 2.1 flow (discovery, dynamic client + #: registration, PKCE). ``None`` keeps static ``headers`` (bearer token) + #: authentication. Only meaningful with ``transport = "http"``. + auth: Literal["oauth"] | None = None # common timeout_s: float = 30.0 enabled: bool = True diff --git a/src/lecode/extras/mcp_auth.py b/src/lecode/extras/mcp_auth.py new file mode 100644 index 0000000..42b60cb --- /dev/null +++ b/src/lecode/extras/mcp_auth.py @@ -0,0 +1,382 @@ +"""OAuth 2.1 support for streamable-HTTP MCP servers. + +The ``mcp`` SDK's ``OAuthClientProvider`` implements the whole protocol +(protected-resource + authorization-server discovery, dynamic client +registration, PKCE, token exchange, refresh). This module supplies the two +application pieces it needs: + +- :class:`FileTokenStorage`: persists tokens + client registration per + endpoint under ``/mcp-auth`` (0600 files, atomic replace). +- :class:`LoopbackAuthCallback`: serves the authorization-code callback on + ``127.0.0.1`` and opens the system browser for interactive ``/mcp auth``. + +Non-interactive connections pass no redirect/callback handlers, so the SDK +reuses cached credentials (including refresh) and never blocks on a browser. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import hashlib +import json +import logging +import os +import tempfile +import time +import webbrowser +from pathlib import Path +from typing import TYPE_CHECKING, Any +from urllib.parse import parse_qs, urlsplit + +from pydantic import AnyUrl + +from lecode.config.loader import config_dir + +if TYPE_CHECKING: # deferred like the rest of the codebase; the SDK is heavy + from mcp.client.auth import AuthorizationCodeResult, OAuthClientProvider + from mcp.shared.auth import OAuthClientInformationFull, OAuthToken + +log = logging.getLogger(__name__) + +#: How long one interactive login may sit waiting on the browser (s). +CALLBACK_TIMEOUT_S = 180.0 + +_AUTH_DIR_NAME = "mcp-auth" + +#: Serializes file updates in-process so concurrent set_tokens/set_client_info +#: calls cannot clobber each other's read-modify-write. +#: ponytail: one global lock; per-endpoint locks if two servers ever contend. +_storage_lock = asyncio.Lock() + +_REDIRECT_CALLBACK = "/callback" + + +# -- helpers ------------------------------------------------------------------ + + +def _redirect_port(server_url: str) -> int: + """Deterministic loopback port per endpoint (stable across sessions). + + A stable port keeps the redirect_uri registered at first login valid on + every later login. Bind failure is self-healing in + :meth:`LoopbackAuthCallback.open` (fresh dynamic registration). + """ + digest = int.from_bytes(hashlib.sha256(server_url.encode()).digest()[:2]) + return 40000 + digest % 20000 + + +def _auth_dir() -> Path: + return config_dir() / _AUTH_DIR_NAME + + +def mcp_auth_file(server_url: str) -> Path: + """The credentials file for one endpoint (public for tests/docs).""" + digest = hashlib.sha256(server_url.encode()).hexdigest()[:32] + return _auth_dir() / f"{digest}.json" + + +class FileTokenStorage: + """Persistence for the SDK's :class:`TokenStorage` protocol. + + One JSON file per endpoint: ``{server_url, tokens, client_info}``. Files + are 0600 from creation, written atomically (temp file + replace); the + directory is 0700. Contents are plaintext JSON. + """ + + def __init__(self, server_url: str) -> None: + self.server_url = server_url + self.path = mcp_auth_file(server_url) + + # -- protocol ----------------------------------------------------------- + + async def get_tokens(self) -> OAuthToken | None: + from mcp.shared.auth import OAuthToken + + data = await self._read() + if data is None or data.get("tokens") is None: + return None + try: + return OAuthToken.model_validate_json(json.dumps(data["tokens"])) + except (ValueError, TypeError) as e: + log.debug("mcp-auth: discarding invalid tokens in %s: %s", self.path, e) + return None + + async def set_tokens(self, tokens: OAuthToken) -> None: + # Stamp when the tokens were acquired: the SDK does not persist (or + # restore) the absolute expiry itself, so this is the only way to know + # later that a stored access token has expired (see token_expiry()). + await self._update(tokens=tokens.model_dump(mode="json"), tokens_acquired_at=time.time()) + + async def get_client_info(self) -> OAuthClientInformationFull | None: + from mcp.shared.auth import OAuthClientInformationFull + + data = await self._read() + if data is None or data.get("client_info") is None: + return None + try: + return OAuthClientInformationFull.model_validate_json(json.dumps(data["client_info"])) + except (ValueError, TypeError) as e: + log.debug("mcp-auth: discarding invalid client_info in %s: %s", self.path, e) + return None + + async def set_client_info(self, client_info: OAuthClientInformationFull) -> None: + await self._update(client_info=client_info.model_dump(mode="json")) + + # -- management ---------------------------------------------------------- + + async def clear(self) -> None: + """Drop all persisted credentials for this endpoint (logout).""" + async with _storage_lock: + + def do_unlink() -> bool: + try: + self.path.unlink() + return True + except FileNotFoundError: + return False + + removed = await asyncio.to_thread(do_unlink) + if removed: + log.debug("mcp-auth: cleared %s", self.path) + + async def clear_client_info(self) -> None: + """Drop only the client registration (forces a fresh one next login).""" + await self._update(client_info=None) + + async def token_expiry(self) -> float | None: + """When the stored access token expires (Unix time); None if unknown. + + None means "no stored tokens", "no expiry info", or a file from before + acquisition stamping — all map to the SDK's own unknown-expiry + behavior (attach and let the server judge). + """ + data = await self._read() + if data is None: + return None + acquired = data.get("tokens_acquired_at") + expires_in = (data.get("tokens") or {}).get("expires_in") + if isinstance(acquired, bool) or not isinstance(acquired, (int, float)): + return None + try: + return acquired + int(expires_in) + except (TypeError, ValueError): + return None + + # -- internals ----------------------------------------------------------- + + async def _read(self) -> dict[str, Any] | None: + def do_read() -> dict[str, Any] | None: + try: + data = json.loads(self.path.read_text(encoding="utf-8")) + except FileNotFoundError: + return None + except (OSError, UnicodeDecodeError, json.JSONDecodeError) as e: + log.debug("mcp-auth: unreadable credentials file %s: %s", self.path, e) + return None + if not isinstance(data, dict): + return None + # fail-safe rebinding: never serve this endpoint another URL's + # credentials (only reachable on a sha256 collision). + if data.get("server_url") != self.server_url: + log.warning("mcp-auth: %s is bound to a different endpoint — ignoring", self.path) + return None + return data + + return await asyncio.to_thread(do_read) + + async def _update(self, **fields: Any) -> None: + async with _storage_lock: + data = await self._read() + merged = data or {} + merged["server_url"] = self.server_url + merged.update(fields) + + def do_write() -> None: + _auth_dir().mkdir(mode=0o700, parents=True, exist_ok=True) + with contextlib.suppress(OSError): # pragma: no cover — best effort + _auth_dir().chmod(0o700) + # mkstemp → 0600 by default; os.replace keeps that mode. + fd, tmp_path = tempfile.mkstemp(dir=_auth_dir(), prefix=".tmp-") + try: + with os.fdopen(fd, "w", encoding="utf-8") as fh: + json.dump(merged, fh) + fh.flush() + os.fsync(fh.fileno()) + os.replace(tmp_path, self.path) + finally: + if os.path.exists(tmp_path): # pragma: no branch + with contextlib.suppress(OSError): # pragma: no cover + os.unlink(tmp_path) + + await asyncio.to_thread(do_write) + + +# -- interactive loopback callback -------------------------------------------- + + +class LoopbackAuthCallback: + """Hosts the authorization-code callback on 127.0.0.1 for one login.""" + + def __init__( + self, + server_url: str, + server: asyncio.Server, + redirect_url: str, + result: asyncio.Future[dict[str, list[str]]], + accepted: set[asyncio.StreamWriter], + ) -> None: + self.server_url = server_url + self._server = server + self.redirect_url = redirect_url + self._result = result + self._accepted = accepted + + @classmethod + async def open(cls, storage: FileTokenStorage) -> LoopbackAuthCallback: + """Bind the callback listener, reusing the port registered previously.""" + port = _redirect_port(storage.server_url) + client_info = await storage.get_client_info() + registered = client_info.redirect_uris if client_info is not None else None + if registered: + registered_port = registered[0].port + if registered_port is not None: + port = registered_port + loop = asyncio.get_running_loop() + result: asyncio.Future[dict[str, list[str]]] = loop.create_future() + accepted: set[asyncio.StreamWriter] = set() + + async def handle(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + accepted.add(writer) + try: + try: + line = await asyncio.wait_for(reader.readline(), timeout=60) + except (TimeoutError, ValueError, EOFError): # slow, overlong, or broken client + line = b"" + # "GET /callback?... HTTP/1.1" — one request per login; junk 404s + parts = line.decode(errors="replace").split(" ") + target = parts[1] if len(parts) > 1 else "" + try: + ok = bool(target) and urlsplit(target).path == _REDIRECT_CALLBACK + except ValueError: # e.g. a malformed IPv6 literal in the target + ok = False + body = ( + b"

lecode

" + b"Authorization received. You can close this window and return to lecode." + b"

" + ) + headers = ( + f"HTTP/1.1 {200 if ok else 404} {'OK' if ok else 'Not Found'}\r\n" + "Content-Type: text/html; charset=utf-8\r\n" + f"Content-Length: {len(body)}\r\n" + "Connection: close\r\n" + "\r\n" + ).encode() + writer.write(headers + body) + try: + await writer.drain() + finally: + writer.close() + if ok and not result.done(): + result.set_result(parse_qs(urlsplit(target).query)) + except Exception: # pragma: no cover — never raise into asyncio's callback task + log.debug("mcp-auth: callback connection failed", exc_info=True) + with contextlib.suppress(Exception): + writer.close() + finally: + accepted.discard(writer) + + try: + server = await asyncio.start_server(handle, "127.0.0.1", port) + except OSError: + # The stable port is taken (rare) — the old registration is then + # useless, so drop it and register afresh on a random free port. + log.debug("mcp-auth: callback port %d busy; re-registering on a free port", port) + await storage.clear_client_info() + server = await asyncio.start_server(handle, "127.0.0.1", 0) + actual_port = server.sockets[0].getsockname()[1] + redirect_url = f"http://127.0.0.1:{actual_port}{_REDIRECT_CALLBACK}" + return cls(storage.server_url, server, redirect_url, result, accepted) + + async def open_browser(self, authorization_url: str) -> None: + """The SDK's redirect_handler: hand the URL to the system browser.""" + from mcp.client.auth import OAuthFlowError + + opened = await asyncio.to_thread(webbrowser.open, authorization_url) + if not opened: + raise OAuthFlowError( + "could not open a browser on this machine — run /mcp auth where one exists" + ) + + async def wait_for_callback(self) -> AuthorizationCodeResult: + """The SDK's callback_handler: yield the redirect parameters.""" + from mcp.client.auth import AuthorizationCodeResult, OAuthFlowError + + try: + params = await asyncio.wait_for(self._result, timeout=CALLBACK_TIMEOUT_S) + except TimeoutError: + raise OAuthFlowError( + "timed out waiting for the authorization redirect " + f"(limit {int(CALLBACK_TIMEOUT_S)}s)" + ) from None + error = params.get("error") or [None] + if error[0]: + description = params.get("error_description") or ["(no description)"] + raise OAuthFlowError(f"authorization denied: {error[0]} — {description[0]}") + code = params.get("code") or [None] + if not code[0]: + raise OAuthFlowError("the authorization redirect carried no code") + return AuthorizationCodeResult( + code=code[0], + state=(params.get("state") or [None])[0], + iss=(params.get("iss") or [None])[0], + ) + + async def aclose(self) -> None: + """Stop serving callbacks; no-op when the flow already completed.""" + self._server.close() + # Close any still-open accepted connection so its handler task and the + # peer both see the shutdown instead of idling until the 60s read timeout. + for writer in list(self._accepted): + with contextlib.suppress(Exception): + writer.close() + await self._server.wait_closed() + if not self._result.done(): + self._result.cancel() + + +async def make_oauth_provider( + server_url: str, + storage: FileTokenStorage, + loopback: LoopbackAuthCallback | None, +) -> OAuthClientProvider: + """Build the SDK auth handler for one endpoint. + + ``loopback`` is ``None`` for automatic connections (cached credentials + and refresh only — no browser); interactive ``/mcp auth`` passes one. + """ + from mcp.client.auth import OAuthClientProvider + from mcp.shared.auth import OAuthClientMetadata + + redirect_url = ( + loopback.redirect_url + if loopback is not None + else (f"http://127.0.0.1:{_redirect_port(server_url)}{_REDIRECT_CALLBACK}") + ) + metadata = OAuthClientMetadata(client_name="lecode", redirect_uris=[AnyUrl(redirect_url)]) + provider = OAuthClientProvider( + server_url=server_url, + client_metadata=metadata, + storage=storage, + redirect_handler=loopback.open_browser if loopback is not None else None, + callback_handler=loopback.wait_for_callback if loopback is not None else None, + ) + # The SDK restores stored tokens but not their absolute expiry, so an + # expired access token would be attached, 401'd, and the automatic flow + # would dead-end into "authentication required" even with a good refresh + # token on disk. Seed the expiry it failed to restore and its own + # proactive refresh path runs — silently, browser-free. + expiry = await storage.token_expiry() + if expiry is not None and expiry <= time.time(): + provider.context.token_expiry_time = expiry + return provider diff --git a/src/lecode/extras/mcp_client.py b/src/lecode/extras/mcp_client.py index 409d3b0..e01b3af 100644 --- a/src/lecode/extras/mcp_client.py +++ b/src/lecode/extras/mcp_client.py @@ -16,9 +16,13 @@ header — Exa's exact header scheme is not pinned down in their docs. - **context7** (default off): ``https://mcp.context7.com/mcp``, no auth. -OAuth for HTTP MCP servers is intentionally out of scope: the SDK supports -it but it needs an interactive browser flow; bearer tokens via -``[mcp.servers.].headers`` are the supported path. +OAuth 2.1 for HTTP servers (``[mcp.servers.] auth = "oauth"``) +rides on the SDK's ``OAuthClientProvider`` (see ``lecode.extras.mcp_auth``): +automatic connections reuse cached credentials and refresh them silently; +interactive login is explicit via ``/mcp auth ``, which opens the +browser and serves the redirect on a loopback port. Static bearer tokens via +``[mcp.servers.].headers`` remain the supported path for servers that +do not speak OAuth. Failed tool calls get exactly one reconnect attempt (fresh session), then an error result. Everything is fail-open: MCP trouble never blocks the agent. @@ -31,10 +35,11 @@ import logging import os from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Literal from lecode.agent.tools.base import Tool, ToolContext, ToolRegistry, ToolResult from lecode.config.models import Config, McpServerConfig +from lecode.extras.mcp_auth import CALLBACK_TIMEOUT_S if TYPE_CHECKING: from mcp import ClientSession @@ -51,6 +56,10 @@ #: Per-server startup budget. CONNECT_TIMEOUT_S = 10.0 +#: Budget for one interactive ``/mcp auth`` login: the loopback callback +#: timeout plus a margin for discovery, registration, and token exchange. +INTERACTIVE_AUTH_BUDGET_S = CALLBACK_TIMEOUT_S + 30.0 + def auto_servers(config: Config, env: dict[str, str] | None = None) -> dict[str, McpServerConfig]: """The auto-configured servers implied by the ``[mcp]`` flags.""" @@ -85,10 +94,15 @@ class ServerStatus: """One server's state for ``/mcp``.""" name: str - state: str # "connected" | "failed" | "disabled" + state: Literal["connected", "failed", "disabled", "auth_required"] tools: int = 0 error: str | None = None + @property + def auth_hint(self) -> str: + """The canonical line for the ``auth_required`` state.""" + return f"{self.name}: authentication required — /mcp auth {self.name}" + @dataclass class _Server: @@ -98,16 +112,60 @@ class _Server: stack: contextlib.AsyncExitStack | None = None tools: list[McpToolDef] = field(default_factory=list) error: str | None = None + auth_required: bool = False + interactive: bool = False @property def status(self) -> ServerStatus: if not self.config.enabled: return ServerStatus(self.name, "disabled") if self.session is None: + if self.auth_required: + return ServerStatus(self.name, "auth_required", error=self.error) return ServerStatus(self.name, "failed", error=self.error) return ServerStatus(self.name, "connected", tools=len(self.tools)) +def _clean_error(e: BaseException) -> str: + """One short line; OAuth errors can wrap server HTML/JSON bodies. + + Server-controlled text (OAuth error descriptions, response bodies) must + never reach the terminal raw: collapse whitespace, drop non-printable + characters (ANSI/control escapes), and cap the length. + """ + text = "".join(ch for ch in " ".join(str(e).split()) if ch.isprintable()) + return text[:300] + ("…" if len(text) > 300 else "") + + +def _unwrap_exceptions(e: BaseException) -> list[BaseException]: + """Flatten ExceptionGroups (anyio wraps transport task errors in them).""" + if isinstance(e, BaseExceptionGroup): + flat: list[BaseException] = [] + for sub in e.exceptions: + flat.extend(_unwrap_exceptions(sub)) + return flat + return [e] + + +def _oauth_leaf(e: BaseException) -> BaseException | None: + """The OAuthFlowError buried in ``e`` (or an ExceptionGroup), if any.""" + from mcp.client.auth import OAuthFlowError + + return next((x for x in _unwrap_exceptions(e) if isinstance(x, OAuthFlowError)), None) + + +def _connect_failure(e: BaseException) -> tuple[str, bool]: + """(error text, needs interactive login) for one failed connection attempt. + + OAuth servers without (working) credentials must not look like a generic + outage: ``/mcp auth`` is the fix, so they get the auth_required state. + """ + oauth_leaf = _oauth_leaf(e) + if oauth_leaf is not None: + return _clean_error(oauth_leaf), True + return f"{type(e).__name__}: {_clean_error(e)}", False + + def _result_text(result: Any) -> str: """Flatten a CallToolResult's content parts to plain text.""" parts = [getattr(part, "text", "") for part in result.content or []] @@ -121,6 +179,7 @@ def __init__(self, config: Config, ctx: ToolContext | None = None) -> None: self._config = config self._ctx = ctx self._servers: dict[str, _Server] = {} + self._registry: ToolRegistry | None = None self._closed = False # -- connection ------------------------------------------------------------ @@ -139,43 +198,76 @@ async def _connect_one(self, server: _Server) -> None: try: await asyncio.wait_for(self._open(server), timeout=CONNECT_TIMEOUT_S) server.error = None + server.auth_required = False log.debug("mcp: %s connected (%d tools)", server.name, len(server.tools)) except Exception as e: server.session = None - server.error = f"{type(e).__name__}: {e}" - log.debug("mcp: %s connect failed: %s", server.name, e) + server.error, server.auth_required = _connect_failure(e) + if server.auth_required: + log.debug("mcp: %s needs OAuth authentication: %s", server.name, e) + else: + log.debug("mcp: %s connect failed: %s", server.name, e) async def _open(self, server: _Server) -> None: from mcp import ClientSession stack = contextlib.AsyncExitStack() try: - read, write = await self._open_transport(server, stack) + http_client = None + if server.config.auth == "oauth" and server.config.transport == "http": + http_client = await self._build_http_client(server, server.config, stack) + await self._preflight_auth(server, http_client) + read, write = await self._open_transport(server, stack, http_client) session = await stack.enter_async_context( ClientSession(read, write, read_timeout_seconds=server.config.timeout_s) ) await session.initialize() result = await session.list_tools() - except Exception: + except BaseException: + # BaseException (not Exception) so a cancelled interactive login + # still closes the loopback listener and HTTP client. await stack.aclose() raise server.session = session server.stack = stack server.tools = list(result.tools) + async def _preflight_auth(self, server: _Server, http_client: Any) -> None: + """One bare POST so the OAuth middleware completes the whole flow here. + + Doing the flow inside ``session.initialize()`` would lose the error: + the SDK's streamable transport dispatches requests in a task group and + answers a failed POST by cancelling the waiting sender, so the + OAuthFlowError would surface only in logs, not to the caller. A plain + httpx call propagates it directly; the status/body are irrelevant — + the auth middleware already attached a token or raised. + """ + from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER + from mcp_types import LATEST_PROTOCOL_VERSION + + await http_client.post( + server.config.url, + headers={ + MCP_PROTOCOL_VERSION_HEADER: LATEST_PROTOCOL_VERSION, + "Accept": "application/json, text/event-stream", + }, + json={}, + ) + async def _open_transport( - self, server: _Server, stack: contextlib.AsyncExitStack + self, + server: _Server, + stack: contextlib.AsyncExitStack, + http_client: Any | None = None, ) -> tuple[Any, Any]: config = server.config if config.transport == "http": if not config.url: raise ValueError(f"mcp server {server.name}: http transport needs a url") - from mcp.client.streamable_http import ( - create_mcp_http_client, - streamable_http_client, - ) + from mcp.client.streamable_http import streamable_http_client - http_client = create_mcp_http_client(headers=config.headers or None) + if http_client is None: + http_client = await self._build_http_client(server, config, stack) return await stack.enter_async_context( streamable_http_client(config.url, http_client=http_client) ) @@ -194,6 +286,40 @@ async def _open_transport( stack.callback(devnull.close) return await stack.enter_async_context(stdio_client(params, errlog=devnull)) + async def _build_http_client( + self, server: _Server, config: McpServerConfig, stack: contextlib.AsyncExitStack + ) -> Any: + """The httpx client for one server, with OAuth attached when configured.""" + from mcp.client.streamable_http import create_mcp_http_client + + url = config.url + if not url: + raise ValueError(f"mcp server {server.name}: http transport needs a url") + headers = config.headers or None + if config.auth != "oauth": + return create_mcp_http_client(headers=headers) + if config.transport != "http": + raise ValueError(f"mcp server {server.name}: auth='oauth' requires transport='http'") + if any(key.lower() == "authorization" for key in (config.headers or {})): + raise ValueError( + f"mcp server {server.name}: auth='oauth' conflicts with a static " + "Authorization header — drop the header, OAuth manages Authorization" + ) + from lecode.extras.mcp_auth import ( + FileTokenStorage, + LoopbackAuthCallback, + make_oauth_provider, + ) + + storage = FileTokenStorage(url) + loopback = None + if server.interactive: + loopback = await LoopbackAuthCallback.open(storage) + stack.push_async_callback(loopback.aclose) + return create_mcp_http_client( + headers=headers, auth=await make_oauth_provider(url, storage, loopback) + ) + async def _close_server(self, server: _Server) -> None: if server.stack is not None: try: @@ -206,15 +332,33 @@ async def _close_server(self, server: _Server) -> None: # -- tools ----------------------------------------------------------------- - def tool_wrappers(self) -> list[Tool]: - """lecode tools wrapping every connected server's discovered tools.""" + def tool_wrappers(self, server_name: str | None = None) -> list[Tool]: + """lecode tools wrapping connected servers' discovered tools.""" return [ McpTool(self, server.name, tool) for server in self._servers.values() - if server.session is not None + if server.session is not None and (server_name is None or server.name == server_name) for tool in server.tools ] + def _sync_tools(self, name: str | None = None) -> None: + """Align the registry with reality: register new wrappers, drop stale ones. + + Called after every (re)connect and after logout. ``None`` covers the + initial attach; a server name covers one-server reconnect/auth. + """ + if self._registry is None: + return + affected = [sn for sn in self._servers if name in (None, sn)] + prefixes = tuple(f"mcp:{sn}:" for sn in affected) + wrappers = self.tool_wrappers(name) + wanted = {tool.name for tool in wrappers} + for existing in self._registry.names(): + if existing.startswith(prefixes) and existing not in wanted: + self._registry.unregister(existing) + for wrapper in wrappers: + self._registry.register(wrapper) + async def call(self, server_name: str, tool_name: str, args: dict[str, Any]) -> ToolResult: """Call an MCP tool; one reconnect attempt on connection failure.""" server = self._servers.get(server_name) @@ -229,6 +373,7 @@ async def call(self, server_name: str, tool_name: str, args: dict[str, Any]) -> log.debug("mcp: %s:%s call failed, reconnecting: %s", server_name, tool_name, e) await self._close_server(server) await self._connect_one(server) + self._sync_tools(server_name) if server.session is None: return ToolResult( f"error: MCP server '{server_name}' unavailable: {server.error or 'not connected'}", @@ -257,6 +402,10 @@ def status(self) -> list[ServerStatus]: """Per-server state, sorted by name.""" return sorted((server.status for server in self._servers.values()), key=lambda s: s.name) + def bind_registry(self, registry: ToolRegistry) -> None: + """Give the manager the registry it keeps in sync (reconnect/auth/logout).""" + self._registry = registry + def server_tools(self, name: str) -> list[McpToolDef] | None: """One server's discovered tools; None when unknown.""" server = self._servers.get(name) @@ -270,6 +419,55 @@ async def reconnect(self, name: str) -> ServerStatus | None: await self._close_server(server) if server.config.enabled: await self._connect_one(server) + self._sync_tools(name) + return server.status + + async def authenticate(self, name: str) -> ServerStatus | None: + """Interactive OAuth login for one server (opens the browser). + + Runs outside the normal per-server connect budget: the user needs time + to approve in the browser. Non-OAuth servers get a plain error state. + """ + server = self._servers.get(name) + if server is None: + return None + if server.config.auth != "oauth": + server.error = "this server is not configured with auth = 'oauth'" + return server.status + await self._close_server(server) + server.error = None + server.auth_required = False + server.interactive = True + try: + await asyncio.wait_for(self._open(server), timeout=INTERACTIVE_AUTH_BUDGET_S) + except asyncio.CancelledError: + raise + except Exception as e: + server.session = None + server.error, server.auth_required = _connect_failure(e) + if not server.auth_required and isinstance(e, TimeoutError): + server.error = f"timed out (limit {INTERACTIVE_AUTH_BUDGET_S}s)" + log.debug("mcp: %s authentication failed: %s", name, e) + finally: + server.interactive = False + self._sync_tools(name) + return server.status + + async def logout(self, name: str) -> ServerStatus | None: + """Drop one server's session and its persisted OAuth credentials.""" + server = self._servers.get(name) + if server is None: + return None + if server.config.auth != "oauth" or not server.config.url: + server.error = "logout applies only to OAuth servers" + return server.status + await self._close_server(server) + from lecode.extras.mcp_auth import FileTokenStorage + + await FileTokenStorage(server.config.url).clear() + server.error = None + server.auth_required = True + self._sync_tools(name) return server.status async def shutdown(self) -> None: @@ -303,14 +501,15 @@ async def attach_mcp(registry: ToolRegistry, ctx: ToolContext) -> McpManager: Never raises: per-server failures are isolated inside the manager, and a wholesale failure just leaves zero MCP tools registered. The manager is - always installed under ``ctx.extras["mcp"]`` so ``/mcp`` can report. + always installed under ``ctx.extras["mcp"]`` so ``/mcp`` can report, and + keeps the registry handle so later reconnects/auth/logout can update it. """ manager = McpManager(ctx.config, ctx) + manager.bind_registry(registry) ctx.extras[MCP_EXTRA] = manager try: await manager.connect() except Exception as e: # belt-and-braces; connect() isolates per server log.debug("mcp: connect failed: %s", e) - for wrapper in manager.tool_wrappers(): - registry.register(wrapper) + manager._sync_tools() return manager diff --git a/src/lecode/slash/handlers.py b/src/lecode/slash/handlers.py index d598a86..2e1e4e5 100644 --- a/src/lecode/slash/handlers.py +++ b/src/lecode/slash/handlers.py @@ -664,6 +664,8 @@ def report(status: str, text: str) -> None: report("ok", f"mcp {s.name}: connected · {s.tools} tools") elif s.state == "failed": report("warn", f"mcp {s.name}: failed — {s.error or 'connect error'}") + elif s.state == "auth_required": + report("warn", f"mcp {s.auth_hint}") else: report("skip", f"mcp {s.name}: disabled") else: @@ -975,7 +977,7 @@ async def cmd_chain(app: TuiApp, args: list[str]) -> None: async def cmd_mcp(app: TuiApp, args: list[str]) -> None: - """``/mcp`` — server states; ``/mcp tools ``; ``/mcp reconnect ``.""" + """``/mcp`` — states; ``tools|reconnect|auth|logout ``.""" manager = app.runtime.ctx.extras.get("mcp") if manager is None or not manager.status(): app.feed.info("no MCP servers configured") @@ -988,15 +990,46 @@ async def cmd_mcp(app: TuiApp, args: list[str]) -> None: if status is None: app.feed.error(f"unknown MCP server: {args[1]}") return - # Pick up tools that (re)appeared; wrappers delegate by name, so the - # existing registrations survive reconnects unchanged. - for wrapper in manager.tool_wrappers(): - app.runtime.registry.register(wrapper) if status.state == "connected": app.feed.info(f"mcp: {status.name} reconnected ({status.tools} tools)") else: app.feed.error(f"mcp: {status.name} reconnect failed: {status.error}") return + if args and args[0] == "auth": + if len(args) < 2: + app.feed.error("usage: /mcp auth ") + return + name = args[1] + if manager.server_tools(name) is None: + app.feed.error(f"unknown MCP server: {name}") + return + app.feed.info(f"mcp: {name}: opening your browser for OAuth login — approve there…") + status = await manager.authenticate(name) + if status.error and status.state == "connected": + # e.g. /mcp auth on a server without auth = "oauth" + app.feed.error(f"mcp: {name}: {status.error}") + elif status.state == "connected": + app.feed.info(f"mcp: {name} authenticated ({status.tools} tools)") + elif status.state == "auth_required": + app.feed.error( + f"mcp: {name} authentication failed: {status.error} — try /mcp auth {name}" + ) + else: + app.feed.error(f"mcp: {name} authentication failed: {status.error}") + return + if args and args[0] == "logout": + if len(args) < 2: + app.feed.error("usage: /mcp logout ") + return + status = await manager.logout(args[1]) + if status is None: + app.feed.error(f"unknown MCP server: {args[1]}") + return + if status.error: + app.feed.error(f"mcp: {status.name}: {status.error}") + else: + app.feed.info(f"mcp: {status.name} logged out — /mcp auth {status.name} to reconnect") + return if args and args[0] == "tools": if len(args) < 2: app.feed.error("usage: /mcp tools ") @@ -1017,6 +1050,8 @@ async def cmd_mcp(app: TuiApp, args: list[str]) -> None: lines.append(f"{status.name}: connected ({status.tools} tools)") elif status.state == "disabled": lines.append(f"{status.name}: disabled") + elif status.state == "auth_required": + lines.append(status.auth_hint) else: lines.append(f"{status.name}: failed — {status.error}") app.feed.info("\n".join(lines)) @@ -1088,8 +1123,10 @@ async def cmd_init(app: TuiApp, args: list[str]) -> None: "mcp": ( "MCP servers are configured under [mcp.servers] (stdio or http); Exa web " "search is auto-configured when EXA_API_KEY is set, context7 with " - "enable_context7 = true. /mcp shows state, /mcp tools , /mcp " - "reconnect . Tools appear as mcp::." + "enable_context7 = true. /mcp shows state; /mcp tools|reconnect|auth|logout " + '. HTTP servers with auth = "oauth" log in via /mcp auth ' + "(opens your browser once; credentials are reused afterwards). Tools " + "appear as mcp::." ), "memory": ( "Persistent markdown memory: MEMORY.md (auto-injected), daily logs, " @@ -1387,7 +1424,7 @@ async def cmd_quit(app: TuiApp, args: list[str]) -> None: "wt-exit": "[--delete] [--force]", "loop": " [max-iterations]", "chain": "", - "mcp": "[tools|reconnect ]", + "mcp": "[tools|reconnect|auth|logout ]", "tutor": "", "review": "[file…]", "notifications": "[on|off]", diff --git a/src/lecode/tui/loading.py b/src/lecode/tui/loading.py index aa7106d..60ba02a 100644 --- a/src/lecode/tui/loading.py +++ b/src/lecode/tui/loading.py @@ -127,6 +127,8 @@ def mcp_step(config: Config, mcp_servers: list[ServerStatus] | None = None) -> L lines.append(f"{s.name}: connected · {s.tools} tools") elif s.state == "failed": lines.append(f"{s.name}: failed — {s.error or 'connect error'}") + elif s.state == "auth_required": + lines.append(s.auth_hint) else: lines.append(f"{s.name}: disabled") exa_missing = ( @@ -138,7 +140,7 @@ def mcp_step(config: Config, mcp_servers: list[ServerStatus] | None = None) -> L lines.append("exa: no EXA_API_KEY") if not lines: return LoadStep("mcp", "no servers", SKIP) - degraded = exa_missing or any(s.state == "failed" for s in mcp_servers) + degraded = exa_missing or any(s.state in ("failed", "auth_required") for s in mcp_servers) return LoadStep("mcp", "\n".join(lines), WARN if degraded else OK) # no live statuses (tests, headless): report the configuration only diff --git a/tests/conftest.py b/tests/conftest.py index 532daeb..260c046 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -16,3 +16,25 @@ def tool_ctx(tmp_path, monkeypatch) -> ToolContext: config = Config() checker = PermissionChecker(config, mode="yolo", cwd=tmp_path) return ToolContext(cwd=tmp_path, config=config, permission_checker=checker, auto_approve=True) + + +@pytest.fixture(autouse=True) +def _reset_sse_starlette_shutdown_state(): + """Reset sse-starlette's process-global shutdown state before each test. + + Its shutdown watcher polls the uvicorn server it captured from the + process-global SIGTERM handler table. When one test's uvicorn fixture + sets ``should_exit`` at teardown and the watcher's 0.5s poll lands before + that event loop closes, it flips the module-global ``AppStatus.should_exit`` + (the library never resets it), and every SSE response in LATER tests is + cancelled at birth — "SSE stream ended without a response" / + "ASGI callable returned without completing response". Pinned by the + shutdown regression pair in tests/test_mcp.py; drop both if sse-starlette + ever scopes this state per event loop. + """ + from sse_starlette.sse import AppStatus, _get_shutdown_state + + AppStatus.should_exit = False + state = _get_shutdown_state() + state.watcher_started = False + state.events.clear() diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 1e1649e..a1db769 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -284,7 +284,7 @@ async def fake_connect_one(self, server): @pytest.fixture -async def http_port(): +async def http_server(): """An in-process streamable-HTTP MCP server on a random localhost port.""" import uvicorn from mcp.server.mcpserver import MCPServer @@ -306,12 +306,17 @@ def ping() -> str: task.cancel() raise RuntimeError("uvicorn MCP server did not start") await asyncio.sleep(0.02) - port = uvicorn_server.servers[0].sockets[0].getsockname()[1] - yield port + yield uvicorn_server uvicorn_server.should_exit = True await asyncio.wait_for(task, timeout=10) +@pytest.fixture +async def http_port(http_server): + """Just the port of ``http_server``.""" + yield http_server.servers[0].sockets[0].getsockname()[1] + + async def test_http_transport_round_trip(http_port): config = Config() config.mcp.enable_exa = False @@ -331,6 +336,61 @@ async def test_http_transport_round_trip(http_port): await mgr.shutdown() +async def test_sse_shutdown_sets_process_global_exit_flag(http_server): + """[regression 1/2] uvicorn teardown poisons sse-starlette's global exit flag. + + sse-starlette's shutdown watcher (started with the first SSE response in + the thread) polls the uvicorn server it captured from the process-global + SIGTERM handler table. When this server exits (``should_exit = True``), + the watcher sets ``AppStatus.should_exit`` — a module global the library + never resets — which makes every LATER EventSourceResponse end its SSE + stream immediately (observed with mcp==2.1.1 → sse-starlette 3.4.8 and + uvicorn 0.52.4, which installs signal handlers via ``signal.signal``). + If this assert ever fails, the leak was fixed upstream: drop this pair + and the reset fixture in tests/conftest.py. + """ + from sse_starlette.sse import AppStatus + + port = http_server.servers[0].sockets[0].getsockname()[1] + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["web"] = McpServerConfig( + transport="http", url=f"http://127.0.0.1:{port}/mcp", timeout_s=5.0 + ) + mgr = McpManager(config) + await mgr.connect() + try: + assert mgr.status()[0].state == "connected" # an SSE response ran → watcher is live + finally: + await mgr.shutdown() + http_server.should_exit = True + await asyncio.sleep(1.0) # watcher polls every 0.5s; keep this loop alive for one poll + assert AppStatus.should_exit is True + + +async def test_http_server_connects_after_earlier_sse_server_shutdown(http_port): + """[regression 2/2] the poisoned-flag victim: must stay connected. + + Without the per-test reset in tests/conftest.py, the flag left behind by + the previous test's server shutdown cancels this initialize POST's SSE + response at birth: "SSE stream ended without a response" and the server + side logs "ASGI callable returned without completing response". + """ + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["web"] = McpServerConfig( + transport="http", url=f"http://127.0.0.1:{http_port}/mcp", timeout_s=5.0 + ) + mgr = McpManager(config) + await mgr.connect() + try: + assert mgr.status()[0].state == "connected" + result = await mgr.call("web", "ping", {}) + assert result.content == "pong" + finally: + await mgr.shutdown() + + # -- /mcp command ----------------------------------------------------------------------- @@ -342,6 +402,7 @@ async def app_with_mcp(tmp_path, monkeypatch): app, _, out = make_app(tmp_path, monkeypatch, []) manager = McpManager(mcp_config()) await manager.connect() + manager.bind_registry(app.runtime.registry) app.runtime.ctx.extras[MCP_EXTRA] = manager yield app, manager, out await manager.shutdown() diff --git a/tests/test_mcp_auth.py b/tests/test_mcp_auth.py new file mode 100644 index 0000000..b687f95 --- /dev/null +++ b/tests/test_mcp_auth.py @@ -0,0 +1,658 @@ +"""OAuth 2.1 support tests: storage, loopback callback, and one end-to-end +flow against a local fake authorization server + MCP server. + +The fake AS implements just enough of RFC 8414/7591/6749 for the SDK's +client: resource/server metadata, dynamic registration, an insta-consent +``/authorize`` redirect, and a token endpoint with refresh. +""" + +from __future__ import annotations + +import asyncio +import stat +import time +from collections.abc import Awaitable, Callable +from typing import Any +from urllib.parse import urlencode + +import httpx2 +import pytest +import uvicorn +from mcp.client.auth import OAuthClientProvider, OAuthFlowError +from mcp.server.mcpserver import MCPServer +from mcp.shared.auth import OAuthClientInformationFull, OAuthToken +from starlette.requests import Request +from starlette.responses import JSONResponse, PlainTextResponse, RedirectResponse + +from lecode.agent.builder import build_runtime +from lecode.config.models import Config, McpServerConfig +from lecode.extras import mcp_auth as mcp_auth_mod +from lecode.extras.mcp_auth import ( + FileTokenStorage, + LoopbackAuthCallback, + _redirect_port, + make_oauth_provider, +) +from lecode.extras.mcp_client import MCP_EXTRA, McpManager, ServerStatus, attach_mcp + +# -- storage ------------------------------------------------------------------ + + +@pytest.fixture +def cfg_env(tmp_path, monkeypatch): + """Isolate credential storage from the user's real config dir.""" + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + return tmp_path + + +async def test_storage_round_trip_and_permissions(cfg_env): + storage = FileTokenStorage("https://auth.example/mcp") + assert await storage.get_tokens() is None + assert await storage.get_client_info() is None + + token = OAuthToken(access_token="acc", refresh_token="ref", expires_in=120) + await storage.set_tokens(token) + await storage.set_client_info(OAuthClientInformationFull(client_id="cli-1")) + + assert await storage.get_tokens() == token + assert (await storage.get_client_info()).client_id == "cli-1" + assert stat.S_IMODE(storage.path.stat().st_mode) == 0o600 + assert stat.S_IMODE(storage.path.parent.stat().st_mode) == 0o700 + + # clearing the client registration keeps the tokens + await storage.clear_client_info() + assert await storage.get_tokens() == token + assert await storage.get_client_info() is None + + await storage.clear() + assert not storage.path.exists() + assert await storage.get_tokens() is None + + +async def test_storage_rejects_wrong_endpoint_file(cfg_env): + storage = FileTokenStorage("https://auth.example/mcp") + await storage.set_tokens(OAuthToken(access_token="acc")) + # a file bound to another endpoint must never be served back + storage.path.write_text( + '{"server_url": "https://other.example/mcp", "tokens": {"access_token": "x"}}' + ) + assert await storage.get_tokens() is None + + +async def test_storage_discards_garbage(cfg_env): + storage = FileTokenStorage("https://auth.example/mcp") + storage.path.parent.mkdir(parents=True, exist_ok=True) + storage.path.write_text("not json at all") + assert await storage.get_tokens() is None + + +# -- loopback callback -------------------------------------------------------- + + +@pytest.mark.parametrize( + "request_line", + [b"BROKEN\r\n", b"GET http://[invalid/callback HTTP/1.1\r\n", b"x" * 70000 + b"\r\n"], +) +async def test_loopback_malformed_request_closes_connection(cfg_env, caplog, request_line): + loopauth = await LoopbackAuthCallback.open(FileTokenStorage("https://auth.example/mcp")) + reader, writer = await asyncio.open_connection( + "127.0.0.1", httpx2.URL(loopauth.redirect_url).port + ) + try: + writer.write(request_line) + await writer.drain() + await asyncio.wait_for(reader.read(), timeout=1) + assert reader.at_eof() + assert not [record for record in caplog.records if record.levelno >= 40] + async with httpx2.AsyncClient(trust_env=False) as client: + response = await client.get(loopauth.redirect_url + "?code=valid&state=s1") + assert response.status_code == 200 + assert (await loopauth.wait_for_callback()).code == "valid" + finally: + writer.close() + await writer.wait_closed() + await loopauth.aclose() + + +async def test_loopback_callback_round_trip(cfg_env): + storage = FileTokenStorage("https://auth.example/mcp") + loopauth = await LoopbackAuthCallback.open(storage) + assert loopauth.redirect_url.startswith("http://127.0.0.1:") + assert loopauth.redirect_url.endswith("/callback") + try: + + def hit(): + with httpx2.Client(trust_env=False) as client: + resp = client.get( + loopauth.redirect_url + "?code=c1&state=s1&iss=https%3A%2F%2Fauth.example" + ) + assert resp.status_code == 200 + + task = asyncio.create_task(asyncio.to_thread(hit)) + result = await loopauth.wait_for_callback() + await task + assert result.code == "c1" + assert result.state == "s1" + assert result.iss == "https://auth.example" + finally: + await loopauth.aclose() + + +async def test_loopback_callback_denied(cfg_env): + storage = FileTokenStorage("https://auth.example/mcp") + loopauth = await LoopbackAuthCallback.open(storage) + try: + + def hit(): + with httpx2.Client(trust_env=False) as client: + client.get(loopauth.redirect_url + "?error=access_denied&error_description=no") + + task = asyncio.create_task(asyncio.to_thread(hit)) + with pytest.raises(OAuthFlowError, match="access_denied"): + await loopauth.wait_for_callback() + await task + finally: + await loopauth.aclose() + + +async def test_loopback_close_cuts_pending_connection(cfg_env): + """aclose while a browser connected but stalled: EOF now, wait cancelled.""" + loopauth = await LoopbackAuthCallback.open(FileTokenStorage("https://auth.example/mcp")) + reader, writer = await asyncio.open_connection( + "127.0.0.1", httpx2.URL(loopauth.redirect_url).port + ) + try: + read_task = asyncio.create_task(reader.read()) + await asyncio.sleep(0.1) # let the server accept the connection + await asyncio.wait_for(loopauth.aclose(), timeout=2) + assert await asyncio.wait_for(read_task, timeout=2) == b"" # EOF, not the 60s idle + with pytest.raises(asyncio.CancelledError): + await loopauth.wait_for_callback() + finally: + writer.close() + await writer.wait_closed() + + +def test_loopback_port_is_stable_per_endpoint(cfg_env): + assert _redirect_port("https://auth.example/mcp") == _redirect_port("https://auth.example/mcp") + + +# -- provider wiring ---------------------------------------------------------- + + +async def test_make_oauth_provider_attaches_handlers(cfg_env): + storage = FileTokenStorage("https://auth.example/mcp") + provider = await make_oauth_provider("https://auth.example/mcp", storage, None) + assert isinstance(provider, OAuthClientProvider) + # automatic connections must never open a browser + assert provider.context.redirect_handler is None + assert provider.context.callback_handler is None + redirect = str(provider.context.client_metadata.redirect_uris[0]) + assert redirect == f"http://127.0.0.1:{_redirect_port('https://auth.example/mcp')}/callback" + # no stored tokens → nothing to seed + assert provider.context.token_expiry_time is None + + +async def test_oauth_conflicts_with_static_authorization_header(cfg_env): + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["x"] = McpServerConfig( + transport="http", + url="https://auth.example/mcp", + auth="oauth", + headers={"Authorization": "Bearer static"}, + ) + manager = McpManager(config) + try: + await manager.connect() # fails fast, before any network I/O + status = manager.status()[0] + assert status.state == "failed" + assert "conflicts" in status.error + finally: + await manager.shutdown() + + +# -- slash commands ----------------------------------------------------------- + + +async def test_mcp_auth_and_logout_commands(tmp_path, monkeypatch): + from tests.test_mcp import mock_server_config + from tests.test_tui_app import make_app + + app, _, out = make_app(tmp_path, monkeypatch, []) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["test"] = mock_server_config() + manager = McpManager(config) + await manager.connect() + app.runtime.ctx.extras[MCP_EXTRA] = manager + try: + + async def fake_authenticate(name): + return ServerStatus(name, "connected", tools=4) + + async def fake_logout(name): + return ServerStatus(name, "auth_required") + + monkeypatch.setattr(manager, "authenticate", fake_authenticate) + monkeypatch.setattr(manager, "logout", fake_logout) + + await app.handle_command("/mcp auth test") + assert "test authenticated (4 tools)" in out.getvalue() + await app.handle_command("/mcp logout test") + assert "test logged out" in out.getvalue() + await app.handle_command("/mcp auth nope") + assert "unknown MCP server: nope" in out.getvalue() + finally: + await manager.shutdown() + + +async def test_mcp_status_lists_auth_required(tmp_path, monkeypatch): + from tests.test_tui_app import make_app + + app, _, out = make_app(tmp_path, monkeypatch, []) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url="https://auth.example/mcp", auth="oauth" + ) + manager = McpManager(config) + + async def fail_open(self, server): + raise OAuthFlowError("no redirect handler provided") + + monkeypatch.setattr(McpManager, "_open", fail_open) + await manager.connect() + app.runtime.ctx.extras[MCP_EXTRA] = manager + try: + status = manager.status()[0] + assert status.state == "auth_required" + await app.handle_command("/mcp") + assert "oauth: authentication required — /mcp auth oauth" in out.getvalue() + finally: + await manager.shutdown() + + +# -- end-to-end against a fake authorization server --------------------------- + + +class FakeAuth: + """Issues tokens for one client; guards /mcp with bearer checks. + + Faithful on expiry: a bearer token is only accepted until its own + ``token_ttl`` elapses (set per-issue via :attr:`token_ttl`), so an + expired cached token really gets a 401. + """ + + def __init__(self) -> None: + self.issued_access: list[str] = [] + self.access_expiry: dict[str, float] = {} + self.issued_refresh: set[str] = set() + self.deny = False + self.reject_refresh = False + self.token_ttl = 3600 + self._handlers: dict[tuple[str, str], Callable[[Request], Awaitable[Any]]] = {} + + @staticmethod + def _base(request) -> str: + return f"{request.url.scheme}://{request.url.netloc}" + + async def _protected_resource(self, request): + base = self._base(request) + return JSONResponse({"resource": f"{base}/mcp", "authorization_servers": [base]}) + + async def _auth_metadata(self, request): + base = self._base(request) + return JSONResponse( + { + "issuer": base, + "authorization_endpoint": f"{base}/authorize", + "token_endpoint": f"{base}/token", + "registration_endpoint": f"{base}/register", + "response_types_supported": ["code"], + "code_challenge_methods_supported": ["S256"], + } + ) + + async def _register(self, request): + body = await request.json() + body["client_id"] = "lecode-test-client" + body["client_secret"] = "lecode-test-secret" + return JSONResponse(body, status_code=201) + + async def _authorize(self, request): + query = request.query_params + redirect = query["redirect_uri"] + if self.deny: + params = { + "state": query.get("state", ""), + "error": "access_denied", + # the description is attacker/server-controlled text; the + # regression pair below pins that control chars never survive + "error_description": "user\x1b said\x07 no", + } + else: + params = { + "state": query.get("state", ""), + "code": "lecode-auth-code", + "iss": self._base(request), + } + separator = "&" if "?" in redirect else "?" + return RedirectResponse(f"{redirect}{separator}{urlencode(params)}", status_code=302) + + async def _token(self, request): + form = await request.form() + if form.get("grant_type") == "refresh_token": + assert form.get("refresh_token") in self.issued_refresh + if self.reject_refresh: + return JSONResponse({"error": "invalid_grant"}, status_code=400) + access = f"access-{len(self.issued_access)}" + refresh = f"refresh-{len(self.issued_access)}" + self.issued_access.append(access) + self.access_expiry[access] = time.time() + self.token_ttl + self.issued_refresh.add(refresh) + body = { + "access_token": access, + "token_type": "Bearer", + "expires_in": self.token_ttl, + "refresh_token": refresh, + } + if form.get("scope"): + body["scope"] = form["scope"] + return JSONResponse(body) + + def build_app(self): + """The fake AS + MCP server as one ASGI app. + + A plain ASGI wrapper (not a Starlette Mount) keeps the MCP app the + served app, so uvicorn runs its lifespan — the SDK's session manager + starts its task group there and refuses requests without it. + """ + mcp = MCPServer("mock-oauth") + + @mcp.tool() + def ping() -> str: + """Ping the server.""" + return "pong" + + inner = mcp.streamable_http_app() + self._handlers = { + ("GET", "/.well-known/oauth-protected-resource"): self._protected_resource, + ("GET", "/.well-known/oauth-authorization-server"): self._auth_metadata, + ("POST", "/register"): self._register, + ("GET", "/authorize"): self._authorize, + ("POST", "/token"): self._token, + } + authz = self + + async def app(scope, receive, send): + if scope["type"] != "http": + await inner(scope, receive, send) + return + request = Request(scope, receive) + path = request.url.path + if path.startswith("/mcp"): + authorization = request.headers.get("authorization", "") + now = time.time() + accepted = any( + authorization == f"Bearer {t}" and now < authz.access_expiry[t] + for t in authz.issued_access + ) + if not accepted: + base = f"{request.url.scheme}://{request.url.netloc}" + response = PlainTextResponse( + "authentication required", + status_code=401, + headers={ + "WWW-Authenticate": ( + f'Bearer resource_metadata="{base}' + '/.well-known/oauth-protected-resource"' + ) + }, + ) + await response(scope, receive, send) + return + handler = authz._handlers.get((scope.get("method"), path)) + if handler is not None: + response = await handler(request) + await response(scope, receive, send) + return + await inner(scope, receive, send) + + return app + + +@pytest.fixture +async def fake_oauth_server(): + """The fake AS + MCP server on a random localhost port.""" + authz = FakeAuth() + config = uvicorn.Config(authz.build_app(), host="127.0.0.1", port=0, log_level="error") + server = uvicorn.Server(config) + task = asyncio.ensure_future(server.serve()) + deadline = asyncio.get_running_loop().time() + 15 + while not server.started: + if asyncio.get_running_loop().time() > deadline: + task.cancel() + raise RuntimeError("fake OAuth server did not start") + await asyncio.sleep(0.02) + port = server.servers[0].sockets[0].getsockname()[1] + yield f"http://127.0.0.1:{port}", authz + server.should_exit = True + await asyncio.wait_for(task, timeout=10) + + +def _fake_browser(url: str) -> bool: + """A stand-in for webbrowser.open: follow the AS redirect into lecode.""" + with httpx2.Client(trust_env=False, follow_redirects=False) as client: + resp = client.get(url) + assert resp.status_code in (302, 303), f"expected a redirect, got {resp.status_code}" + location = resp.headers["location"] + assert location.startswith("http://127.0.0.1:"), "redirect must land on loopback" + landed = client.get(location) + assert landed.status_code == 200 + return True + + +async def test_oauth_end_to_end(tmp_path, monkeypatch, fake_oauth_server): + base, authz = fake_oauth_server + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + + browser_calls: list[str] = [] + monkeypatch.setattr( + mcp_auth_mod.webbrowser, + "open", + lambda url: (browser_calls.append(url), _fake_browser(url))[1], + ) + + runtime = build_runtime(config, tmp_path, auto_approve=True) + manager = await attach_mcp(runtime.registry, runtime.ctx) + try: + # automatic connect never opens a browser: it reports auth_required + status = manager.status()[0] + assert status.state == "auth_required" + assert browser_calls == [] + + # interactive login: browser once, then connected with tools; the AS + # mints a short-lived access token so the live-session refresh path is + # exercised next (the SDK only refreshes tokens it minted itself) + authz.token_ttl = 1 + status = await manager.authenticate("oauth") + assert status.state == "connected" + assert status.tools == 1 + assert len(browser_calls) == 1 + assert "mcp:oauth:ping" in runtime.registry.names() + + # a real tool call goes over the authenticated transport + message = await runtime.registry.dispatch("call-1", "mcp:oauth:ping", "{}", runtime.ctx) + assert message["content"] == "pong" + + # once the short-lived token expires, the next call refreshes it + # silently (no browser, no user interaction) + authz.token_ttl = 3600 + await asyncio.sleep(1.2) + issued_before = len(authz.issued_access) + message = await runtime.registry.dispatch("call-2", "mcp:oauth:ping", "{}", runtime.ctx) + assert message["content"] == "pong" + assert len(authz.issued_access) > issued_before + assert len(browser_calls) == 1 + + # "restart": a fresh manager connects with cached credentials only + fresh = McpManager(config) + await fresh.connect() + assert fresh.status()[0].state == "connected" + assert len(browser_calls) == 1 + await fresh.shutdown() + + # logout clears persisted credentials and deregisters the tools + storage = FileTokenStorage(f"{base}/mcp") + status = await manager.logout("oauth") + assert status.state == "auth_required" + assert await storage.get_tokens() is None + assert "mcp:oauth:ping" not in runtime.registry.names() + + # after logout, reconnect and call recovery stay non-interactive: + # they must never open a browser, only report authentication required + status = await manager.reconnect("oauth") + assert status.state == "auth_required" + assert len(browser_calls) == 1 + result = await manager.call("oauth", "ping", {}) + assert result.is_error + assert "unavailable" in result.content + assert len(browser_calls) == 1 + finally: + await manager.shutdown() + + +async def test_oauth_restart_with_expired_token_refreshes_silently( + tmp_path, monkeypatch, fake_oauth_server +): + """Cached credentials past their expiry must refresh on restart, not 401. + + The SDK restores stored tokens but not their absolute expiry (only the + ones it minted in-session carry that), so without the expiry seeding in + ``make_oauth_provider`` an expired access token gets attached, the server + 401s it, and the automatic flow dead-ends into auth_required even though + a perfectly good refresh token sits on disk. + """ + base, authz = fake_oauth_server + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + browser_calls: list[str] = [] + monkeypatch.setattr( + mcp_auth_mod.webbrowser, + "open", + lambda url: (browser_calls.append(url), _fake_browser(url))[1], + ) + + manager = McpManager(config) + await manager.connect() + assert manager.status()[0].state == "auth_required" + authz.token_ttl = 1 # the access token dies one second after login + assert (await manager.authenticate("oauth")).state == "connected" + await manager.shutdown() + + await asyncio.sleep(1.2) # lecode is "down" while the access token expires + authz.token_ttl = 3600 + issued_before = len(authz.issued_access) + + # restart at the attach_mcp seam: connect + tool bridge through the registry + runtime = build_runtime(config, tmp_path, auto_approve=True) + fresh = await attach_mcp(runtime.registry, runtime.ctx) + try: + assert fresh.status()[0].state == "connected" # silent refresh, not auth_required + assert len(authz.issued_access) > issued_before # a new token was minted + assert len(browser_calls) == 1 # restart never opens a browser + message = await runtime.registry.dispatch("call-1", "mcp:oauth:ping", "{}", runtime.ctx) + assert message["content"] == "pong" + finally: + await fresh.shutdown() + + +async def test_oauth_restart_with_rejected_refresh_stays_noninteractive( + tmp_path, monkeypatch, fake_oauth_server +): + """A dead refresh token on restart: auth_required, never a browser.""" + base, authz = fake_oauth_server + authz.reject_refresh = True + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + browser_calls: list[str] = [] + monkeypatch.setattr( + mcp_auth_mod.webbrowser, + "open", + lambda url: (browser_calls.append(url), _fake_browser(url))[1], + ) + + manager = McpManager(config) + await manager.connect() + authz.token_ttl = 1 + assert (await manager.authenticate("oauth")).state == "connected" + await manager.shutdown() + + await asyncio.sleep(1.2) + + fresh = McpManager(config) + await fresh.connect() + try: + assert fresh.status()[0].state == "auth_required" + assert len(browser_calls) == 1 # a failed refresh must not open a browser + finally: + await fresh.shutdown() + + +async def test_oauth_denied_consent(tmp_path, monkeypatch, fake_oauth_server): + base, authz = fake_oauth_server + authz.deny = True + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + monkeypatch.setattr(mcp_auth_mod.webbrowser, "open", _fake_browser) + + manager = McpManager(config) + try: + await manager.connect() + status = await manager.authenticate("oauth") + assert status.state == "auth_required" + assert "access_denied" in status.error + # the AS-controlled description must not carry terminal control chars + assert "\x1b" not in status.error + assert "\x07" not in status.error + assert "user said no" in status.error + finally: + await manager.shutdown() + + +async def test_oauth_no_browser_available(tmp_path, monkeypatch, fake_oauth_server): + base, _ = fake_oauth_server + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + monkeypatch.setattr(mcp_auth_mod.webbrowser, "open", lambda url: False) + + manager = McpManager(config) + try: + await manager.connect() + status = await manager.authenticate("oauth") + assert status.state == "auth_required" + assert "browser" in status.error + finally: + await manager.shutdown() From 8bd505937cb67e9177a453b11f02d8d83256b0e9 Mon Sep 17 00:00:00 2001 From: Mouhand-Kaddo Date: Tue, 8 Sep 2026 10:20:47 +0400 Subject: [PATCH 2/2] fix: show the OAuth authorization URL and quiet the startup flow traceback MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit /mcp auth now prints the authorization URL in the feed right before the browser opens, so it can be pasted into a different browser; if no browser can be opened at all, the error carries the URL to open by hand. The 'No redirect handler provided' OAuthFlowError the user saw is the expected startup dead end (no credentials, no browser at startup): the SDK logs it at ERROR with a full traceback, splashing raw stderr around the TUI at every launch. Automatic connects now run with the mcp.client.auth loggers muted; the failure is still classified and shown as 'authentication required — /mcp auth '. --- docs/mcp.md | 7 ++-- src/lecode/extras/mcp_auth.py | 21 +++++++++-- src/lecode/extras/mcp_client.py | 41 ++++++++++++++++++-- src/lecode/slash/handlers.py | 7 +++- tests/test_mcp_auth.py | 67 ++++++++++++++++++++++++++++++++- 5 files changed, 130 insertions(+), 13 deletions(-) diff --git a/docs/mcp.md b/docs/mcp.md index ad87276..73aed86 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -51,9 +51,10 @@ auth = "oauth" access token is refreshed silently from its refresh token; a missing or rejected refresh shows `authentication required`. Startup never blocks on a browser. -- `/mcp auth glitchtip` runs the interactive login: it opens your browser, - the server walks you through approval, and the redirect lands back on a - loopback port lecode serves (`http://127.0.0.1:/callback`). +- `/mcp auth glitchtip` runs the interactive login: it opens your browser + and prints the authorization URL in the feed (paste it into a different + browser if you prefer); the redirect lands back on a loopback port lecode + serves (`http://127.0.0.1:/callback`). - `/mcp logout glitchtip` drops the session and the persisted credentials. - Credentials (access + refresh tokens, client registration) are stored per endpoint under `~/.config/lecode/mcp-auth/` (0600 files, 0700 directory, diff --git a/src/lecode/extras/mcp_auth.py b/src/lecode/extras/mcp_auth.py index 42b60cb..f4f4f90 100644 --- a/src/lecode/extras/mcp_auth.py +++ b/src/lecode/extras/mcp_auth.py @@ -25,6 +25,7 @@ import tempfile import time import webbrowser +from collections.abc import Callable from pathlib import Path from typing import TYPE_CHECKING, Any from urllib.parse import parse_qs, urlsplit @@ -225,16 +226,25 @@ def __init__( redirect_url: str, result: asyncio.Future[dict[str, list[str]]], accepted: set[asyncio.StreamWriter], + announce: Callable[[str], Any] | None = None, ) -> None: self.server_url = server_url self._server = server self.redirect_url = redirect_url self._result = result self._accepted = accepted + self._announce = announce @classmethod - async def open(cls, storage: FileTokenStorage) -> LoopbackAuthCallback: - """Bind the callback listener, reusing the port registered previously.""" + async def open( + cls, storage: FileTokenStorage, announce: Callable[[str], Any] | None = None + ) -> LoopbackAuthCallback: + """Bind the callback listener, reusing the port registered previously. + + ``announce`` (optional) is called with the authorization URL right + before the browser opens — the UI uses it to show the link so the + user can paste it into a different browser. + """ port = _redirect_port(storage.server_url) client_info = await storage.get_client_info() registered = client_info.redirect_uris if client_info is not None else None @@ -296,16 +306,19 @@ async def handle(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> server = await asyncio.start_server(handle, "127.0.0.1", 0) actual_port = server.sockets[0].getsockname()[1] redirect_url = f"http://127.0.0.1:{actual_port}{_REDIRECT_CALLBACK}" - return cls(storage.server_url, server, redirect_url, result, accepted) + return cls(storage.server_url, server, redirect_url, result, accepted, announce) async def open_browser(self, authorization_url: str) -> None: """The SDK's redirect_handler: hand the URL to the system browser.""" from mcp.client.auth import OAuthFlowError + if self._announce is not None: + with contextlib.suppress(Exception): # announcing is best-effort + self._announce(authorization_url) opened = await asyncio.to_thread(webbrowser.open, authorization_url) if not opened: raise OAuthFlowError( - "could not open a browser on this machine — run /mcp auth where one exists" + f"could not open a browser — open this URL to continue: {authorization_url}" ) async def wait_for_callback(self) -> AuthorizationCodeResult: diff --git a/src/lecode/extras/mcp_client.py b/src/lecode/extras/mcp_client.py index e01b3af..90188f9 100644 --- a/src/lecode/extras/mcp_client.py +++ b/src/lecode/extras/mcp_client.py @@ -34,6 +34,7 @@ import contextlib import logging import os +from collections.abc import Callable, Iterator from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Literal @@ -114,6 +115,8 @@ class _Server: error: str | None = None auth_required: bool = False interactive: bool = False + #: called with the authorization URL during an interactive login (UI hint) + announce: Callable[[str], Any] | None = None @property def status(self) -> ServerStatus: @@ -166,6 +169,27 @@ def _connect_failure(e: BaseException) -> tuple[str, bool]: return f"{type(e).__name__}: {_clean_error(e)}", False +@contextlib.contextmanager +def _quiet_oauth_flow_logs() -> Iterator[None]: + """Mute the SDK's expected dead-end OAuth logging during automatic connects. + + Without stored credentials the SDK still attempts the authorization-code + grant, hits "No redirect handler provided", and logs it at ERROR with a + full traceback — splashing raw stderr around the TUI at every startup + until the user runs ``/mcp auth``. The failure is expected and already + classified as auth_required; only the noise is lost. + """ + loggers = [logging.getLogger(n) for n in ("mcp.client.auth", "mcp.client.auth.oauth2")] + saved = [(lg, lg.level) for lg in loggers] + for lg in loggers: + lg.setLevel(logging.CRITICAL) + try: + yield + finally: + for lg, level in saved: + lg.setLevel(level) + + def _result_text(result: Any) -> str: """Flatten a CallToolResult's content parts to plain text.""" parts = [getattr(part, "text", "") for part in result.content or []] @@ -195,8 +219,10 @@ async def connect(self) -> None: async def _connect_one(self, server: _Server) -> None: if not server.config.enabled: return + quiet = server.config.auth == "oauth" and not server.interactive try: - await asyncio.wait_for(self._open(server), timeout=CONNECT_TIMEOUT_S) + with _quiet_oauth_flow_logs() if quiet else contextlib.nullcontext(): + await asyncio.wait_for(self._open(server), timeout=CONNECT_TIMEOUT_S) server.error = None server.auth_required = False log.debug("mcp: %s connected (%d tools)", server.name, len(server.tools)) @@ -314,7 +340,7 @@ async def _build_http_client( storage = FileTokenStorage(url) loopback = None if server.interactive: - loopback = await LoopbackAuthCallback.open(storage) + loopback = await LoopbackAuthCallback.open(storage, announce=server.announce) stack.push_async_callback(loopback.aclose) return create_mcp_http_client( headers=headers, auth=await make_oauth_provider(url, storage, loopback) @@ -422,11 +448,16 @@ async def reconnect(self, name: str) -> ServerStatus | None: self._sync_tools(name) return server.status - async def authenticate(self, name: str) -> ServerStatus | None: + async def authenticate( + self, name: str, announce: Callable[[str], Any] | None = None + ) -> ServerStatus | None: """Interactive OAuth login for one server (opens the browser). Runs outside the normal per-server connect budget: the user needs time - to approve in the browser. Non-OAuth servers get a plain error state. + to approve in the browser. ``announce`` is called with the + authorization URL just before the browser opens, so the UI can show + the link (a different browser can be used with it). Non-OAuth servers + get a plain error state. """ server = self._servers.get(name) if server is None: @@ -438,6 +469,7 @@ async def authenticate(self, name: str) -> ServerStatus | None: server.error = None server.auth_required = False server.interactive = True + server.announce = announce try: await asyncio.wait_for(self._open(server), timeout=INTERACTIVE_AUTH_BUDGET_S) except asyncio.CancelledError: @@ -450,6 +482,7 @@ async def authenticate(self, name: str) -> ServerStatus | None: log.debug("mcp: %s authentication failed: %s", name, e) finally: server.interactive = False + server.announce = None self._sync_tools(name) return server.status diff --git a/src/lecode/slash/handlers.py b/src/lecode/slash/handlers.py index 2e1e4e5..bdde530 100644 --- a/src/lecode/slash/handlers.py +++ b/src/lecode/slash/handlers.py @@ -1004,7 +1004,12 @@ async def cmd_mcp(app: TuiApp, args: list[str]) -> None: app.feed.error(f"unknown MCP server: {name}") return app.feed.info(f"mcp: {name}: opening your browser for OAuth login — approve there…") - status = await manager.authenticate(name) + status = await manager.authenticate( + name, + announce=lambda url: app.feed.info( + f"mcp: {name}: authorization URL (a different browser works too): {url}" + ), + ) if status.error and status.state == "connected": # e.g. /mcp auth on a server without auth = "oauth" app.feed.error(f"mcp: {name}: {status.error}") diff --git a/tests/test_mcp_auth.py b/tests/test_mcp_auth.py index b687f95..8cd7419 100644 --- a/tests/test_mcp_auth.py +++ b/tests/test_mcp_auth.py @@ -9,6 +9,7 @@ from __future__ import annotations import asyncio +import logging import stat import time from collections.abc import Awaitable, Callable @@ -228,7 +229,9 @@ async def test_mcp_auth_and_logout_commands(tmp_path, monkeypatch): app.runtime.ctx.extras[MCP_EXTRA] = manager try: - async def fake_authenticate(name): + async def fake_authenticate(name, announce=None): + if announce is not None: + announce("https://auth.example/authorize?code_challenge=x") return ServerStatus(name, "connected", tools=4) async def fake_logout(name): @@ -239,6 +242,8 @@ async def fake_logout(name): await app.handle_command("/mcp auth test") assert "test authenticated (4 tools)" in out.getvalue() + assert "authorization URL" in out.getvalue() + assert "https://auth.example/authorize?code_challenge=x" in out.getvalue() await app.handle_command("/mcp logout test") assert "test logged out" in out.getvalue() await app.handle_command("/mcp auth nope") @@ -654,5 +659,65 @@ async def test_oauth_no_browser_available(tmp_path, monkeypatch, fake_oauth_serv status = await manager.authenticate("oauth") assert status.state == "auth_required" assert "browser" in status.error + # when no browser can be opened, the error carries the URL to open by hand + assert f"{base}/authorize?" in status.error + finally: + await manager.shutdown() + + +async def test_oauth_auth_announces_authorization_url(tmp_path, monkeypatch, fake_oauth_server): + """/mcp auth must surface the authorization URL so another browser can be used.""" + base, _ = fake_oauth_server + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + monkeypatch.setattr(mcp_auth_mod.webbrowser, "open", _fake_browser) + + manager = McpManager(config) + try: + await manager.connect() + announced: list[str] = [] + status = await manager.authenticate("oauth", announce=announced.append) + assert status.state == "connected" + assert len(announced) == 1 # once, before the browser opened + assert announced[0].startswith(f"{base}/authorize?") + assert "127.0.0.1" in announced[0] # the loopback redirect is in the URL + finally: + await manager.shutdown() + + +async def test_automatic_connect_does_not_splash_oauth_flow_traceback( + tmp_path, monkeypatch, fake_oauth_server, caplog +): + """The expected no-credentials dead end at startup must not hit the terminal. + + Without stored tokens the SDK still attempts the authorization-code grant, + hits "No redirect handler provided", and logs it at ERROR with a full + traceback — raw stderr splashing around the TUI at every startup until the + user runs /mcp auth. The manager already classifies it as auth_required. + """ + base, _ = fake_oauth_server + monkeypatch.setenv("LECODE_CONFIG_DIR", str(tmp_path / "lecode-config")) + config = Config() + config.mcp.enable_exa = False + config.mcp.servers["oauth"] = McpServerConfig( + transport="http", url=f"{base}/mcp", auth="oauth", timeout_s=5.0 + ) + monkeypatch.setattr(mcp_auth_mod.webbrowser, "open", _fake_browser) + + manager = McpManager(config) + try: + with caplog.at_level(logging.DEBUG, logger="mcp.client.auth"): + await manager.connect() + assert manager.status()[0].state == "auth_required" + loud = [ + r + for r in caplog.records + if r.name.startswith("mcp.client.auth") and r.levelno >= logging.ERROR + ] + assert loud == [] finally: await manager.shutdown()