From 397722f9b5e5246100120b6a6b8e914b4261d8ae Mon Sep 17 00:00:00 2001 From: "timothy.tamm" Date: Tue, 6 Oct 2026 03:39:17 +0000 Subject: [PATCH 1/2] Memoize Databricks CLI tokens per process and skip impossible --no-browser re-auth A ug claude launch calls get_databricks_token from many places, and a host whose first-matching CLI profile has a stale refresh token paid fail -> auth login --no-browser -> fail on every call. Current CLIs reject --no-browser, so the re-auth never did anything. Memoize tokens per workspace and resolved profile until five minutes before their earliest stated expiry, or for 60 seconds without expiry. Cache definitive failures briefly, but never cache token-fetch or re-auth timeouts and lock contention. Forced refresh bypasses the memo and does not replace a success with failure. Auth validation shares successful tokens and logins clear the memo. Only attempt non-interactive re-auth when the CLI advertises --no-browser, and skip the second fetch after a failed re-auth. Ported from databricks-eng/universe#2742026 onto Unity Gateway main at 4d8ff8d. Coordinate concurrent fetches per key, reject in-flight memo writes after login invalidation, and cover failed help probes, expiry, refresh, and concurrent callers without changing credential/profile routing. Co-authored-by: Isaac Signed-off-by: Harry Yao --- src/ucode/custom_oauth.py | 6 +- src/ucode/databricks.py | 385 ++++++++++++++++++----- tests/README.md | 8 + tests/conftest.py | 3 + tests/integration/README.md | 6 + tests/test_databricks.py | 583 +++++++++++++++++++++++++++++++++++ tests/test_mcp_web_search.py | 23 +- 7 files changed, 937 insertions(+), 77 deletions(-) diff --git a/src/ucode/custom_oauth.py b/src/ucode/custom_oauth.py index 14b571591..a5f6c4746 100644 --- a/src/ucode/custom_oauth.py +++ b/src/ucode/custom_oauth.py @@ -17,6 +17,7 @@ from ucode.constants import LOCALHOST, LOOPBACK_HOST from ucode.databricks import ( build_auth_token_argv, + clear_databricks_token_cache, databricks_cli_path, ensure_databricks_cli_version, external_bearer_configured, @@ -168,7 +169,10 @@ def ensure_custom_oauth_cli_token( "--scopes", ",".join(scope for scope in config["scopes"] if scope != "offline_access"), ] - run(login_args, timeout=CUSTOM_OAUTH_TIMEOUT_MS // 1000) + try: + run(login_args, timeout=CUSTOM_OAUTH_TIMEOUT_MS // 1000) + finally: + clear_databricks_token_cache() return get_databricks_token(workspace, profile) diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index d079cfd2d..d56ccce0f 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -22,6 +22,7 @@ import time from collections.abc import Callable from dataclasses import dataclass, field +from datetime import datetime from decimal import Decimal, InvalidOperation from email.message import Message from enum import Enum @@ -68,6 +69,14 @@ # we retry rather than treat them as an expired session. _TOKEN_CACHE_LOCK_MARKERS = ("cache update", "exit status 45") _TOKEN_FETCH_MAX_ATTEMPTS = 4 +# `get_databricks_token` is called from many places during one launch, and each +# miss costs a CLI subprocess (plus a doomed re-auth when the profile is stale). +# A minted token is reused until this close to its expiry; a token whose JSON +# carries no expiry, and a failure, are only remembered briefly so long-lived +# processes (MCP/relay proxies) still notice a re-login done elsewhere. +_TOKEN_MEMO_EXPIRY_MARGIN_S = 300 +_TOKEN_MEMO_DEFAULT_TTL_S = 60 +_TOKEN_MEMO_FAILURE_TTL_S = 60 _HTTP_GET_RETRYABLE_STATUS_CODES = frozenset({429}) _HTTP_GET_RETRY_BASE_SECONDS = 1.0 _HTTP_GET_RETRY_MAX_SECONDS = 5.0 @@ -815,6 +824,8 @@ def clear_databricks_cli_cache() -> None: """Forget cached CLI discovery/resolution (used by tests, and after an install/upgrade).""" global _DISCOVERED_DATABRICKS_CLIS_ORDERED _DISCOVERED_DATABRICKS_CLIS_ORDERED = None + with _AUTH_LOGIN_NO_BROWSER_SUPPORT_LOCK: + _AUTH_LOGIN_NO_BROWSER_SUPPORT.clear() def databricks_cli_installed() -> bool: @@ -1093,36 +1104,53 @@ def has_valid_databricks_auth(workspace: str, profile: str | None = None) -> boo # profiles for the same host, `databricks auth token --host …` refuses # to disambiguate without --profile, so resolve it from the host here. profile = profile or find_profile_name_for_host(workspace) - try: - env = build_databricks_cli_env(workspace, profile) - result = run( - [ - databricks_cli_path(), - "auth", - "token", - "--host", - workspace, - *_profile_args(profile), - "--output", - "json", - ], - check=False, - capture_output=True, - text=True, - env=env, - timeout=15, - ) - _debug( - "has_valid_databricks_auth", - _format_subprocess_result(result), - ) - if result.returncode != 0: + # Only a memoized token short-circuits: a memoized failure must not stop the + # caller's interactive login from being offered. + memo_key = _token_memo_key(workspace, profile) + memoized = _memoized_token(memo_key) + if memoized is not None and memoized.token: + return True + with _token_inflight_lock(memo_key): + # A negative memo is deliberately ignored, but another successful + # fetch may have completed while this caller waited for the key. + memoized = _memoized_token(memo_key) + if memoized is not None and memoized.token: + return True + epoch = _token_cache_epoch() + try: + env = build_databricks_cli_env(workspace, profile) + result = run( + [ + databricks_cli_path(), + "auth", + "token", + "--host", + workspace, + *_profile_args(profile), + "--output", + "json", + ], + check=False, + capture_output=True, + text=True, + env=env, + timeout=15, + ) + _debug( + "has_valid_databricks_auth", + _format_subprocess_result(result), + ) + if result.returncode != 0: + return False + data = json.loads(result.stdout or "{}") + token = data.get("access_token") + if not token: + return False + _remember_token(memo_key, token, _token_lifetime_s(data), _token_epoch=epoch) + return True + except (json.JSONDecodeError, OSError, subprocess.TimeoutExpired) as exc: + _debug("has_valid_databricks_auth", f"exception: {type(exc).__name__}: {exc}") return False - data = json.loads(result.stdout or "{}") - return bool(data.get("access_token")) - except (json.JSONDecodeError, OSError, subprocess.TimeoutExpired) as exc: - _debug("has_valid_databricks_auth", f"exception: {type(exc).__name__}: {exc}") - return False def list_profile_entries() -> list[dict]: @@ -1292,6 +1320,9 @@ def run_databricks_login(workspace: str, profile: str | None = None) -> None: raise RuntimeError("`databricks auth login` failed.") from exc except subprocess.TimeoutExpired as exc: raise RuntimeError("`databricks auth login` timed out.") from exc + finally: + # Even a failed login may have replaced the profile's tokens. + clear_databricks_token_cache() print_success("Databricks authentication complete") @@ -1353,38 +1384,197 @@ def _bearer_from_command(command: str) -> str: raise RuntimeError(f"DATABRICKS_BEARER_COMMAND {reason}. Command: {command}.{detail}") -def get_databricks_token( +@dataclass(frozen=True) +class _MemoizedToken: + token: str | None + error: str | None + valid_until: float + + +_TOKEN_MEMO: dict[tuple[str, str | None], _MemoizedToken] = {} +_TOKEN_MEMO_LOCK = threading.Lock() +_TOKEN_IN_FLIGHT_LOCKS: dict[tuple[str, str | None], threading.Lock] = {} +_TOKEN_CACHE_EPOCH = 0 +# Keyed by CLI path: whether `databricks auth login` accepts `--no-browser`. +_AUTH_LOGIN_NO_BROWSER_SUPPORT: dict[str, bool] = {} +_AUTH_LOGIN_NO_BROWSER_SUPPORT_LOCK = threading.Lock() + + +def clear_databricks_token_cache() -> None: + """Forget memoized tokens and failures (after a login, and between tests).""" + global _TOKEN_CACHE_EPOCH + with _TOKEN_MEMO_LOCK: + _TOKEN_CACHE_EPOCH += 1 + _TOKEN_MEMO.clear() + + +def _token_memo_key(workspace: str, profile: str | None) -> tuple[str, str | None]: + return workspace.rstrip("/"), profile + + +def _token_inflight_lock(key: tuple[str, str | None]) -> threading.Lock: + """Return the persistent single-flight lock for one workspace/profile key.""" + with _TOKEN_MEMO_LOCK: + lock = _TOKEN_IN_FLIGHT_LOCKS.get(key) + if lock is None: + lock = threading.Lock() + _TOKEN_IN_FLIGHT_LOCKS[key] = lock + return lock + + +def _token_cache_epoch() -> int: + with _TOKEN_MEMO_LOCK: + return _TOKEN_CACHE_EPOCH + + +def _memoized_token(key: tuple[str, str | None]) -> _MemoizedToken | None: + with _TOKEN_MEMO_LOCK: + entry = _TOKEN_MEMO.get(key) + if entry is None: + return None + if time.monotonic() >= entry.valid_until: + del _TOKEN_MEMO[key] + return None + return entry + + +def _token_lifetime_s(payload: dict) -> float | None: + """Seconds until a token expires, using the earliest stated deadline. + + The earlier of ``expires_in`` and the absolute ``expiry`` wins, so a CLI + that echoes the as-issued ``expires_in`` cannot stretch the memo past the + token. + """ + lifetimes: list[float] = [] + expires_in = payload.get("expires_in") + if isinstance(expires_in, int | float) and not isinstance(expires_in, bool): + lifetimes.append(float(expires_in)) + expiry = payload.get("expiry") + if isinstance(expiry, str) and expiry: + try: + parsed = datetime.fromisoformat(expiry.replace("Z", "+00:00")) + except ValueError: + parsed = None + if parsed is not None and parsed.tzinfo is not None: + lifetimes.append(parsed.timestamp() - time.time()) + return min(lifetimes) if lifetimes else None + + +def _remember_token( + key: tuple[str, str | None], + token: str, + lifetime_s: float | None, + *, + _token_epoch: int | None = None, +) -> None: + ttl = ( + lifetime_s - _TOKEN_MEMO_EXPIRY_MARGIN_S + if lifetime_s is not None + else _TOKEN_MEMO_DEFAULT_TTL_S + ) + with _TOKEN_MEMO_LOCK: + if _token_epoch is not None and _token_epoch != _TOKEN_CACHE_EPOCH: + return + if ttl <= 0: + # A forced refresh can replace a still-valid memo with a token that + # is already too close to expiry to reuse. Do not leave the older + # token behind for subsequent non-forced callers. + _TOKEN_MEMO.pop(key, None) + return + _TOKEN_MEMO[key] = _MemoizedToken(token, None, time.monotonic() + ttl) + + +def _remember_token_failure( + key: tuple[str, str | None], error: str, *, _token_epoch: int | None = None +) -> None: + with _TOKEN_MEMO_LOCK: + if _token_epoch is not None and _token_epoch != _TOKEN_CACHE_EPOCH: + return + _TOKEN_MEMO[key] = _MemoizedToken(None, error, time.monotonic() + _TOKEN_MEMO_FAILURE_TTL_S) + + +def _auth_login_supports_no_browser(cli: str) -> bool: + """Whether this CLI's `auth login` accepts ``--no-browser``; checked once per process. + + Current CLIs (v1.17-v1.19) reject the flag outright, so running the re-auth + just spawns a guaranteed `unknown flag` failure. Fail closed: if ``--help`` + can't be read or fails, a login with the flag would not have worked either. + """ + with _AUTH_LOGIN_NO_BROWSER_SUPPORT_LOCK: + supported = _AUTH_LOGIN_NO_BROWSER_SUPPORT.get(cli) + if supported is None: + try: + result = run( + [cli, "auth", "login", "--help"], + check=False, + capture_output=True, + text=True, + timeout=10, + ) + supported = result.returncode == 0 and "--no-browser" in ( + f"{result.stdout or ''}{result.stderr or ''}" + ) + except (OSError, subprocess.TimeoutExpired) as exc: + _debug("auth login --help", f"exception: {type(exc).__name__}: {exc}") + supported = False + _AUTH_LOGIN_NO_BROWSER_SUPPORT[cli] = supported + return supported + + +def _get_databricks_token_for_profile( workspace: str, - profile: str | None = None, + profile: str | None, *, - force_refresh: bool = False, + force_refresh: bool, ) -> str: - # ``DATABRICKS_BEARER`` is the CI escape hatch: when set, skip the - # `databricks auth token` subprocess entirely and return the pre-fetched - # bearer directly. Used by the e2e job, where the protected runner has - # no `databricks auth login` cache and `databricks auth token` only knows - # how to read user-OAuth caches (not M2M client_credentials). Mirrors the - # same short-circuit baked into ``build_auth_shell_command``. - bearer = os.environ.get("DATABRICKS_BEARER", "").strip() - if bearer: - _debug("get_databricks_token", "using DATABRICKS_BEARER env var") - return bearer + memo_key = _token_memo_key(workspace, profile) + if not force_refresh: + memoized = _memoized_token(memo_key) + if memoized is not None: + _debug( + "get_databricks_token", + f"memoized {'token' if memoized.token else 'failure'} " + f"profile={profile or ''}", + ) + if memoized.token: + return memoized.token + raise RuntimeError(memoized.error) + + # The second memo check is required after waiting for another caller's + # fetch. The lock is per workspace/profile, so unrelated profiles proceed. + with _token_inflight_lock(memo_key): + if not force_refresh: + memoized = _memoized_token(memo_key) + if memoized is not None: + _debug( + "get_databricks_token", + f"memoized {'token' if memoized.token else 'failure'} " + f"profile={profile or ''}", + ) + if memoized.token: + return memoized.token + raise RuntimeError(memoized.error) + token_epoch = _token_cache_epoch() + return _fetch_databricks_token_for_profile( + workspace, + profile, + force_refresh=force_refresh, + token_epoch=token_epoch, + ) - # ``DATABRICKS_BEARER_COMMAND`` is the same escape hatch in command form, - # for callers whose bearer expires and has to be re-minted (an external - # credential broker, a sidecar). A static env var cannot be rewritten in a - # running process, so the command is re-run on every fetch instead. - command = os.environ.get("DATABRICKS_BEARER_COMMAND", "").strip() - if command: - return _bearer_from_command(command) - _log_auth_diagnostics() - # See has_valid_databricks_auth: resolve the profile from the host when - # the caller didn't supply one, so duplicate-host cfgs don't break us. - profile = profile or find_profile_name_for_host(workspace) +def _fetch_databricks_token_for_profile( + workspace: str, + profile: str | None, + *, + force_refresh: bool, + token_epoch: int, +) -> str: + memo_key = _token_memo_key(workspace, profile) env = build_databricks_cli_env(workspace, profile) + cli = databricks_cli_path() cmd = [ - databricks_cli_path(), + cli, "auth", "token", "--host", @@ -1403,8 +1593,14 @@ def get_databricks_token( + f" profile={profile or ''}", ) - def _fetch() -> tuple[str, str]: - """Return (access_token, stderr). token is '' on any failure.""" + # A timeout or lost cache lock says nothing about the profile, so only a + # definitive refusal is memoized. + transient_failure = False + + def _fetch() -> tuple[str, float | None, str]: + """Return (access_token, lifetime_s, stderr). token is '' on any failure.""" + nonlocal transient_failure + transient_failure = False try: result = run( cmd, @@ -1416,13 +1612,15 @@ def _fetch() -> tuple[str, str]: ) _debug("auth token", _format_subprocess_result(result)) if result.returncode == 0: - return json.loads(result.stdout or "{}").get("access_token", ""), "" - return "", result.stderr or "" + payload = json.loads(result.stdout or "{}") + return payload.get("access_token", ""), _token_lifetime_s(payload), "" + return "", None, result.stderr or "" except (subprocess.TimeoutExpired, json.JSONDecodeError) as exc: _debug("auth token", f"exception: {type(exc).__name__}: {exc}") - return "", str(exc) + transient_failure = isinstance(exc, subprocess.TimeoutExpired) + return "", None, str(exc) - def _fetch_with_lock_retry() -> str: + def _fetch_with_lock_retry() -> tuple[str, float | None]: """Mint a token, retrying transient token-cache lock contention. Concurrent `databricks auth token` calls racing on the shared cache fail @@ -1430,25 +1628,27 @@ def _fetch_with_lock_retry() -> str: held only for the brief cache write, so a short jittered backoff almost always wins the next attempt. A non-lock failure returns '' immediately so the caller can fall through to the re-auth path.""" + nonlocal transient_failure for attempt in range(_TOKEN_FETCH_MAX_ATTEMPTS): - token, stderr = _fetch() + token, lifetime_s, stderr = _fetch() if token: - return token + return token, lifetime_s if not any(marker in stderr.lower() for marker in _TOKEN_CACHE_LOCK_MARKERS): - return "" + return "", None _debug("auth token", f"cache-lock contention (attempt {attempt + 1}); retrying") + transient_failure = True if attempt < _TOKEN_FETCH_MAX_ATTEMPTS - 1: time.sleep(random.uniform(0.05, 0.1 * (2**attempt))) - return "" + return "", None - token = _fetch_with_lock_retry() - if not token: + token, lifetime_s = _fetch_with_lock_retry() + if not token and _auth_login_supports_no_browser(cli): # Session may have expired — attempt non-interactive re-auth and retry once. _debug("auth token", "empty on first fetch; attempting auth login --no-browser") try: reauth = run( [ - databricks_cli_path(), + cli, "auth", "login", "--host", @@ -1463,8 +1663,17 @@ def _fetch_with_lock_retry() -> str: ) _debug("auth login", _format_subprocess_result(reauth)) except (subprocess.CalledProcessError, subprocess.TimeoutExpired) as exc: + # A failed re-auth changed nothing, so re-fetching would fail the same way. + if isinstance(exc, subprocess.TimeoutExpired): + transient_failure = True _debug("auth login", f"exception: {type(exc).__name__}: {exc}") - token = _fetch_with_lock_retry() + else: + if reauth.returncode == 0: + token, lifetime_s = _fetch_with_lock_retry() + else: + _debug("auth login", f"returned non-zero status {reauth.returncode}") + elif not token: + _debug("auth token", "empty on first fetch; CLI has no `auth login --no-browser`") if not token: profile_name = profile or find_profile_name_for_host(workspace) @@ -1475,14 +1684,52 @@ def _fetch_with_lock_retry() -> str: f" databricks auth logout --profile {profile_name}\n" f" databricks auth login --host {workspace} --profile {profile_name}" ) - raise RuntimeError( + error = ( f"Databricks CLI returned no access token for {workspace}. " "Run `databricks auth login` to re-authenticate." f"{stale_profile_hint}" ) + # A forced refresh is a retry after a rejected token; its failure must not + # shadow an earlier token that non-forced callers may still be using. + if not force_refresh and not transient_failure: + _remember_token_failure(memo_key, error, _token_epoch=token_epoch) + raise RuntimeError(error) + _remember_token(memo_key, token, lifetime_s, _token_epoch=token_epoch) return token +def get_databricks_token( + workspace: str, + profile: str | None = None, + *, + force_refresh: bool = False, +) -> str: + # ``DATABRICKS_BEARER`` is the CI escape hatch: when set, skip the + # `databricks auth token` subprocess entirely and return the pre-fetched + # bearer directly. Used by the e2e job, where the protected runner has + # no `databricks auth login` cache and `databricks auth token` only knows + # how to read user-OAuth caches (not M2M client_credentials). Mirrors the + # same short-circuit baked into ``build_auth_shell_command``. + bearer = os.environ.get("DATABRICKS_BEARER", "").strip() + if bearer: + _debug("get_databricks_token", "using DATABRICKS_BEARER env var") + return bearer + + # ``DATABRICKS_BEARER_COMMAND`` is the same escape hatch in command form, + # for callers whose bearer expires and has to be re-minted (an external + # credential broker, a sidecar). A static env var cannot be rewritten in a + # running process, so the command is re-run on every fetch instead. + command = os.environ.get("DATABRICKS_BEARER_COMMAND", "").strip() + if command: + return _bearer_from_command(command) + + _log_auth_diagnostics() + # See has_valid_databricks_auth: resolve the profile from the host when + # the caller didn't supply one, so duplicate-host cfgs don't break us. + profile = profile or find_profile_name_for_host(workspace) + return _get_databricks_token_for_profile(workspace, profile, force_refresh=force_refresh) + + def _extract_apps_payload(payload: object) -> list[dict]: if isinstance(payload, list): return [item for item in payload if isinstance(item, dict)] diff --git a/tests/README.md b/tests/README.md index c4450b54d..555980b52 100644 --- a/tests/README.md +++ b/tests/README.md @@ -62,6 +62,14 @@ selection, and errors without browser consent through the MCP handler. These component checks replace external auth/network boundaries; they do not establish live search, parent/child discovery, or classifier permission behavior. +`test_databricks.py` also covers process-local CLI token memoization, expiry +margins, short-lived failure entries, forced refresh, and login invalidation. +Concurrent callers share a per-profile fetch, and an in-flight result cannot +restore a memo invalidated by login. +Fake CLI subprocesses check that re-auth runs only when a successful help probe +advertises `--no-browser`. Search component tests verify token reuse and transient +timeout retries. These checks do not establish live login or launch-time savings. + `test_mcp_web_search_concurrency.py` drives the real stdio dispatcher with controlled HTTP and authentication boundaries. It covers concurrent results and catalog requests, the four-worker limit, active and queued cancellation, isolated worker errors, and diff --git a/tests/conftest.py b/tests/conftest.py index 2b0c9f090..a3f816161 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -84,6 +84,9 @@ def reject_privileged_write(path, _desired_text): # process; a real resolution (or a prior test's patched one) must not leak # into a test that assumes a bare "databricks" or a specific fake CLI. databricks_mod.clear_databricks_cli_cache() + # Tokens and token failures are memoized for the process; a fake CLI's answer + # in one test must not satisfy (or fail) a fetch in the next. + databricks_mod.clear_databricks_token_cache() def _workspace() -> str: diff --git a/tests/integration/README.md b/tests/integration/README.md index 542309acc..95f62ccde 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -54,6 +54,12 @@ tests in `../test_claude_search_discovery.py`, including legacy catalog fallback an explicit model override, and no available GPT model. Live search remains outside this integration suite. +CLI token memoization and capability-gated `--no-browser` re-auth are covered by +component tests in `../test_databricks.py` and `../test_mcp_web_search.py`, including +expiry, forced refresh, concurrent fetches, login invalidation, and transient failures. This suite +does not assert live token-fetch counts, stale-profile recovery, or launch-time +savings; credential/profile routing is unchanged. + The sudo-session regression checks live in `../test_managed_files.py` and `../test_cli.py`. They cover shared-worker invocation counts, shutdown/cancellation, and temporary-file failure handling without sudo. diff --git a/tests/test_databricks.py b/tests/test_databricks.py index ca619be8b..05d288a9b 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -14,6 +14,7 @@ import pytest +import ucode.custom_oauth as custom_oauth_mod import ucode.databricks as db_mod from ucode.databricks import ( CODING_AGENT_RECOMMEND_MODEL_PATH, @@ -2126,6 +2127,24 @@ def _fake_databricks(self, tmp_path, script: str) -> dict: # resolved by the kernel (not PATH lookup), so this stays runnable. return {**os.environ, "PATH": str(tmp_path)} + def _logging_fake(self, tmp_path, script: str) -> tuple[dict, Path]: + """Fake CLI that appends each invocation's argv (one line) to a log.""" + log = tmp_path / "argv.log" + env = self._fake_databricks(tmp_path, f'printf "%s\\n" "$*" >> {log}\n{script}') + return env, log + + # `auth login --help` of CLI v1.17-v1.19: no --no-browser. + _HELP_WITHOUT_NO_BROWSER = ( + ' "auth login --help") echo " --timeout duration Timeout"; exit 0 ;;\n' + ) + _HELP_WITH_NO_BROWSER = ( + ' "auth login --help") echo " --no-browser Do not open a browser"; exit 0 ;;\n' + ) + _STALE_TOKEN = ( + ' *"auth token"*) echo \'{"error_code": "UNAUTHENTICATED", "message": ' + '"the refresh token is invalid"}\'; exit 1 ;;\n' + ) + def test_returns_token_on_success(self, tmp_path, monkeypatch): env = self._fake_databricks( tmp_path, @@ -2170,6 +2189,9 @@ def test_reauths_and_retries_when_token_empty(self, tmp_path, monkeypatch): call_count.write_text("0") env = self._fake_databricks( tmp_path, + 'case "$*" in\n' + ' "auth login --help") echo " --no-browser Do not open a browser"; exit 0 ;;\n' + "esac\n" f"count=$(cat {call_count})\n" f"echo $((count + 1)) > {call_count}\n" 'case "$*" in\n' @@ -2252,6 +2274,567 @@ def test_error_suggests_logout_when_matching_profile_exists(self, tmp_path, monk assert "databricks auth logout --profile example-profile" in message assert f"databricks auth login --host {WS} --profile example-profile" in message + def test_stale_profile_is_tried_once_and_hint_survives_memoized_failure( + self, tmp_path, monkeypatch + ): + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + self._HELP_WITHOUT_NO_BROWSER + self._STALE_TOKEN + "esac\nexit 2", + ) + monkeypatch.setattr("os.environ", env) + + messages = [] + for _ in range(3): + with pytest.raises(RuntimeError) as exc_info: + get_databricks_token(WS, "stale-profile") + messages.append(str(exc_info.value)) + + calls = log.read_text().splitlines() + assert sum("auth token" in call for call in calls) == 1 + assert calls.count("auth login --help") == 1 + assert not any("--no-browser" in call for call in calls) + for message in messages: + assert "stale or invalid" in message + assert "databricks auth logout --profile stale-profile" in message + assert f"databricks auth login --host {WS} --profile stale-profile" in message + + def test_failed_help_does_not_advertise_no_browser(self, tmp_path, monkeypatch): + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + ' "auth login --help") echo "error: --no-browser is unavailable" >&2; exit 1 ;;\n' + + self._STALE_TOKEN + + ' *"--no-browser"*) echo "unexpected reauth" >&2; exit 1 ;;\n' + + "esac\nexit 2", + ) + monkeypatch.setattr("os.environ", env) + + with pytest.raises(RuntimeError, match="no access token"): + get_databricks_token(WS, "failed-help") + + calls = log.read_text().splitlines() + assert calls.count("auth login --help") == 1 + assert not any("--no-browser" in call for call in calls) + + def test_failed_supported_reauth_does_not_refetch(self, tmp_path, monkeypatch): + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + + self._HELP_WITH_NO_BROWSER + + self._STALE_TOKEN + + ' *"--no-browser"*) echo "reauth rejected" >&2; exit 1 ;;\n' + + "esac\nexit 2", + ) + monkeypatch.setattr("os.environ", env) + + with pytest.raises(RuntimeError, match="no access token"): + get_databricks_token(WS, "failed-reauth") + + calls = log.read_text().splitlines() + assert sum("auth token" in call for call in calls) == 1 + assert calls.count("auth login --help") == 1 + assert sum("--no-browser" in call for call in calls) == 1 + + def test_skips_no_browser_reauth_when_cli_lacks_the_flag(self, tmp_path, monkeypatch): + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + + self._HELP_WITHOUT_NO_BROWSER + + ' *"--no-browser"*) echo "Error: unknown flag: --no-browser" >&2; exit 1 ;;\n' + + "esac\n" + 'echo \'{"access_token": ""}\'\n', + ) + monkeypatch.setattr("os.environ", env) + + with pytest.raises(RuntimeError, match="no access token"): + get_databricks_token(WS, "p1") + with pytest.raises(RuntimeError, match="no access token"): + get_databricks_token(WS, "p2") + + calls = log.read_text().splitlines() + # One fetch per profile; the help probe runs once per process, not per profile. + assert sum("auth token" in call for call in calls) == 2 + assert calls.count("auth login --help") == 1 + assert not any("--no-browser" in call for call in calls) + + def test_successful_token_is_memoized_and_force_refresh_bypasses_it( + self, tmp_path, monkeypatch + ): + counter = tmp_path / "count" + counter.write_text("0") + env, log = self._logging_fake( + tmp_path, + f"read n < {counter}; n=$((n + 1)); echo $n > {counter}\n" + 'echo "{\\"access_token\\": \\"token-$n\\", \\"expires_in\\": 3600}"', + ) + monkeypatch.setattr("os.environ", env) + + assert get_databricks_token(WS, "p") == "token-1" + assert get_databricks_token(WS, "p") == "token-1" + assert get_databricks_token(WS, "p", force_refresh=True) == "token-2" + # The forced mint replaces the memo for later non-forced callers. + assert get_databricks_token(WS, "p") == "token-2" + # Profiles are memoized independently. + assert get_databricks_token(WS, "other") == "token-3" + + calls = log.read_text().splitlines() + assert len(calls) == 3 + assert "--force-refresh" in calls[1] + + def test_concurrent_same_key_misses_are_single_flight(self, monkeypatch): + fetch_started = threading.Event() + release_fetch = threading.Event() + barrier = threading.Barrier(2) + calls = [] + results = {} + errors = [] + calls_lock = threading.Lock() + results_lock = threading.Lock() + lock_requests_lock = threading.Lock() + second_lock_requested = threading.Event() + lock_requests = 0 + target_key = db_mod._token_memo_key(WS, "same-profile") + original_inflight_lock = db_mod._token_inflight_lock + + def fake_run(command, **kwargs): + with calls_lock: + calls.append(command) + fetch_started.set() + release_fetch.wait(timeout=2) + return subprocess.CompletedProcess( + command, 0, json.dumps({"access_token": "shared", "expires_in": 3600}), "" + ) + + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(db_mod, "run", fake_run) + + def observe_inflight_lock(key): + nonlocal lock_requests + lock = original_inflight_lock(key) + if key == target_key: + with lock_requests_lock: + lock_requests += 1 + if lock_requests == 2: + second_lock_requested.set() + return lock + + monkeypatch.setattr(db_mod, "_token_inflight_lock", observe_inflight_lock) + + def worker(name, operation): + try: + barrier.wait(timeout=2) + value = operation() + with results_lock: + results[name] = value + except BaseException as exc: # pragma: no cover - surfaced below + with results_lock: + errors.append(exc) + + threads = [ + threading.Thread( + target=worker, + args=("get", lambda: get_databricks_token(WS, "same-profile")), + ), + threading.Thread( + target=worker, + args=("has_valid", lambda: db_mod.has_valid_databricks_auth(WS, "same-profile")), + ), + ] + started_threads = [] + fetch_seen = False + second_lock_seen = False + try: + for thread in threads: + thread.start() + started_threads.append(thread) + fetch_seen = fetch_started.wait(timeout=2) + second_lock_seen = second_lock_requested.wait(timeout=2) + finally: + release_fetch.set() + for thread in started_threads: + thread.join(timeout=2) + + assert fetch_seen + assert second_lock_seen + assert all(not thread.is_alive() for thread in threads) + assert errors == [] + assert results == {"get": "shared", "has_valid": True} + with calls_lock: + assert len(calls) == 1 + + def test_waiting_failed_caller_cannot_shadow_successful_memo(self, monkeypatch): + fetch_started = threading.Event() + release_first_fetch = threading.Event() + success_stored = threading.Event() + calls = [] + results = [] + errors = [] + calls_lock = threading.Lock() + results_lock = threading.Lock() + original_remember = db_mod._remember_token + lock_requests_lock = threading.Lock() + second_lock_requested = threading.Event() + lock_requests = 0 + target_key = db_mod._token_memo_key(WS, "race-profile") + original_inflight_lock = db_mod._token_inflight_lock + + def fake_run(command, **kwargs): + with calls_lock: + call_number = len(calls) + 1 + calls.append(command) + if call_number == 1: + fetch_started.set() + release_first_fetch.wait(timeout=2) + return subprocess.CompletedProcess( + command, 0, json.dumps({"access_token": "winner", "expires_in": 3600}), "" + ) + success_stored.wait(timeout=2) + return subprocess.CompletedProcess(command, 1, "invalid refresh token", "") + + def remember_token(*args, **kwargs): + result = original_remember(*args, **kwargs) + if args[1] == "winner": + success_stored.set() + return result + + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(db_mod, "run", fake_run) + monkeypatch.setattr(db_mod, "_remember_token", remember_token) + + def observe_inflight_lock(key): + nonlocal lock_requests + lock = original_inflight_lock(key) + if key == target_key: + with lock_requests_lock: + lock_requests += 1 + if lock_requests == 2: + second_lock_requested.set() + return lock + + monkeypatch.setattr(db_mod, "_token_inflight_lock", observe_inflight_lock) + + def first_worker(): + try: + assert get_databricks_token(WS, "race-profile") == "winner" + with results_lock: + results.append("first") + except BaseException as exc: # pragma: no cover - surfaced below + with results_lock: + errors.append(exc) + + def second_worker(): + try: + assert get_databricks_token(WS, "race-profile") == "winner" + with results_lock: + results.append("second") + except BaseException as exc: # pragma: no cover - surfaced below + with results_lock: + errors.append(exc) + + first = threading.Thread(target=first_worker) + second = threading.Thread(target=second_worker) + started_threads = [] + fetch_seen = False + second_lock_seen = False + try: + first.start() + started_threads.append(first) + fetch_seen = fetch_started.wait(timeout=2) + second.start() + started_threads.append(second) + second_lock_seen = second_lock_requested.wait(timeout=2) + finally: + release_first_fetch.set() + for thread in started_threads: + thread.join(timeout=2) + + assert fetch_seen + assert second_lock_seen + assert all(not thread.is_alive() for thread in (first, second)) + assert errors == [] + assert sorted(results) == ["first", "second"] + with calls_lock: + assert len(calls) == 1 + + def test_inflight_fetch_cannot_restore_cache_after_login_clear(self, monkeypatch): + fetch_started = threading.Event() + release_fetch = threading.Event() + calls = [] + calls_lock = threading.Lock() + result = [] + errors = [] + + def fake_run(command, **kwargs): + with calls_lock: + calls.append(command) + call_number = len(calls) + if call_number == 1: + fetch_started.set() + release_fetch.wait(timeout=2) + token = "old-inflight" + else: + token = "fresh-after-login" + return subprocess.CompletedProcess( + command, 0, json.dumps({"access_token": token, "expires_in": 3600}), "" + ) + + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(db_mod, "run", fake_run) + + def worker(): + try: + result.append(get_databricks_token(WS, "login-profile")) + except BaseException as exc: # pragma: no cover - surfaced below + errors.append(exc) + + thread = threading.Thread(target=worker) + thread.start() + assert fetch_started.wait(timeout=2) + db_mod.clear_databricks_token_cache() + release_fetch.set() + thread.join(timeout=2) + + assert not thread.is_alive() + assert errors == [] + assert result == ["old-inflight"] + assert get_databricks_token(WS, "login-profile") == "fresh-after-login" + assert len(calls) == 2 + + def test_supported_reauth_timeout_is_transient_across_retries(self, monkeypatch): + calls = [] + + def fake_run(command, **kwargs): + calls.append((command, kwargs.get("timeout"))) + if command[-1] == "--help": + return subprocess.CompletedProcess(command, 0, "--no-browser", "") + if "--no-browser" in command: + raise subprocess.TimeoutExpired(command, kwargs["timeout"]) + return subprocess.CompletedProcess(command, 0, '{"access_token": ""}', "") + + monkeypatch.delenv("DATABRICKS_BEARER", raising=False) + monkeypatch.delenv("DATABRICKS_BEARER_COMMAND", raising=False) + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(db_mod, "run", fake_run) + + for _ in range(2): + with pytest.raises(RuntimeError, match="no access token"): + get_databricks_token(WS, "timeout-profile") + + assert [timeout for _, timeout in calls] == [15, 10, 30, 15, 30] + assert sum("auth token" in " ".join(command) for command, _ in calls) == 2 + assert sum("--no-browser" in command for command, _ in calls) == 2 + + def test_failed_force_refresh_retains_a_still_valid_success(self, tmp_path, monkeypatch): + counter = tmp_path / "count" + counter.write_text("0") + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + + self._HELP_WITHOUT_NO_BROWSER + + "esac\n" + + f"read n < {counter}; n=$((n + 1)); echo $n > {counter}\n" + + 'if [ "$n" -eq 1 ]; then echo \'{"access_token": "old", "expires_in": 3600}\'; ' + + 'else echo "refresh failed" >&2; exit 1; fi', + ) + monkeypatch.setattr("os.environ", env) + + assert get_databricks_token(WS, "p") == "old" + with pytest.raises(RuntimeError, match="no access token"): + get_databricks_token(WS, "p", force_refresh=True) + assert get_databricks_token(WS, "p") == "old" + + calls = log.read_text().splitlines() + assert sum("auth token" in call for call in calls) == 2 + assert calls.count("auth login --help") == 1 + + def test_force_refresh_retries_a_memoized_failure(self, tmp_path, monkeypatch): + flag = tmp_path / "logged-in" + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + + self._HELP_WITHOUT_NO_BROWSER + + "esac\n" + + f'if [ -f {flag} ]; then echo \'{{"access_token": "fresh"}}\'; ' + + 'else echo "invalid refresh token" >&2; exit 1; fi', + ) + monkeypatch.setattr("os.environ", env) + + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") + flag.write_text("") + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") + assert get_databricks_token(WS, "p", force_refresh=True) == "fresh" + assert sum("auth token" in call for call in log.read_text().splitlines()) == 2 + + def test_forced_near_expiry_token_evicts_an_older_memo(self, tmp_path, monkeypatch): + counter = tmp_path / "count" + counter.write_text("0") + env, log = self._logging_fake( + tmp_path, + f"read n < {counter}; n=$((n + 1)); echo $n > {counter}\n" + 'case "$n" in\n' + ' 1) echo \'{"access_token": "old", "expires_in": 3600}\' ;;\n' + ' 2) echo \'{"access_token": "near", "expires_in": 120}\' ;;\n' + ' *) echo \'{"access_token": "next", "expires_in": 3600}\' ;;\n' + "esac", + ) + monkeypatch.setattr("os.environ", env) + + assert get_databricks_token(WS, "p") == "old" + assert get_databricks_token(WS, "p", force_refresh=True) == "near" + assert get_databricks_token(WS, "p") == "next" + assert sum("auth token" in call for call in log.read_text().splitlines()) == 3 + + @pytest.mark.parametrize( + "payload", + [ + {"access_token": "t", "expires_in": 120}, + {"access_token": "t", "expiry": "2000-01-01T00:00:00.123456789Z"}, + ], + ids=["expires_in-within-margin", "expiry-in-the-past"], + ) + def test_token_near_expiry_is_not_memoized(self, tmp_path, monkeypatch, payload): + env, log = self._logging_fake(tmp_path, f"echo '{json.dumps(payload)}'") + monkeypatch.setattr("os.environ", env) + + assert get_databricks_token(WS, "p") == "t" + assert get_databricks_token(WS, "p") == "t" + assert len(log.read_text().splitlines()) == 2 + + def test_token_without_expiry_memo_expires_after_default_ttl(self, tmp_path, monkeypatch): + now = [1000.0] + monkeypatch.setattr(db_mod.time, "monotonic", lambda: now[0]) + counter = tmp_path / "count" + counter.write_text("0") + env, log = self._logging_fake( + tmp_path, + f"read n < {counter}; n=$((n + 1)); echo $n > {counter}\n" + 'echo "{\\"access_token\\": \\"token-$n\\"}"', + ) + monkeypatch.setattr("os.environ", env) + + assert get_databricks_token(WS, "p") == "token-1" + assert get_databricks_token(WS, "p") == "token-1" + now[0] += db_mod._TOKEN_MEMO_DEFAULT_TTL_S + assert get_databricks_token(WS, "p") == "token-2" + assert sum("auth token" in call for call in log.read_text().splitlines()) == 2 + + def test_parses_cli_expiry_timestamp(self, monkeypatch): + monkeypatch.setattr(db_mod.time, "time", lambda: 1_000_000_000.0) + # 2001-09-09T01:46:40Z is epoch 1_000_000_000; the CLI emits nanoseconds. + assert db_mod._token_lifetime_s({"expiry": "2001-09-09T02:46:40.925704274Z"}) == ( + pytest.approx(3600.925704) + ) + assert db_mod._token_lifetime_s({"expires_in": 86400, "expiry": "bogus"}) == 86400 + assert db_mod._token_lifetime_s( + {"expires_in": 86400, "expiry": "2001-09-09T01:56:40Z"} + ) == pytest.approx(600) + assert db_mod._token_lifetime_s({"expiry": "bogus"}) is None + assert db_mod._token_lifetime_s({}) is None + + def test_lock_contention_failure_is_not_memoized(self, tmp_path, monkeypatch): + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + self._HELP_WITHOUT_NO_BROWSER + "esac\n" + 'echo "Error: cache update: exit status 45" >&2; exit 1', + ) + monkeypatch.setattr("os.environ", env) + monkeypatch.setattr(db_mod.time, "sleep", lambda _s: None) + + for _ in range(2): + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") + calls = log.read_text().splitlines() + assert sum("auth token" in call for call in calls) == 2 * db_mod._TOKEN_FETCH_MAX_ATTEMPTS + + def test_memoized_failure_expires(self, tmp_path, monkeypatch): + env, log = self._logging_fake( + tmp_path, 'case "$*" in\n' + self._HELP_WITHOUT_NO_BROWSER + "esac\nexit 1" + ) + monkeypatch.setattr("os.environ", env) + monkeypatch.setattr(db_mod, "_TOKEN_MEMO_FAILURE_TTL_S", 0) + + for _ in range(2): + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") + assert sum("auth token" in call for call in log.read_text().splitlines()) == 2 + + def test_has_valid_auth_shares_the_memo_but_not_its_failures(self, tmp_path, monkeypatch): + flag = tmp_path / "logged-in" + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + + self._HELP_WITHOUT_NO_BROWSER + + "esac\n" + + f'if [ -f {flag} ]; then echo \'{{"access_token": "ok", "expires_in": 3600}}\'; ' + + "else exit 1; fi", + ) + monkeypatch.setattr("os.environ", env) + + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") + flag.write_text("") + # A memoized failure must not suppress the real check that decides on login. + assert db_mod.has_valid_databricks_auth(WS, "p") + assert get_databricks_token(WS, "p") == "ok" + assert db_mod.has_valid_databricks_auth(WS, "p") + assert sum("auth token" in call for call in log.read_text().splitlines()) == 2 + + def test_interactive_login_clears_memoized_failure(self, tmp_path, monkeypatch): + flag = tmp_path / "logged-in" + env, log = self._logging_fake( + tmp_path, + 'case "$*" in\n' + + self._HELP_WITHOUT_NO_BROWSER + + f' "auth login --host"*) : > {flag}; exit 0 ;;\n' + + "esac\n" + + f'if [ -f {flag} ]; then echo \'{{"access_token": "after-login"}}\'; ' + + "else exit 1; fi", + ) + monkeypatch.setattr("os.environ", env) + + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") + db_mod.run_databricks_login(WS, "p") + assert get_databricks_token(WS, "p") == "after-login" + + def test_custom_oauth_cli_login_clears_shared_token_memo(self, monkeypatch): + monkeypatch.delenv("DATABRICKS_BEARER", raising=False) + monkeypatch.delenv("DATABRICKS_BEARER_COMMAND", raising=False) + config = { + "client_id": "custom-client", + "redirect_url": "http://localhost:8020/callback", + "scopes": ["offline_access", "all-apis"], + "profile": "custom-profile", + } + db_mod._remember_token((WS, "custom-profile"), "stale", 3600) + login_calls = [] + token_calls = [] + + monkeypatch.setattr(custom_oauth_mod, "ensure_databricks_cli_version", lambda *_: None) + monkeypatch.setattr(custom_oauth_mod, "has_valid_databricks_auth", lambda *_: False) + monkeypatch.setattr(custom_oauth_mod, "databricks_cli_path", lambda: "databricks") + + def login_run(command, **kwargs): + login_calls.append(command) + return subprocess.CompletedProcess(command, 0, "", "") + + monkeypatch.setattr(custom_oauth_mod, "run", login_run) + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + + def token_run(command, **kwargs): + token_calls.append(command) + return subprocess.CompletedProcess( + command, 0, json.dumps({"access_token": "fresh", "expires_in": 3600}), "" + ) + + monkeypatch.setattr(db_mod, "run", token_run) + + assert custom_oauth_mod.ensure_custom_oauth_cli_token(WS, config) == "fresh" + assert len(login_calls) == 1 + assert len(token_calls) == 1 + class TestGetDatabricksProfiles: def _patched_run(self, monkeypatch, payload: dict, returncode: int = 0) -> None: diff --git a/tests/test_mcp_web_search.py b/tests/test_mcp_web_search.py index ca151149e..be14a7c06 100644 --- a/tests/test_mcp_web_search.py +++ b/tests/test_mcp_web_search.py @@ -308,13 +308,19 @@ def test_search_refreshes_through_selected_cli_profile(self, monkeypatch, profil def cli_token(command, **kwargs): calls.append(command) return subprocess.CompletedProcess( - command, 0, json.dumps({"access_token": f"token-{len(calls)}"}), "" + command, + 0, + json.dumps({"access_token": f"token-{len(calls)}", "expires_in": 3600}), + "", ) monkeypatch.setattr(databricks, "run", cli_token) assert self.search() == {"content": [{"type": "text", "text": "Search result"}]} + # A still-valid token is reused; once forgotten the CLI profile mints a new one. assert self.search() == {"content": [{"type": "text", "text": "Search result"}]} - assert self.auth_headers == ["Bearer token-1", "Bearer token-2"] + databricks.clear_databricks_token_cache() + assert self.search() == {"content": [{"type": "text", "text": "Search result"}]} + assert self.auth_headers == ["Bearer token-1", "Bearer token-1", "Bearer token-2"] assert [command[command.index("--profile") + 1] for command in calls] == [ profile, profile, @@ -333,11 +339,14 @@ def timed_out(command, **kwargs): assert result["isError"] is True assert "Failed to acquire Databricks token" in result["content"][0]["text"] assert self.auth_headers == [] - assert [timeout for _, timeout in calls] == [15, 30, 15] - assert "--no-browser" in calls[1][0] - assert all( - command[command.index("--profile") + 1] == "custom-profile" for command, _ in calls - ) + # The `auth login --help` probe times out too, so no `--no-browser` re-auth is attempted. + assert [timeout for _, timeout in calls] == [15, 10] + assert calls[1][0] == ["databricks", "auth", "login", "--help"] + token_command = calls[0][0] + assert token_command[token_command.index("--profile") + 1] == "custom-profile" + # A timeout is transient, so the next search asks the CLI again rather than a memo. + self.search() + assert [timeout for _, timeout in calls] == [15, 10, 15] @pytest.mark.parametrize("profile", ["workspace-profile", "custom-profile"]) def test_profile_auth_preserves_explicit_bearer_override(self, monkeypatch, profile): From ef1d2693ff8da309d7b8c058cbf9d1499f54dfe0 Mon Sep 17 00:00:00 2001 From: Harry Yao Date: Thu, 8 Oct 2026 06:06:15 +0000 Subject: [PATCH 2/2] Address token memoization review feedback Require explicit invalid or missing credential diagnostics before negative caching, expire memos when either wall or monotonic time reaches its deadline, and clear the shared token memo after every attempted MCP connection login. Add regressions for transient failure recovery, suspend and clock rollback, login failure paths, and fresh-token use on sync and async proxy retries. Update the component coverage notes. Signed-off-by: Harry Yao --- src/ucode/databricks.py | 73 +++++++++----- src/ucode/mcp_connection_login.py | 8 +- tests/README.md | 10 +- tests/integration/README.md | 5 +- tests/test_databricks.py | 156 ++++++++++++++++++++++++++--- tests/test_mcp_connection_login.py | 30 ++++++ tests/test_mcp_proxy.py | 63 +++++++++++- 7 files changed, 304 insertions(+), 41 deletions(-) diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index d56ccce0f..080422dd0 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -69,6 +69,17 @@ # we retry rather than treat them as an expired session. _TOKEN_CACHE_LOCK_MARKERS = ("cache update", "exit status 45") _TOKEN_FETCH_MAX_ATTEMPTS = 4 +# Only explicit credential refusals justify a negative memo; unknown CLI errors +# can be network, VPN, or service failures that the next call should retry. +_TOKEN_CREDENTIAL_FAILURE_MARKERS = ( + "invalid_grant", + "invalid refresh token", + "refresh token is invalid", + "expired refresh token", + "refresh token has expired", + "no cached token", + "no cached credentials", +) # `get_databricks_token` is called from many places during one launch, and each # miss costs a CLI subprocess (plus a doomed re-auth when the profile is stale). # A minted token is reused until this close to its expiry; a token whose JSON @@ -1389,6 +1400,7 @@ class _MemoizedToken: token: str | None error: str | None valid_until: float + wall_valid_until: float _TOKEN_MEMO: dict[tuple[str, str | None], _MemoizedToken] = {} @@ -1432,7 +1444,9 @@ def _memoized_token(key: tuple[str, str | None]) -> _MemoizedToken | None: entry = _TOKEN_MEMO.get(key) if entry is None: return None - if time.monotonic() >= entry.valid_until: + # Monotonic time may pause during suspend; wall time catches expiry + # during sleep, while monotonic time bounds reuse after clock rollback. + if time.monotonic() >= entry.valid_until or time.time() >= entry.wall_valid_until: del _TOKEN_MEMO[key] return None return entry @@ -1481,7 +1495,7 @@ def _remember_token( # token behind for subsequent non-forced callers. _TOKEN_MEMO.pop(key, None) return - _TOKEN_MEMO[key] = _MemoizedToken(token, None, time.monotonic() + ttl) + _TOKEN_MEMO[key] = _MemoizedToken(token, None, time.monotonic() + ttl, time.time() + ttl) def _remember_token_failure( @@ -1490,7 +1504,12 @@ def _remember_token_failure( with _TOKEN_MEMO_LOCK: if _token_epoch is not None and _token_epoch != _TOKEN_CACHE_EPOCH: return - _TOKEN_MEMO[key] = _MemoizedToken(None, error, time.monotonic() + _TOKEN_MEMO_FAILURE_TTL_S) + _TOKEN_MEMO[key] = _MemoizedToken( + None, + error, + time.monotonic() + _TOKEN_MEMO_FAILURE_TTL_S, + time.time() + _TOKEN_MEMO_FAILURE_TTL_S, + ) def _auth_login_supports_no_browser(cli: str) -> bool: @@ -1593,14 +1612,14 @@ def _fetch_databricks_token_for_profile( + f" profile={profile or ''}", ) - # A timeout or lost cache lock says nothing about the profile, so only a - # definitive refusal is memoized. - transient_failure = False + # Treat unknown failures as transient unless the CLI explicitly identifies + # invalid or missing credentials. + transient_failure = True def _fetch() -> tuple[str, float | None, str]: - """Return (access_token, lifetime_s, stderr). token is '' on any failure.""" + """Return (access_token, lifetime_s, diagnostics). token is '' on failure.""" nonlocal transient_failure - transient_failure = False + transient_failure = True try: result = run( cmd, @@ -1613,11 +1632,16 @@ def _fetch() -> tuple[str, float | None, str]: _debug("auth token", _format_subprocess_result(result)) if result.returncode == 0: payload = json.loads(result.stdout or "{}") + if not isinstance(payload, dict): + return "", None, "Databricks CLI token response is not a JSON object" return payload.get("access_token", ""), _token_lifetime_s(payload), "" - return "", None, result.stderr or "" + diagnostics = f"{result.stderr or ''}\n{result.stdout or ''}" + transient_failure = not any( + marker in diagnostics.lower() for marker in _TOKEN_CREDENTIAL_FAILURE_MARKERS + ) + return "", None, diagnostics except (subprocess.TimeoutExpired, json.JSONDecodeError) as exc: _debug("auth token", f"exception: {type(exc).__name__}: {exc}") - transient_failure = isinstance(exc, subprocess.TimeoutExpired) return "", None, str(exc) def _fetch_with_lock_retry() -> tuple[str, float | None]: @@ -1630,10 +1654,10 @@ def _fetch_with_lock_retry() -> tuple[str, float | None]: so the caller can fall through to the re-auth path.""" nonlocal transient_failure for attempt in range(_TOKEN_FETCH_MAX_ATTEMPTS): - token, lifetime_s, stderr = _fetch() + token, lifetime_s, diagnostics = _fetch() if token: return token, lifetime_s - if not any(marker in stderr.lower() for marker in _TOKEN_CACHE_LOCK_MARKERS): + if not any(marker in diagnostics.lower() for marker in _TOKEN_CACHE_LOCK_MARKERS): return "", None _debug("auth token", f"cache-lock contention (attempt {attempt + 1}); retrying") transient_failure = True @@ -1676,19 +1700,20 @@ def _fetch_with_lock_retry() -> tuple[str, float | None]: _debug("auth token", "empty on first fetch; CLI has no `auth login --no-browser`") if not token: - profile_name = profile or find_profile_name_for_host(workspace) - stale_profile_hint = "" - if profile_name: - stale_profile_hint = ( - " The saved Databricks CLI profile may be stale or invalid. Try:\n" - f" databricks auth logout --profile {profile_name}\n" - f" databricks auth login --host {workspace} --profile {profile_name}" + error = f"Databricks CLI returned no access token for {workspace}. " + if transient_failure: + error += ( + "Retry the request and check network connectivity or Databricks CLI diagnostics." ) - error = ( - f"Databricks CLI returned no access token for {workspace}. " - "Run `databricks auth login` to re-authenticate." - f"{stale_profile_hint}" - ) + else: + error += "Run `databricks auth login` to re-authenticate." + profile_name = profile or find_profile_name_for_host(workspace) + if profile_name: + error += ( + " The saved Databricks CLI profile may be stale or invalid. Try:\n" + f" databricks auth logout --profile {profile_name}\n" + f" databricks auth login --host {workspace} --profile {profile_name}" + ) # A forced refresh is a retry after a rejected token; its failure must not # shadow an earlier token that non-forced callers may still be using. if not force_refresh and not transient_failure: diff --git a/src/ucode/mcp_connection_login.py b/src/ucode/mcp_connection_login.py index e793afd8d..b4dc44a18 100644 --- a/src/ucode/mcp_connection_login.py +++ b/src/ucode/mcp_connection_login.py @@ -21,7 +21,11 @@ import subprocess import sys -from ucode.databricks import AIGW_MCP_SERVICES_SEGMENT, databricks_cli_path +from ucode.databricks import ( + AIGW_MCP_SERVICES_SEGMENT, + clear_databricks_token_cache, + databricks_cli_path, +) from ucode.os_compatibility import subprocess_cross_os # Login can pop a browser and wait for the user to complete the SaaS login, so @@ -128,6 +132,8 @@ def run_connection_login( 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" + finally: + clear_databricks_token_cache() if result.returncode == 0: return True, "signed in" return ( diff --git a/tests/README.md b/tests/README.md index 555980b52..04c050757 100644 --- a/tests/README.md +++ b/tests/README.md @@ -63,12 +63,18 @@ component checks replace external auth/network boundaries; they do not establish live search, parent/child discovery, or classifier permission behavior. `test_databricks.py` also covers process-local CLI token memoization, expiry -margins, short-lived failure entries, forced refresh, and login invalidation. +margins on both wall and monotonic clocks (including suspend and clock rollback), +short-lived failure entries only for confirmed invalid/missing credentials, +forced refresh, and login invalidation. Unknown CLI failures, network errors, +and malformed token responses retry on the next call without a stale-profile hint. Concurrent callers share a per-profile fetch, and an in-flight result cannot restore a memo invalidated by login. Fake CLI subprocesses check that re-auth runs only when a successful help probe advertises `--no-browser`. Search component tests verify token reuse and transient -timeout retries. These checks do not establish live login or launch-time savings. +timeout retries. `test_mcp_connection_login.py` verifies token/failure invalidation +after successful and failed login attempts; `test_mcp_proxy.py` checks fresh tokens +on sync and async post-login retries. These checks do not establish live login or +launch-time savings. `test_mcp_web_search_concurrency.py` drives the real stdio dispatcher with controlled HTTP and authentication boundaries. It covers concurrent results and catalog requests, diff --git a/tests/integration/README.md b/tests/integration/README.md index 95f62ccde..735eb0627 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -56,7 +56,10 @@ outside this integration suite. CLI token memoization and capability-gated `--no-browser` re-auth are covered by component tests in `../test_databricks.py` and `../test_mcp_web_search.py`, including -expiry, forced refresh, concurrent fetches, login invalidation, and transient failures. This suite +wall/monotonic expiry across suspend or clock rollback, forced refresh, concurrent +fetches, login invalidation, confirmed credential-error caching, and transient-failure +recovery. Connection-login component tests also check invalidation on every login +attempt and fresh tokens on sync/async proxy retries. This suite does not assert live token-fetch counts, stale-profile recovery, or launch-time savings; credential/profile routing is unchanged. diff --git a/tests/test_databricks.py b/tests/test_databricks.py index 05d288a9b..6152f6049 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -2260,9 +2260,7 @@ def test_error_suggests_logout_when_matching_profile_exists(self, tmp_path, monk ' *"auth profiles"*) echo \'{"profiles": [{"host": "' + WS + '", "name": "example-profile", "auth_type": "databricks-cli"}]}\'; exit 0 ;;\n' - ' *"auth login"*) exit 0 ;;\n' - "esac\n" - 'echo \'{"access_token": "", "token_type": "Bearer"}\'', + ' *"auth login"*) exit 0 ;;\n' + self._STALE_TOKEN + "esac\nexit 2", ) monkeypatch.setattr("os.environ", env) @@ -2298,6 +2296,95 @@ def test_stale_profile_is_tried_once_and_hint_survives_memoized_failure( assert "databricks auth logout --profile stale-profile" in message assert f"databricks auth login --host {WS} --profile stale-profile" in message + @pytest.mark.parametrize("stream", ["stdout", "stderr"]) + @pytest.mark.parametrize( + "diagnostic", + [ + "Error: INVALID_GRANT", + "invalid refresh token", + '{"error_code": "UNAUTHENTICATED", "message": "the refresh token is invalid"}', + "no cached token for this host", + "no cached credentials for this host", + ], + ) + def test_only_confirmed_credential_errors_are_memoized(self, monkeypatch, stream, diagnostic): + calls = [] + + def fake_run(command, **kwargs): + calls.append(command) + if command[-1] == "--help": + return subprocess.CompletedProcess(command, 0, "--timeout duration", "") + output = {"stdout": "", "stderr": "", stream: diagnostic} + return subprocess.CompletedProcess(command, 1, **output) + + monkeypatch.delenv("DATABRICKS_BEARER", raising=False) + monkeypatch.delenv("DATABRICKS_BEARER_COMMAND", raising=False) + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(db_mod, "run", fake_run) + + for _ in range(3): + with pytest.raises(RuntimeError, match="stale or invalid"): + get_databricks_token(WS, "p") + + assert sum(command[1:3] == ["auth", "token"] for command in calls) == 1 + assert sum(command[-1] == "--help" for command in calls) == 1 + + @pytest.mark.parametrize( + ("returncode", "stdout", "stderr"), + [ + (1, "", "dial tcp: connection refused"), + (1, "", "lookup host: no such host"), + (1, "temporary VPN failure", ""), + (1, "", ""), + (1, "", '{"error_code": "UNAUTHENTICATED"}'), + (0, "not JSON", ""), + (0, "", ""), + (0, '{"access_token": ""}', ""), + (0, "[]", ""), + (0, "null", ""), + ], + ids=[ + "network", + "dns", + "vpn-stdout", + "unknown", + "unclassified-auth-error", + "malformed-json", + "empty-output", + "empty-token", + "non-object-json", + "null-json", + ], + ) + def test_transient_token_failure_recovers_on_the_next_call( + self, monkeypatch, returncode, stdout, stderr + ): + token_calls = [] + + def fake_run(command, **kwargs): + if command[-1] == "--help": + return subprocess.CompletedProcess(command, 0, "--timeout duration", "") + token_calls.append(command) + if len(token_calls) == 1: + return subprocess.CompletedProcess(command, returncode, stdout, stderr) + return subprocess.CompletedProcess( + command, 0, '{"access_token": "recovered", "expires_in": 3600}', "" + ) + + monkeypatch.delenv("DATABRICKS_BEARER", raising=False) + monkeypatch.delenv("DATABRICKS_BEARER_COMMAND", raising=False) + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(db_mod, "run", fake_run) + + with pytest.raises(RuntimeError, match="no access token") as exc_info: + get_databricks_token(WS, "p") + assert "stale or invalid" not in str(exc_info.value) + assert "auth logout" not in str(exc_info.value) + assert "auth login" not in str(exc_info.value) + assert get_databricks_token(WS, "p") == "recovered" + assert get_databricks_token(WS, "p") == "recovered" + assert len(token_calls) == 2 + def test_failed_help_does_not_advertise_no_browser(self, tmp_path, monkeypatch): env, log = self._logging_fake( tmp_path, @@ -2702,9 +2789,13 @@ def test_token_near_expiry_is_not_memoized(self, tmp_path, monkeypatch, payload) assert get_databricks_token(WS, "p") == "t" assert len(log.read_text().splitlines()) == 2 - def test_token_without_expiry_memo_expires_after_default_ttl(self, tmp_path, monkeypatch): - now = [1000.0] - monkeypatch.setattr(db_mod.time, "monotonic", lambda: now[0]) + @pytest.mark.parametrize("expired_clock", ["monotonic", "wall"]) + def test_token_without_expiry_memo_expires_after_default_ttl( + self, tmp_path, monkeypatch, expired_clock + ): + now = {"monotonic": 1000.0, "wall": 1_000_000_000.0} + monkeypatch.setattr(db_mod.time, "monotonic", lambda: now["monotonic"]) + monkeypatch.setattr(db_mod.time, "time", lambda: now["wall"]) counter = tmp_path / "count" counter.write_text("0") env, log = self._logging_fake( @@ -2716,10 +2807,42 @@ def test_token_without_expiry_memo_expires_after_default_ttl(self, tmp_path, mon assert get_databricks_token(WS, "p") == "token-1" assert get_databricks_token(WS, "p") == "token-1" - now[0] += db_mod._TOKEN_MEMO_DEFAULT_TTL_S + now[expired_clock] += db_mod._TOKEN_MEMO_DEFAULT_TTL_S assert get_databricks_token(WS, "p") == "token-2" assert sum("auth token" in call for call in log.read_text().splitlines()) == 2 + @pytest.mark.parametrize("expired_clock", ["monotonic", "wall"]) + @pytest.mark.parametrize( + "expiry_fields", + [ + {"expires_in": 3600}, + {"expiry": "2001-09-09T02:46:40Z"}, + {"expires_in": 86400, "expiry": "2001-09-09T02:46:40Z"}, + ], + ids=["relative-expiry", "absolute-expiry", "earliest-expiry"], + ) + def test_token_memo_expires_when_either_clock_reaches_the_deadline( + self, tmp_path, monkeypatch, expired_clock, expiry_fields + ): + now = {"monotonic": 1000.0, "wall": 1_000_000_000.0} + monkeypatch.setattr(db_mod.time, "monotonic", lambda: now["monotonic"]) + monkeypatch.setattr(db_mod.time, "time", lambda: now["wall"]) + payload = {"access_token": "t", **expiry_fields} + env, log = self._logging_fake(tmp_path, f"echo '{json.dumps(payload)}'") + monkeypatch.setattr("os.environ", env) + + assert get_databricks_token(WS, "p") == "t" + now[expired_clock] += 3600 - db_mod._TOKEN_MEMO_EXPIRY_MARGIN_S - 1 + assert get_databricks_token(WS, "p") == "t" + assert len(log.read_text().splitlines()) == 1 + # Simulate suspend with a frozen monotonic clock, or a wall clock moved + # backwards while monotonic time still bounds the memo's lifetime. + if expired_clock == "monotonic": + now["wall"] -= 3600 + now[expired_clock] += 1 + assert get_databricks_token(WS, "p") == "t" + assert len(log.read_text().splitlines()) == 2 + def test_parses_cli_expiry_timestamp(self, monkeypatch): monkeypatch.setattr(db_mod.time, "time", lambda: 1_000_000_000.0) # 2001-09-09T01:46:40Z is epoch 1_000_000_000; the CLI emits nanoseconds. @@ -2748,16 +2871,24 @@ def test_lock_contention_failure_is_not_memoized(self, tmp_path, monkeypatch): calls = log.read_text().splitlines() assert sum("auth token" in call for call in calls) == 2 * db_mod._TOKEN_FETCH_MAX_ATTEMPTS - def test_memoized_failure_expires(self, tmp_path, monkeypatch): + @pytest.mark.parametrize("expired_clock", ["monotonic", "wall"]) + def test_memoized_failure_expires(self, tmp_path, monkeypatch, expired_clock): + now = {"monotonic": 1000.0, "wall": 1_000_000_000.0} + monkeypatch.setattr(db_mod.time, "monotonic", lambda: now["monotonic"]) + monkeypatch.setattr(db_mod.time, "time", lambda: now["wall"]) env, log = self._logging_fake( - tmp_path, 'case "$*" in\n' + self._HELP_WITHOUT_NO_BROWSER + "esac\nexit 1" + tmp_path, + 'case "$*" in\n' + self._HELP_WITHOUT_NO_BROWSER + self._STALE_TOKEN + "esac\nexit 2", ) monkeypatch.setattr("os.environ", env) - monkeypatch.setattr(db_mod, "_TOKEN_MEMO_FAILURE_TTL_S", 0) for _ in range(2): with pytest.raises(RuntimeError): get_databricks_token(WS, "p") + assert sum("auth token" in call for call in log.read_text().splitlines()) == 1 + now[expired_clock] += db_mod._TOKEN_MEMO_FAILURE_TTL_S + with pytest.raises(RuntimeError): + get_databricks_token(WS, "p") assert sum("auth token" in call for call in log.read_text().splitlines()) == 2 def test_has_valid_auth_shares_the_memo_but_not_its_failures(self, tmp_path, monkeypatch): @@ -2768,12 +2899,13 @@ def test_has_valid_auth_shares_the_memo_but_not_its_failures(self, tmp_path, mon + self._HELP_WITHOUT_NO_BROWSER + "esac\n" + f'if [ -f {flag} ]; then echo \'{{"access_token": "ok", "expires_in": 3600}}\'; ' - + "else exit 1; fi", + + 'else echo "invalid refresh token" >&2; exit 1; fi', ) monkeypatch.setattr("os.environ", env) with pytest.raises(RuntimeError): get_databricks_token(WS, "p") + assert db_mod._memoized_token(db_mod._token_memo_key(WS, "p")) is not None flag.write_text("") # A memoized failure must not suppress the real check that decides on login. assert db_mod.has_valid_databricks_auth(WS, "p") @@ -2790,7 +2922,7 @@ def test_interactive_login_clears_memoized_failure(self, tmp_path, monkeypatch): + f' "auth login --host"*) : > {flag}; exit 0 ;;\n' + "esac\n" + f'if [ -f {flag} ]; then echo \'{{"access_token": "after-login"}}\'; ' - + "else exit 1; fi", + + 'else echo "invalid refresh token" >&2; exit 1; fi', ) monkeypatch.setattr("os.environ", env) diff --git a/tests/test_mcp_connection_login.py b/tests/test_mcp_connection_login.py index 5180d8dd1..d78f1e465 100644 --- a/tests/test_mcp_connection_login.py +++ b/tests/test_mcp_connection_login.py @@ -10,6 +10,7 @@ import pytest +from ucode import databricks as db_mod from ucode import mcp_connection_login as mcl WS = "https://ws.staging.cloud.databricks.com" @@ -107,6 +108,32 @@ def _run(argv, **kwargs): ok, message = mcl.run_connection_login(AIGW_URL, WS) assert not ok and "could not run" in message + @pytest.mark.parametrize("outcome", ["success", "nonzero", "timeout", "oserror"]) + def test_login_attempt_clears_shared_tokens_and_failures(self, monkeypatch, outcome): + token_key = db_mod._token_memo_key(WS, "p") + failure_key = db_mod._token_memo_key(WS, "failed-profile") + db_mod._remember_token(token_key, "before-login", 3600) + db_mod._remember_token_failure(failure_key, "invalid refresh token") + assert db_mod._memoized_token(token_key) is not None + assert db_mod._memoized_token(failure_key) is not None + + def fake_run(argv, **kwargs): + if "--help" in argv: + return subprocess.CompletedProcess(argv, 0, "--resource", "") + if outcome == "timeout": + raise subprocess.TimeoutExpired(argv, 300) + if outcome == "oserror": + raise OSError("not found") + return subprocess.CompletedProcess(argv, 0 if outcome == "success" else 1) + + monkeypatch.setattr(mcl.subprocess, "run", fake_run) + + ok, _ = mcl.run_connection_login(AIGW_URL, WS, profile="p") + + assert ok == (outcome == "success") + assert db_mod._memoized_token(token_key) is None + assert db_mod._memoized_token(failure_key) is None + 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). @@ -118,6 +145,9 @@ def _run(argv, **kwargs): raise AssertionError("login must not run when --resource is unsupported") monkeypatch.setattr(mcl.subprocess, "run", _run) + token_key = db_mod._token_memo_key(WS, "p") + db_mod._remember_token(token_key, "unchanged", 3600) ok, message = mcl.run_connection_login(AIGW_URL, WS) assert not ok assert "--resource" in message and "Upgrade" in message + assert db_mod._memoized_token(token_key) is not None diff --git a/tests/test_mcp_proxy.py b/tests/test_mcp_proxy.py index 20e11dd4b..dcb534aab 100644 --- a/tests/test_mcp_proxy.py +++ b/tests/test_mcp_proxy.py @@ -2,6 +2,7 @@ from __future__ import annotations +import subprocess import tomllib from contextlib import asynccontextmanager from pathlib import Path @@ -10,7 +11,8 @@ import httpx import pytest -from ucode import mcp_proxy +from ucode import databricks as db_mod +from ucode import mcp_connection_login, mcp_proxy WS = "https://example.databricks.com" URL = f"{WS}/api/2.0/mcp/functions/system/ai" @@ -156,6 +158,65 @@ def test_on_401_from_connection_service_runs_login_then_retries(self, monkeypatc assert len(yielded) == 2 # retried after signing in assert yielded[1].headers["Authorization"] == "Bearer tok" + @pytest.mark.parametrize("async_flow", [False, True], ids=["sync", "async"]) + def test_connection_login_retry_uses_the_new_profiles_token(self, monkeypatch, async_flow): + monkeypatch.delenv("DATABRICKS_BEARER", raising=False) + monkeypatch.delenv("DATABRICKS_BEARER_COMMAND", raising=False) + monkeypatch.setattr(db_mod, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(mcp_connection_login, "databricks_cli_path", lambda: "databricks") + monkeypatch.setattr(mcp_proxy, "mcp_service_needs_connection_login", lambda *a, **k: True) + db_mod._remember_token(db_mod._token_memo_key(WS, "p"), "before-login", 3600) + token_calls = [] + login_calls = [] + + def token_run(argv, **kwargs): + token_calls.append(argv) + return subprocess.CompletedProcess( + argv, 0, '{"access_token": "after-login", "expires_in": 3600}', "" + ) + + def login_run(argv, **kwargs): + if "--help" in argv: + return subprocess.CompletedProcess(argv, 0, "--resource", "") + login_calls.append(argv) + return subprocess.CompletedProcess(argv, 0) + + monkeypatch.setattr(db_mod, "run", token_run) + monkeypatch.setattr(mcp_connection_login.subprocess_cross_os, "run", login_run) + auth = mcp_proxy._build_token_auth(CONN_URL, WS, "p") + request = httpx.Request("POST", CONN_URL) + + if async_flow: + + async def scenario(): + gen = auth.async_auth_flow(request) + try: + first = await gen.__anext__() + assert first.headers["Authorization"] == "Bearer before-login" + retried = await gen.asend(httpx.Response(401)) + assert retried.headers["Authorization"] == "Bearer after-login" + finally: + await gen.aclose() + + anyio.run(scenario) + else: + gen = auth.auth_flow(request) + try: + assert next(gen).headers["Authorization"] == "Bearer before-login" + assert gen.send(httpx.Response(401)).headers["Authorization"] == ( + "Bearer after-login" + ) + finally: + gen.close() + + assert len(login_calls) == 1 + assert "--resource" in login_calls[0] + assert login_calls[0][-2:] == ["--profile", "p"] + assert len(token_calls) == 1 + assert token_calls[0][1:3] == ["auth", "token"] + assert "--force-refresh" not in token_calls[0] + assert token_calls[0][token_calls[0].index("--profile") + 1] == "p" + def test_non_connection_401_is_not_retried(self, monkeypatch): monkeypatch.setattr(mcp_proxy, "get_databricks_token", lambda ws, profile: "tok") logins: list = []