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..080422dd0 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,25 @@ # 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 +# 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 +835,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 +1115,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 +1331,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 +1395,205 @@ 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 + wall_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 + # 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 + + +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, time.time() + 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, + time.time() + _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 +1612,14 @@ def get_databricks_token( + f" profile={profile or ''}", ) - def _fetch() -> tuple[str, str]: - """Return (access_token, stderr). token is '' on any failure.""" + # 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, diagnostics). token is '' on failure.""" + nonlocal transient_failure + transient_failure = True try: result = run( cmd, @@ -1416,13 +1631,20 @@ 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 "{}") + 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), "" + 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}") - return "", str(exc) + 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 +1652,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, diagnostics = _fetch() if token: - return token - if not any(marker in stderr.lower() for marker in _TOKEN_CACHE_LOCK_MARKERS): - return "" + return token, lifetime_s + 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 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,26 +1687,74 @@ 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) - 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." ) - raise RuntimeError( - 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: + _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/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 c4450b54d..04c050757 100644 --- a/tests/README.md +++ b/tests/README.md @@ -62,6 +62,20 @@ 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 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. `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, 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..735eb0627 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -54,6 +54,15 @@ 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 +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. + 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..6152f6049 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' @@ -2238,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) @@ -2252,6 +2272,701 @@ 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 + + @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, + '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 + + @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( + 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[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. + 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 + + @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 + self._STALE_TOKEN + "esac\nexit 2", + ) + monkeypatch.setattr("os.environ", env) + + 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): + 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 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") + 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 echo "invalid refresh token" >&2; 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_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 = [] 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):