diff --git a/docs/mcp.md b/docs/mcp.md index 4daa3a6..73aed86 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -36,8 +36,35 @@ 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 + 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, + 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 +95,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..f4f4f90 --- /dev/null +++ b/src/lecode/extras/mcp_auth.py @@ -0,0 +1,395 @@ +"""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 collections.abc import Callable +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], + 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, 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 + 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, 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( + f"could not open a browser — open this URL to continue: {authorization_url}" + ) + + 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..90188f9 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. @@ -30,11 +34,13 @@ import contextlib import logging import os +from collections.abc import Callable, Iterator 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 +57,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 +95,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 +113,83 @@ 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 + #: called with the authorization URL during an interactive login (UI hint) + announce: Callable[[str], Any] | None = None @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 + + +@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 []] @@ -121,6 +203,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 ------------------------------------------------------------ @@ -136,46 +219,81 @@ 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)) 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 +312,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, announce=server.announce) + 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 +358,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 +399,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 +428,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 +445,62 @@ 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, 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. ``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: + 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 + server.announce = announce + 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 + server.announce = None + 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 +534,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..bdde530 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,51 @@ 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, + 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}") + 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 +1055,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 +1128,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 +1429,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..8cd7419 --- /dev/null +++ b/tests/test_mcp_auth.py @@ -0,0 +1,723 @@ +"""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 logging +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, 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): + 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() + 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") + 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 + # 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()