diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index a4c27e772..fd73981e9 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -1975,6 +1975,13 @@ def delete_coding_agent_config(workspace: str, token: str, name: str) -> str | N # --- MCP services (parallel to model services) ----------------------------- +# Canonical path segment of an AI Gateway MCP-services endpoint +# (``https:///ai-gateway/mcp-services/..``). This is +# the single source of truth: URL building (below), connection-backed detection +# (`mcp_connection_login.connection_from_url`), and URL-shape classification +# (`mcp.py`) all reference this one constant. +AIGW_MCP_SERVICES_SEGMENT = "/ai-gateway/mcp-services/" + _MCP_SERVICE_NAME_PREFIX = "mcp-services/" @@ -2042,7 +2049,7 @@ def list_mcp_services( def build_mcp_service_url(workspace: str, full_name: str) -> str: - return f"{workspace}/ai-gateway/mcp-services/{full_name}" + return f"{workspace}{AIGW_MCP_SERVICES_SEGMENT}{full_name}" def build_skills_mcp_url(workspace: str, locations: list[str]) -> str: diff --git a/src/ucode/mcp.py b/src/ucode/mcp.py index ec2eed073..88046969c 100644 --- a/src/ucode/mcp.py +++ b/src/ucode/mcp.py @@ -19,6 +19,7 @@ from ucode.config_io import restore_file from ucode.constants import MCP_CLEANUP_SCOPES, MCP_USER_SCOPE from ucode.databricks import ( + AIGW_MCP_SERVICES_SEGMENT, PermissionDeniedError, apply_pat_environment, build_mcp_proxy_argv, @@ -31,6 +32,7 @@ list_mcp_services, workspace_hostname, ) +from ucode.mcp_connection_login import connection_from_url from ucode.mcp_oauth import ( CLAUDE_CODE_OAUTH_CLIENT_ID, CURSOR_OAUTH_CLIENT_ID, @@ -57,10 +59,6 @@ spinner, ) -# AI Gateway MCP-services endpoints carry this path segment. These are the -# connection-backed services that need a per-user connection login. -AIGW_MCP_SERVICES_PATH = "/ai-gateway/mcp-services/" - # Workspace-relative path fragments for the V2 AI Gateway MCP endpoints, shared by the URL-shape # checks (`_is_app_mcp_server`, `_mcp_server_location`) so the set stays in one place. MCP_EXTERNAL_PATH = "/api/2.0/mcp/external/" @@ -259,7 +257,7 @@ def _oauth_http_client(client: str, workspace: str, *, use_pat: bool) -> str | N This is the single source of truth for that choice, shared by the per-server (:func:`configure_client_mcp_server`) and batched (:func:`_managed_mcp_entry`) paths. Whether it applies to a *specific* server additionally requires a connection-backed mcp-services URL - (``AIGW_MCP_SERVICES_PATH``), which the caller checks per server. It is URL-independent, so + (``AIGW_MCP_SERVICES_SEGMENT``), which the caller checks per server. It is URL-independent, so callers can compute it once per (client, workspace) rather than once per server.""" oauth_client = AGENT_OAUTH_CLIENT.get(client) if oauth_client is not None and not use_pat and oauth_client_available(workspace, oauth_client): @@ -282,7 +280,7 @@ def configure_client_mcp_server( # proxy: non-connection MCPs, the skills registry, PAT auth, agents without a mapped OAuth # client, and workspaces where the mapped client isn't published. http_client = _oauth_http_client(client, workspace, use_pat=use_pat) - if http_client is not None and AIGW_MCP_SERVICES_PATH in url: + if http_client is not None and AIGW_MCP_SERVICES_SEGMENT in url: if client == "claude": removed_scopes = [ scope @@ -753,7 +751,7 @@ def _is_app_mcp_server(server: dict) -> bool: return False stripped = url.rstrip("/") known = ( - AIGW_MCP_SERVICES_PATH, + AIGW_MCP_SERVICES_SEGMENT, MCP_EXTERNAL_PATH, MCP_GENIE_PATH, MCP_VECTOR_SEARCH_PATH, @@ -867,7 +865,7 @@ def _agent_managed_file_entries( if agent == "claude": # A native HTTP+OAuth entry is only valid for a connection-backed mcp-services URL; # anything else stays on the stdio proxy so both delivery paths resolve identically. - if AIGW_MCP_SERVICES_PATH not in url: + if AIGW_MCP_SERVICES_SEGMENT not in url: continue entries[name] = claude.managed_mcp_entry(url) elif agent == "codex": @@ -1199,7 +1197,7 @@ def _managed_mcp_entry( workspace) (or ``None``); the caller computes it once per client so the batch loop doesn't re-probe ``oauth_client_available`` per server. HTTP+OAuth applies here only when that client is set AND this server's URL is a connection-backed mcp-services URL; otherwise the stdio proxy.""" - use_http = http_client is not None and AIGW_MCP_SERVICES_PATH in url + use_http = http_client is not None and AIGW_MCP_SERVICES_SEGMENT in url if client == "claude": if use_http: return claude.managed_mcp_entry(url) @@ -2147,8 +2145,8 @@ def _mcp_server_location(server: dict) -> str: return "skills" url = str(server.get("url") or "") stripped = url.rstrip("/") - if AIGW_MCP_SERVICES_PATH in url: - return url.split(AIGW_MCP_SERVICES_PATH, 1)[1] or "mcp-service" + if AIGW_MCP_SERVICES_SEGMENT in url: + return connection_from_url(url) or "mcp-service" if MCP_EXTERNAL_PATH in url: return f"connection:{stripped.rsplit('/', 1)[-1]}" if MCP_GENIE_PATH in url: diff --git a/src/ucode/mcp_connection_login.py b/src/ucode/mcp_connection_login.py new file mode 100644 index 000000000..d1ae3c2a7 --- /dev/null +++ b/src/ucode/mcp_connection_login.py @@ -0,0 +1,140 @@ +"""Per-connection login for AI Gateway MCP services, driven by the proxy on a 401. + +A connection-backed AI Gateway MCP service (e.g. ``system.ai.github``) needs a +per-user connection credential before its tools can be used. Until the user has +logged in to the underlying SaaS, AI Gateway answers requests with an HTTP 401 +(RFC 9728 ``WWW-Authenticate``). + +The ``ug mcp-proxy`` bridge (see ``mcp_proxy``) sees that 401 in its httpx auth +flow and, for a connection-backed service, runs the Databricks CLI U2M login +with an RFC 8707 ``resource`` indicator naming the service, then retries the +request. A resource-aware ``/oidc`` drives the connection's own SaaS login +before minting the token, so the credential exists on retry — transparently to +the coding agent, which just sees the connection authenticate and succeed. This +is the behaviour of a generic OAuth MCP bridge (e.g. ``mcp-remote``), done in +ucode with the Databricks CLI so no extra library or per-agent OAuth app is +needed. Requires the CLI ``--resource`` flag (databricks/cli#6621). +""" + +from __future__ import annotations + +import subprocess +import sys + +from ucode.databricks import AIGW_MCP_SERVICES_SEGMENT + +# Login can pop a browser and wait for the user to complete the SaaS login, so +# allow generously more than a token refresh would take. +_LOGIN_TIMEOUT_SECONDS = 300 + + +def connection_from_url(url: str) -> str | None: + """Return the connection FQN of an AI Gateway MCP service URL, or ``None``. + + ``https://ws/ai-gateway/mcp-services/system.ai.github`` -> ``system.ai.github``. + A URL that is not an mcp-services endpoint (or names no service) returns + ``None`` — only connection-backed services get the login-on-401 treatment. + """ + marker = url.find(AIGW_MCP_SERVICES_SEGMENT) + if marker == -1: + return None + tail = url[marker + len(AIGW_MCP_SERVICES_SEGMENT) :] + # Strip any trailing path (``/tools/list``), query, or fragment. + connection = tail.split("/")[0].split("?")[0].split("#")[0] + return connection or None + + +def _cli_supports_resource_flag(login_binary: str) -> bool: + """Whether `` auth login`` advertises the ``--resource`` flag. + + The connection sign-in needs a Databricks CLI with ``--resource`` + (databricks/cli#6621). An older CLI rejects the flag and the login exits with + a cryptic parse error, so we check ``--help`` up front to give a clear message + instead. Fail-open (assume supported) if ``--help`` can't be run — the real + login attempt will surface any genuine failure.""" + try: + result = subprocess.run( + [login_binary, "auth", "login", "--help"], + check=False, + timeout=20, + capture_output=True, + text=True, + ) + except (OSError, subprocess.TimeoutExpired): + return True + return "--resource" in f"{result.stdout or ''}{result.stderr or ''}" + + +def run_connection_login( + resource_url: str, + workspace: str, + *, + profile: str | None = None, + login_binary: str = "databricks", +) -> tuple[bool, str]: + """Run the CLI U2M login with an RFC 8707 resource indicator for this service. + + ``resource_url`` is the MCP service endpoint (also the proxy's upstream URL); + it is sent as ``--resource`` so a resource-aware ``/oidc`` drives the + connection's SaaS login before issuing the token. Uses the Databricks CLI's + own default client, whose loopback redirect is already registered — no + ``--client-id`` needed. + + The CLI opens the browser to complete the login and prints the authorize URL. + We route its output to **stderr** (never stdout — that is the proxy's MCP + JSON-RPC wire), so a coding agent surfaces it in the server's log and the URL + stays visible when the browser can't open (e.g. a headless remote). Returns + ``(ok, message)``; on failure ``message`` points at that log. + """ + connection = connection_from_url(resource_url) or resource_url + if not _cli_supports_resource_flag(login_binary): + return False, ( + f"the Databricks CLI ('{login_binary}') has no `--resource` flag, so the " + f"'{connection}' connection sign-in can't run. Upgrade the CLI " + "(databricks/cli#6621) and retry." + ) + argv = [ + login_binary, + "auth", + "login", + "--host", + workspace.rstrip("/"), + "--resource", + resource_url, + ] + if profile: + argv += ["--profile", profile] + print( + f"Signing in to '{connection}' — opening your browser to complete the connection login; " + "if it doesn't open, use the authorization URL printed below.", + file=sys.stderr, + flush=True, + ) + try: + # stdout -> stderr: the CLI's prompts and authorize URL reach the agent's + # MCP log (fd 2) without corrupting this process's stdout (fd 1, the MCP + # JSON-RPC stream). stdin is closed since the flow is browser-driven. + result = subprocess.run( + argv, + check=False, + timeout=_LOGIN_TIMEOUT_SECONDS, + stdin=subprocess.DEVNULL, + stdout=sys.stderr, + stderr=sys.stderr, + ) + except OSError as exc: + return False, f"could not run '{login_binary} auth login': {exc}" + except subprocess.TimeoutExpired: + return False, "connection sign-in timed out waiting for the browser flow to complete" + if result.returncode == 0: + return True, "signed in" + return ( + False, + f"connection sign-in did not complete (CLI exited {result.returncode}; see the log above)", + ) + + +__all__ = [ + "connection_from_url", + "run_connection_login", +] diff --git a/src/ucode/mcp_proxy.py b/src/ucode/mcp_proxy.py index c01741698..95ad38d8f 100644 --- a/src/ucode/mcp_proxy.py +++ b/src/ucode/mcp_proxy.py @@ -39,10 +39,12 @@ from typing import Protocol, Self import anyio +from anyio import to_thread from mcp.client.streamable_http import streamable_http_client from mcp.server.stdio import stdio_server from ucode.databricks import ensure_pat_bearer, get_databricks_token +from ucode.mcp_connection_login import connection_from_url, run_connection_login # Exit code used when the proxy cannot continue. MCP clients surface a non-zero # exit far more usefully than a timeout, so bail out instead of hanging. @@ -111,29 +113,72 @@ def _fail_fast(message: str) -> None: raise SystemExit(AUTH_FAILURE_EXIT_CODE) -def _build_token_auth(workspace: str, profile: str | None): - """Build an httpx ``Auth`` that injects a fresh bearer on every request. - - The base class comes from whichever httpx the SDK uses (see ``_httpx``), so - the returned auth is accepted by that SDK's ``AsyncClient``. Behaviour is - identical across flavours — ``Auth.auth_flow`` has the same generator - contract in httpx and httpx2.""" +def _build_token_auth(url: str, workspace: str, profile: str | None, *, use_pat: bool = False): + """Build an httpx ``Auth`` that injects a fresh bearer and drives the + per-connection login **lazily, on a 401** — not eagerly at startup. + + The bearer is the Databricks *workspace* token (read from the CLI session, + refreshed as it nears expiry). For a connection-backed AI Gateway service, + the gateway answers with an RFC 9728 ``401`` until the user holds the per-user + connection credential. Rather than block startup logging in every configured + server (which pops N browsers and stalls the agent's MCP startup), we let the + bridge come up immediately — ``tools/list`` needs no credential — and only + run ``databricks auth login --resource`` when a request actually gets a 401. + So a browser opens only for a service you actually *use*, one at a time, and + an already-signed-in service never prompts. The login runs at most once per + session; PAT profiles have no connection OAuth to drive, so they never do it. + + The base class comes from whichever httpx the SDK uses (see ``_httpx``); the + sync and async flavours share the same generator contract. The async client + uses ``async_auth_flow``, so the blocking login runs off the event loop.""" httpx = _httpx() + # PAT auth has no interactive OAuth to drive, so never treat it as connection-backed. + connection = None if use_pat else connection_from_url(url) + login = {"attempted": False} + + def _mint(request): + # get_databricks_token honors the DATABRICKS_BEARER short-circuit and PAT + # profiles internally. A RuntimeError means the session is dead (expired + # refresh token, logged-out profile); raising from inside auth_flow would + # tear through the transport's task group and stall the process, so + # translate it into a terminal ProxyAuthError. + try: + token = get_databricks_token(workspace, profile) + except RuntimeError as exc: + raise ProxyAuthError(str(exc)) from exc + request.headers["Authorization"] = f"Bearer {token}" + + def _needs_connection_login(response) -> bool: + # A 401 from a connection-backed service means the per-user credential is + # missing. Only drive the login once per session — if it doesn't resolve + # the 401, re-running it on every subsequent request would loop browsers. + return connection is not None and response.status_code == 401 and not login["attempted"] + + def _connection_login_or_fail() -> None: + login["attempted"] = True + ok, detail = run_connection_login(url, workspace, profile=profile) + if not ok: + raise ProxyAuthError(f"connection login for '{connection}' failed: {detail}") class _DatabricksTokenAuth(httpx.Auth): def auth_flow(self, request): - # get_databricks_token honors the DATABRICKS_BEARER short-circuit and - # PAT profiles internally; --use-pat is surfaced via the env ucode set. - # A RuntimeError here means auth is dead (expired refresh token, - # logged-out profile). Raising it from inside auth_flow would tear - # through the transport's task group and stall the process until the - # client times out, so translate it into a terminal ProxyAuthError the - # caller reports cleanly. - try: - token = get_databricks_token(workspace, profile) - except RuntimeError as exc: - raise ProxyAuthError(str(exc)) from exc - request.headers["Authorization"] = f"Bearer {token}" + _mint(request) + response = yield request + if not _needs_connection_login(response): + return + _connection_login_or_fail() + _mint(request) + yield request + + async def async_auth_flow(self, request): + _mint(request) + response = yield request + if not _needs_connection_login(response): + return + # Run the blocking browser login off the event loop so the bridge's + # other pumps aren't starved while the user completes the sign-in. + await to_thread.run_sync(_connection_login_or_fail) + _mint(request) yield request return _DatabricksTokenAuth() @@ -166,9 +211,9 @@ async def _pump_upstream[T]( raise ProxyTransportError("upstream MCP transport closed unexpectedly") -async def _run(url: str, workspace: str, profile: str | None) -> None: +async def _run(url: str, workspace: str, profile: str | None, use_pat: bool = False) -> None: httpx = _httpx() - auth = _build_token_auth(workspace, profile) + auth = _build_token_auth(url, workspace, profile, use_pat=use_pat) # 2.x-native shape: hand the transport a pre-built AsyncClient carrying our # per-request auth. Works on mcp 1.28+ and 2.x; `streamable_http_client` # yields a (read, write) pair in both. @@ -237,6 +282,13 @@ def serve(url: str, workspace: str, profile: str | None = None, *, use_pat: bool "Set DATABRICKS_BEARER, or reconfigure the profile." ) + # The per-connection sign-in is NOT driven here. Doing it eagerly at startup + # blocks the agent's MCP startup and, with N configured connection-backed + # servers, pops N browsers at once and stalls until each times out. Instead + # the bridge opens immediately (tools/list needs no credential) and the login + # is driven lazily, on a 401, from the httpx auth flow (see _build_token_auth) + # — so a browser opens only for a service whose tool you actually call. + # Pre-flight the token before opening the bridge. Without this, the first # token failure surfaces from inside the transport's task group, where it can # stall the process instead of erroring out. @@ -246,7 +298,7 @@ def serve(url: str, workspace: str, profile: str | None = None, *, use_pat: bool _fail_fast(str(exc)) try: - anyio.run(_run, url, workspace, profile) + anyio.run(_run, url, workspace, profile, use_pat) except BaseException as exc: # noqa: BLE001 - re-raised unless it's a known proxy failure # Errors raised inside the transport arrive wrapped by its task group. # Report expected auth/transport failures without hiding programming bugs. diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 4007199d6..b66f8c984 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -3233,7 +3233,7 @@ def test_eligible_agents_use_managed_file_not_user_scope(self, monkeypatch): resolved=[ { "name": "system-ai-github", - "url": f"https://host{mcp.AIGW_MCP_SERVICES_PATH}github", + "url": f"https://host{mcp.AIGW_MCP_SERVICES_SEGMENT}github", } ], applied=applied, @@ -3257,7 +3257,7 @@ def test_claude_non_mcp_services_url_falls_back_while_mcp_services_goes_native( ): # Claude's managed file only takes native mcp-services entries; any other URL falls back to # the user-scope proxy (matching configure_client_mcp_server). Codex proxies both. - github_url = f"https://host{mcp.AIGW_MCP_SERVICES_PATH}github" + github_url = f"https://host{mcp.AIGW_MCP_SERVICES_SEGMENT}github" custom_url = "https://apps.example/custom/mcp" state = {"workspace": WS, "profile": None, "managed_mcp_servers": []} applied: dict = {} @@ -3353,7 +3353,7 @@ def test_failed_managed_write_falls_back_to_user_scope(self, monkeypatch): def test_migration_removes_prior_user_scope_for_managed_file_agents(self, monkeypatch): # Servers a prior configure registered at user scope for an agent now on the managed file are # unregistered from that agent. - sg_url = f"https://host{mcp.AIGW_MCP_SERVICES_PATH}sg" + sg_url = f"https://host{mcp.AIGW_MCP_SERVICES_SEGMENT}sg" previous = [{"name": "sg", "url": sg_url, "clients": ["claude", "codex"]}] state = {"workspace": WS, "profile": None, "managed_mcp_servers": list(previous)} applied: dict = {} diff --git a/tests/test_mcp_connection_login.py b/tests/test_mcp_connection_login.py new file mode 100644 index 000000000..bfb662fd5 --- /dev/null +++ b/tests/test_mcp_connection_login.py @@ -0,0 +1,115 @@ +"""Tests for the per-connection MCP login helpers (mcp_connection_login). + +Network-free: URL/connection parsing and the login runner with the subprocess +monkeypatched. +""" + +from __future__ import annotations + +import subprocess + +from ucode import mcp_connection_login as mcl + +WS = "https://ws.staging.cloud.databricks.com" +AIGW_URL = f"{WS}/ai-gateway/mcp-services/system.ai.github" + + +class TestConnectionFromUrl: + def test_plain_endpoint(self): + assert mcl.connection_from_url(AIGW_URL) == "system.ai.github" + + def test_with_trailing_path_and_query(self): + assert mcl.connection_from_url(f"{AIGW_URL}/tools/list?x=1") == "system.ai.github" + + def test_non_aigw_url_is_none(self): + assert mcl.connection_from_url(f"{WS}/api/2.0/mcp/functions/system/ai") is None + + def test_missing_service_is_none(self): + assert mcl.connection_from_url(f"{WS}/ai-gateway/mcp-services/") is None + + +class TestRunConnectionLogin: + def _fake_run(self, captured, *, returncode, stderr=""): + def _run(argv, **kwargs): + # The `--resource`-support pre-check runs `auth login --help` first. + if "--help" in argv: + return subprocess.CompletedProcess( + argv, 0, stdout="--resource stringArray", stderr="" + ) + captured.append(argv) + return subprocess.CompletedProcess(argv, returncode, stdout="", stderr=stderr) + + return _run + + def test_success_sends_resource_and_host_without_client_id(self, monkeypatch): + captured: list[list[str]] = [] + monkeypatch.setattr(mcl.subprocess, "run", self._fake_run(captured, returncode=0)) + + ok, message = mcl.run_connection_login(AIGW_URL, WS, profile="p") + + assert ok and message == "signed in" + argv = captured[0] + assert argv[:3] == ["databricks", "auth", "login"] + assert "--resource" in argv and AIGW_URL in argv + assert "--host" in argv and WS in argv + assert "--profile" in argv and "p" in argv + # Uses the CLI's default client (its own registered redirect), so no --client-id. + assert "--client-id" not in argv + + def test_nonzero_exit_reports_failure(self, monkeypatch): + # The CLI's own output streams live to stderr (the agent's MCP log), so on + # failure we return a pointer to that log rather than captured text. + captured: list[list[str]] = [] + monkeypatch.setattr(mcl.subprocess, "run", self._fake_run(captured, returncode=1)) + ok, message = mcl.run_connection_login(AIGW_URL, WS) + assert not ok + assert "did not complete" in message and "1" in message + + def test_output_is_routed_to_stderr_not_stdout(self, monkeypatch): + # stdout must never be captured to the proxy's stdout (the MCP wire); the + # CLI's URL/prompts go to this process's stderr. + seen: dict = {} + + def _run(argv, **kw): + if "--help" in argv: # the --resource pre-check; not the login call under test + return subprocess.CompletedProcess(argv, 0, stdout="--resource", stderr="") + seen.update(kw) + return subprocess.CompletedProcess(argv, 0) + + monkeypatch.setattr(mcl.subprocess, "run", _run) + ok, _ = mcl.run_connection_login(AIGW_URL, WS) + assert ok + assert seen.get("stdout") is mcl.sys.stderr + assert seen.get("stderr") is mcl.sys.stderr + assert "capture_output" not in seen + + def test_timeout_is_reported(self, monkeypatch): + def _run(argv, **kwargs): + raise subprocess.TimeoutExpired(argv, 1) + + monkeypatch.setattr(mcl.subprocess, "run", _run) + ok, message = mcl.run_connection_login(AIGW_URL, WS) + assert not ok and "timed out" in message + + def test_binary_missing_is_reported(self, monkeypatch): + def _run(argv, **kwargs): + raise OSError("not found") + + monkeypatch.setattr(mcl.subprocess, "run", _run) + ok, message = mcl.run_connection_login(AIGW_URL, WS) + assert not ok and "could not run" in message + + def test_old_cli_without_resource_flag_reports_clearly(self, monkeypatch): + # `auth login --help` lacking `--resource` => an old CLI (no databricks/cli#6621). + # We must report that clearly and never attempt the login (the flag would error). + def _run(argv, **kwargs): + if "--help" in argv: + return subprocess.CompletedProcess( + argv, 0, stdout="usage: login [--host]", stderr="" + ) + raise AssertionError("login must not run when --resource is unsupported") + + monkeypatch.setattr(mcl.subprocess, "run", _run) + ok, message = mcl.run_connection_login(AIGW_URL, WS) + assert not ok + assert "--resource" in message and "Upgrade" in message diff --git a/tests/test_mcp_proxy.py b/tests/test_mcp_proxy.py index 0ded63d9a..3b2e0b97b 100644 --- a/tests/test_mcp_proxy.py +++ b/tests/test_mcp_proxy.py @@ -61,7 +61,7 @@ def test_proxy_imports_the_streamable_http_client_shared_by_both_majors(): class TestDatabricksTokenAuth: def test_injects_bearer_from_minted_token(self, monkeypatch): monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok-123") - auth = mcp_proxy._build_token_auth(WS, "uc-dogfood") + auth = mcp_proxy._build_token_auth(URL, WS, "uc-dogfood") request = httpx.Request("POST", URL) # auth_flow is a generator that yields the (mutated) request. @@ -73,7 +73,7 @@ def test_auth_is_an_instance_of_the_selected_httpx_auth(self, monkeypatch): # The auth must subclass the *same* httpx flavor's Auth as the transport, # or the SDK's AsyncClient won't accept it. monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "t") - auth = mcp_proxy._build_token_auth(WS, None) + auth = mcp_proxy._build_token_auth(URL, WS, None) assert isinstance(auth, mcp_proxy._httpx().Auth) @@ -84,7 +84,7 @@ def test_calls_get_token_with_workspace_and_profile(self, monkeypatch): "get_databricks_token", lambda ws, profile: calls.append((ws, profile)) or "t", ) - auth = mcp_proxy._build_token_auth(WS, "myprofile") + auth = mcp_proxy._build_token_auth(URL, WS, "myprofile") list(auth.auth_flow(httpx.Request("POST", URL))) @@ -95,7 +95,7 @@ def test_mints_a_fresh_token_per_request(self, monkeypatch): # picked up mid-session without the proxy tracking expiry itself. tokens = iter(["first", "second"]) monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: next(tokens)) - auth = mcp_proxy._build_token_auth(WS, None) + auth = mcp_proxy._build_token_auth(URL, WS, None) r1 = httpx.Request("POST", URL) r2 = httpx.Request("POST", URL) @@ -107,7 +107,7 @@ def test_mints_a_fresh_token_per_request(self, monkeypatch): def test_auth_flow_yields_the_same_request(self, monkeypatch): monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "t") - auth = mcp_proxy._build_token_auth(WS, None) + auth = mcp_proxy._build_token_auth(URL, WS, None) request = httpx.Request("POST", URL) yielded = list(auth.auth_flow(request)) @@ -122,11 +122,117 @@ def boom(ws, profile): raise RuntimeError("no access token; run `databricks auth login`") monkeypatch.setattr(mcp_proxy, "get_databricks_token", boom) - auth = mcp_proxy._build_token_auth(WS, "p") + auth = mcp_proxy._build_token_auth(URL, WS, "p") with pytest.raises(mcp_proxy.ProxyAuthError, match="databricks auth login"): list(auth.auth_flow(httpx.Request("POST", URL))) + # --- lazy on-401 connection login --------------------------------------- + + @staticmethod + def _run_flow(auth, request, response): + """Drive the sync auth_flow, feeding `response` after the first yield. + Returns the yielded requests (1 = no retry, 2 = logged in and retried).""" + gen = auth.auth_flow(request) + yielded = [next(gen)] + try: + yielded.append(gen.send(response)) + except StopIteration: + pass + return yielded + + def test_on_401_from_connection_service_runs_login_then_retries(self, monkeypatch): + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + logins: list = [] + monkeypatch.setattr( + mcp_proxy, + "run_connection_login", + lambda url, ws, **k: logins.append((url, k.get("profile"))) or (True, "signed in"), + ) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p") + yielded = self._run_flow(auth, httpx.Request("POST", CONN_URL), httpx.Response(401)) + assert logins == [(CONN_URL, "p")] # login driven once, only on the 401 + assert len(yielded) == 2 # retried after signing in + assert yielded[1].headers["Authorization"] == "Bearer tok" + + def test_non_connection_401_is_not_retried(self, monkeypatch): + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + logins: list = [] + monkeypatch.setattr( + mcp_proxy, "run_connection_login", lambda *a, **k: logins.append(1) or (True, "") + ) + auth = mcp_proxy._build_token_auth(URL, WS, "p") # URL is not an mcp-services endpoint + yielded = self._run_flow(auth, httpx.Request("POST", URL), httpx.Response(401)) + assert logins == [] and len(yielded) == 1 + + def test_success_response_never_logs_in(self, monkeypatch): + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + logins: list = [] + monkeypatch.setattr( + mcp_proxy, "run_connection_login", lambda *a, **k: logins.append(1) or (True, "") + ) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p") + yielded = self._run_flow(auth, httpx.Request("POST", CONN_URL), httpx.Response(200)) + assert logins == [] and len(yielded) == 1 # tools/list etc. never trigger a browser + + def test_login_runs_at_most_once_per_session(self, monkeypatch): + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + logins: list = [] + monkeypatch.setattr( + mcp_proxy, "run_connection_login", lambda *a, **k: logins.append(1) or (True, "") + ) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p") + # Two 401s in the same session — the browser login must fire only once. + self._run_flow(auth, httpx.Request("POST", CONN_URL), httpx.Response(401)) + self._run_flow(auth, httpx.Request("POST", CONN_URL), httpx.Response(401)) + assert logins == [1] + + def test_use_pat_never_drives_connection_login(self, monkeypatch): + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + logins: list = [] + monkeypatch.setattr( + mcp_proxy, "run_connection_login", lambda *a, **k: logins.append(1) or (True, "") + ) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p", use_pat=True) + yielded = self._run_flow(auth, httpx.Request("POST", CONN_URL), httpx.Response(401)) + assert logins == [] and len(yielded) == 1 # PAT has no connection OAuth to drive + + def test_login_failure_becomes_a_proxy_auth_error(self, monkeypatch): + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + monkeypatch.setattr( + mcp_proxy, "run_connection_login", lambda *a, **k: (False, "user cancelled") + ) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p") + gen = auth.auth_flow(httpx.Request("POST", CONN_URL)) + next(gen) + with pytest.raises(mcp_proxy.ProxyAuthError, match="user cancelled"): + gen.send(httpx.Response(401)) + + def test_async_auth_flow_drives_login_on_401(self, monkeypatch): + # The proxy uses the async client, so async_auth_flow is the real path; the + # blocking login runs off the event loop via anyio.to_thread. + monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") + logins: list = [] + monkeypatch.setattr( + mcp_proxy, "run_connection_login", lambda *a, **k: logins.append(1) or (True, "") + ) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p") + + async def scenario(): + gen = auth.async_auth_flow(httpx.Request("POST", CONN_URL)) + await gen.__anext__() + try: + await gen.asend(httpx.Response(401)) + except StopAsyncIteration: + pass + await gen.aclose() + return logins + + assert anyio.run(scenario) == [1] + + +CONN_URL = f"{WS}/ai-gateway/mcp-services/system.ai.github" + class TestPump: def test_forwards_all_messages_in_order(self): @@ -230,7 +336,7 @@ async def stop_bridge(*args, **kwargs): yield monkeypatch.setattr(httpx_module, "AsyncClient", CapturingClient) - monkeypatch.setattr(mcp_proxy, "_build_token_auth", lambda *args: object()) + monkeypatch.setattr(mcp_proxy, "_build_token_auth", lambda *args, **kwargs: object()) monkeypatch.setattr(mcp_proxy, "streamable_http_client", stop_bridge) with pytest.raises(StopBridge): @@ -255,7 +361,7 @@ def fake_run(func, *args): assert captured["func"] is mcp_proxy._run # PAT is resolved inside the token mint, so _run takes no use_pat arg. - assert captured["args"] == (URL, WS, "uc-dogfood") + assert captured["args"] == (URL, WS, "uc-dogfood", False) def test_defaults_profile_none(self, monkeypatch): captured: dict = {} @@ -264,7 +370,7 @@ def test_defaults_profile_none(self, monkeypatch): mcp_proxy.serve(URL, WS) - assert captured["args"] == (URL, WS, None) + assert captured["args"] == (URL, WS, None, False) def test_use_pat_exports_the_bearer_before_serving(self, monkeypatch): # PAT auth: the profile's static token must be exported (ensure_pat_bearer)