From 005ed6babb033a6909e3c8e7f310437958f4bccb Mon Sep 17 00:00:00 2001 From: Anjali Sujithan Date: Mon, 27 Jul 2026 00:05:04 +0000 Subject: [PATCH 1/2] Add relayed auth MPS support for Claude Max/Enterprise --- src/ucode/agents/__init__.py | 33 ++++--- src/ucode/agents/claude.py | 170 ++++++++++++++++++++++++++++++++-- src/ucode/cli.py | 31 +++++-- src/ucode/databricks.py | 14 ++- src/ucode/gateway_proxy.py | 174 +++++++++++++++++++++++++++++++++++ tests/test_agent_claude.py | 64 +++++++++++++ tests/test_agents_init.py | 34 ++++++- tests/test_cli.py | 25 +++++ tests/test_databricks.py | 74 +++++++++------ tests/test_gateway_proxy.py | 120 ++++++++++++++++++++++++ 10 files changed, 677 insertions(+), 62 deletions(-) create mode 100644 src/ucode/gateway_proxy.py create mode 100644 tests/test_gateway_proxy.py diff --git a/src/ucode/agents/__init__.py b/src/ucode/agents/__init__.py index 88cb9e0..4bee914 100644 --- a/src/ucode/agents/__init__.py +++ b/src/ucode/agents/__init__.py @@ -278,25 +278,27 @@ def resolve_launch_model( def resolve_provider_models( tool: str, state: dict, provider: str | None -) -> tuple[dict | None, str | None]: +) -> tuple[dict | None, str | None, bool]: """Validate ``provider`` for ``tool`` and return the model ids to pin. - Returns ``(provider_models, error)``. ``provider_models`` is a + Returns ``(provider_models, error, relayed)``. ``provider_models`` is a ``{family: model_id}`` dict for a Bedrock-backed claude service (whose provider-side ids must be pinned explicitly), or None for an Anthropic/ - canonical service or when ``provider`` is None. A non-None ``error`` means - the provider is invalid for the tool (wrong type, missing, feature off, or a - Bedrock service with no Claude models) and the caller should not launch. + canonical service or when ``provider`` is None. ``relayed`` is True for a + credential-less Anthropic subscription relay, which the launch path wires + with the relayed overlay + refresh proxy. A non-None ``error`` means the + provider is invalid for the tool and the caller should not launch. """ if not provider: - return None, None + return None, None, False token = get_databricks_token(state["workspace"], state.get("profile")) service, error = resolve_provider_service(tool, provider, state["workspace"], token) if error or service is None: - return None, error + return None, error, False + relayed = bool(service.get("relayed")) if service["provider_type"] in BEDROCK_PROVIDER_TYPES: - return map_bedrock_claude_models(service.get("targets") or []), None - return None, None + return map_bedrock_claude_models(service.get("targets") or []), None, relayed + return None, None, relayed def configure_tool( @@ -305,6 +307,7 @@ def configure_tool( model: str | None = None, provider: str | None = None, provider_models: dict[str, str] | None = None, + relayed: bool = False, ) -> dict: result: dict | tuple[dict, str] if tool == "codex": @@ -315,7 +318,7 @@ def configure_tool( if not model and not provider: raise RuntimeError(f"A {tool} model must be selected before configuration.") result = claude.write_tool_config( - state, model, provider=provider, provider_models=provider_models + state, model, provider=provider, provider_models=provider_models, relayed=relayed ) else: # provider routing is claude/codex-only; every other tool needs a model. @@ -405,10 +408,12 @@ def configure_single_tool(tool: str, state: dict) -> dict: def _configure_one(tool: str, state: dict, provider: str | None) -> dict: """Write one tool's config, routing through ``provider`` when set.""" if provider: - provider_models, error = resolve_provider_models(tool, state, provider) + provider_models, error, relayed = resolve_provider_models(tool, state, provider) if error: raise RuntimeError(error) - return configure_tool(tool, state, None, provider=provider, provider_models=provider_models) + return configure_tool( + tool, state, None, provider=provider, provider_models=provider_models, relayed=relayed + ) if tool == "codex": return configure_tool("codex", state) state, model = resolve_launch_model(tool, state, None) @@ -477,6 +482,10 @@ def validate_tool(tool: str) -> tuple[bool, str]: spec = TOOL_SPECS[tool] binary = spec["binary"] module = _MODULES[tool] + # Some configs (e.g. claude relayed) can't be probed with a live message — + # the proxy + subscription login only exist at launch. Trust the written config. + if hasattr(module, "skip_validation") and module.skip_validation(load_state()): + return True, "" cmd = module.validate_cmd(binary) env = None if hasattr(module, "validate_env"): diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index 2e22b24..0190578 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -7,7 +7,10 @@ import os import re import shutil +import signal +import socket import subprocess +import threading from pathlib import Path from typing import cast @@ -110,6 +113,22 @@ def _resolve_web_search_model(state: dict) -> str | None: # must be replaced, not just left alone. MAXIMUM_MLFLOW_VERSION = (3, 12) +# Relayed drops the user scope to deliberately omit the stale apiKeyHelper. Only applied to relayed +# launches — normal launches keep loading user settings (hooks/permissions) as before. +_RELAYED_SETTING_SOURCES = "project,local" + + +def relayed_proxy_base_url(state: dict) -> str: + """Loopback base URL for the relayed refresh proxy, allocating a free port + on first call and caching it in state so config and launch agree.""" + port = state.get("relayed_proxy_port") + if not isinstance(port, int): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + state["relayed_proxy_port"] = port + return f"http://127.0.0.1:{port}" + def _web_search_mcp_entry(workspace: str, search_model: str, profile: str | None = None) -> dict: """Stdio MCP server entry pointing at `ucode mcp web-search`. Resolves @@ -140,6 +159,8 @@ def render_overlay( provider: str | None = None, provider_models: dict[str, str] | None = None, fable_enabled: bool = False, + relayed: bool = False, + relayed_base_url: str | None = None, ) -> tuple[dict, list[list[str]]]: """Return (overlay, managed_key_paths) for Claude settings.json. @@ -153,8 +174,20 @@ def render_overlay( understands Claude Code's own canonical model names, so no model id is pinned. A Bedrock-backed provider exposes different model ids (e.g. `us.anthropic.claude-sonnet-4-6`), passed in `provider_models` by family — - those get pinned via the `ANTHROPIC_DEFAULT_*_MODEL` env vars.""" - base_url = build_tool_base_url("claude", workspace) + those get pinned via the `ANTHROPIC_DEFAULT_*_MODEL` env vars. + + When `relayed` is set (a credential-less Anthropic subscription-relay MPS, + Claude Max/Team/Enterprise), Claude Code's own keychain OAuth must remain the + `Authorization` credential, so no `apiKeyHelper` is written (it would outrank + the subscription OAuth). The Databricks credential rides in the + `X-Databricks-AI-Gateway-Token` swap header, injected per request by a local + refresh proxy at `relayed_base_url` — not written here.""" + if relayed: + if not relayed_base_url: + raise RuntimeError("Relayed launch requires a proxy base URL.") + base_url = relayed_base_url + else: + base_url = build_tool_base_url("claude", workspace) # ANTHROPIC_CUSTOM_HEADERS is parsed as `key: value` pairs separated by # newlines (Anthropic SDK convention). Setting User-Agent here overrides # the SDK's default UA on outbound requests so the gateway can attribute @@ -165,6 +198,8 @@ def render_overlay( ] if provider: header_lines.append(f"Databricks-Model-Provider-Service: {provider}") + # Relayed: the X-Databricks-AI-Gateway-Token swap header is added per request + # by the refresh proxy, not here — a static value would go stale mid-session. custom_headers = "\n".join(header_lines) env: dict[str, str] = { "ANTHROPIC_BASE_URL": base_url, @@ -217,11 +252,14 @@ def render_overlay( env["ANTHROPIC_DEFAULT_SONNET_MODEL"] = _maybe_add_1m_suffix(claude_models["sonnet"]) if claude_models.get("haiku"): env["ANTHROPIC_DEFAULT_HAIKU_MODEL"] = claude_models["haiku"] - overlay: dict = { - "apiKeyHelper": build_auth_shell_command(workspace, profile, use_pat=use_pat), - "env": env, - } - keys: list[list[str]] = [["apiKeyHelper"]] + [["env", k] for k in env] + # Relayed omits apiKeyHelper so Claude Code's subscription OAuth stays the + # Authorization credential; every other path uses it as the gateway auth. + overlay: dict = {"env": env} + if relayed: + keys = [["env", k] for k in env] + else: + overlay["apiKeyHelper"] = build_auth_shell_command(workspace, profile, use_pat=use_pat) + keys = [["apiKeyHelper"]] + [["env", k] for k in env] # Disable Claude Code's built-in WebSearch: it declares Anthropic's hosted # `web_search_20250305` server tool, which the Databricks gateway rejects @@ -303,9 +341,13 @@ def write_tool_config( model: str | None, provider: str | None = None, provider_models: dict[str, str] | None = None, + relayed: bool = False, ) -> dict: backup_existing_file(CLAUDE_SETTINGS_PATH, CLAUDE_BACKUP_PATH) web_search_model = _resolve_web_search_model(state) + # Relayed inference points at a local refresh proxy; its loopback base URL is + # recorded in state so launch starts the proxy on the matching port. + relayed_base_url = relayed_proxy_base_url(state) if relayed else None overlay, managed_keys = render_overlay( state["workspace"], model, @@ -316,6 +358,8 @@ def write_tool_config( provider=provider, provider_models=provider_models, fable_enabled=bool(state.get("fable_enabled")), + relayed=relayed, + relayed_base_url=relayed_base_url, ) tracing_env_vars = tracing_env(state, "claude") stop_hook_command = claude_tracing_stop_hook_command() if tracing_env_vars else None @@ -334,6 +378,10 @@ def write_tool_config( existing = read_json_safe(CLAUDE_SETTINGS_PATH) merged = deep_merge_dict(existing, overlay) + # Drop any apiKeyHelper a prior non-relayed launch left in the file; relayed + # must not carry one (it would outrank the subscription OAuth). + if relayed: + merged.pop("apiKeyHelper", None) if tracing_env_vars and stop_hook_command: _upsert_tracing_stop_hook(merged, stop_hook_command) if not tracing_env_vars: @@ -361,6 +409,13 @@ def write_tool_config( if web_search_model: _register_web_search_mcp(state["workspace"], web_search_model, state.get("profile")) + # Persist relayed mode + proxy port so launch() wires the refresh proxy and + # subscription login; cleared on a non-relayed launch. + if relayed: + state["claude_relayed"] = True + else: + state.pop("claude_relayed", None) + state.pop("relayed_proxy_port", None) state = mark_tool_managed(state, "claude", managed_keys) save_state(state) return state @@ -628,7 +683,7 @@ def _merge_claude_settings(base: dict, overlay: dict) -> dict: return merged -def _build_claude_argv(binary: str, tool_args: list[str]) -> list[str]: +def _build_claude_argv(binary: str, tool_args: list[str], relayed: bool = False) -> list[str]: """Build the ``claude`` argv, composing any caller ``--settings`` with ucode's managed settings. @@ -644,24 +699,112 @@ def _build_claude_argv(binary: str, tool_args: list[str]) -> list[str]: accumulate one another's hooks. A caller ``--settings`` value ucode cannot resolve raises (see :func:`_load_caller_settings`) rather than being passed through as a second, colliding flag. + + ``relayed`` adds ``--setting-sources`` to exclude the user scope (see + :data:`_RELAYED_SETTING_SOURCES`), so a stale user-scope apiKeyHelper cannot + filter through and shadow the subscription OAuth. """ + source_args = ["--setting-sources", _RELAYED_SETTING_SOURCES] if relayed else [] caller_values, remaining = _extract_caller_settings(tool_args) if not caller_values: # No caller --settings: hand Claude ucode's settings file directly (the # common path; behavior unchanged). - return [binary, "--settings", str(CLAUDE_SETTINGS_PATH), *tool_args] + return [binary, *source_args, "--settings", str(CLAUDE_SETTINGS_PATH), *tool_args] caller_settings: dict = {} for value in caller_values: caller_settings = _merge_claude_settings(caller_settings, _load_caller_settings(value)) # ucode wins over the caller for conflicting keys (protects gateway auth); # hooks from both sides survive. merged = _merge_claude_settings(caller_settings, read_json_safe(CLAUDE_SETTINGS_PATH)) - return [binary, "--settings", json.dumps(merged, separators=(",", ":")), *remaining] + return [binary, *source_args, "--settings", json.dumps(merged, separators=(",", ":")), *remaining] + + +def _has_subscription_login() -> bool: + """True when Claude Code already holds a subscription login (`claude auth + status` exits 0). Never inspects or captures the credential itself.""" + try: + result = subprocess.run( + [SPEC["binary"], "auth", "status"], + check=False, + capture_output=True, + text=True, + timeout=30, + ) + except (OSError, subprocess.TimeoutExpired): + return False + return result.returncode == 0 + + +def _ensure_subscription_login() -> None: + """Ensure Claude Code has a persisted subscription login, running the browser + flow via `claude auth login` if not. ucode never sees or stores the token — + Claude Code persists it to its own secure store and refreshes it natively.""" + if _has_subscription_login(): + return + print_note("Opening browser to sign in with your Claude subscription...") + try: + subprocess.run([SPEC["binary"], "auth", "login"], check=True, timeout=300) + except subprocess.CalledProcessError as exc: + raise RuntimeError("`claude auth login` failed.") from exc + except subprocess.TimeoutExpired as exc: + raise RuntimeError("`claude auth login` timed out.") from exc + print_success("Claude subscription authenticated") + + +def _rewrite_relayed_port(state: dict, port: int) -> None: + """Point the persisted config + state at ``port`` after the proxy had to bind + a different port than the cached one. Keeps ANTHROPIC_BASE_URL (which Claude + Code reads) in sync with the live proxy so requests reach it.""" + state["relayed_proxy_port"] = port + save_state(state) + settings = read_json_safe(CLAUDE_SETTINGS_PATH) + env = settings.get("env") + if isinstance(env, dict): + env["ANTHROPIC_BASE_URL"] = f"http://127.0.0.1:{port}" + write_json_file(CLAUDE_SETTINGS_PATH, settings) + + +def _launch_relayed(state: dict, binary: str, tool_args: list[str]) -> None: + """Relayed launch: sign into the Claude subscription, start the loopback + refresh proxy, then run Claude Code alongside it (the proxy must outlive the + exec, so we spawn-and-wait rather than replacing the process).""" + from ucode.gateway_proxy import start_proxy + + _ensure_subscription_login() + workspace = state["workspace"] + port = state.get("relayed_proxy_port") + if not isinstance(port, int): + raise RuntimeError("Relayed proxy port was not configured; re-run `ucode claude`.") + + server, cache = start_proxy(workspace, state.get("profile"), port) + # start_proxy falls back to an OS-assigned port when the cached one is taken + # (stale proxy from a killed session). Reconcile settings + state to whatever + # it actually bound, so Claude Code connects to the live port. + bound_port = server.server_address[1] + if bound_port != port: + _rewrite_relayed_port(state, bound_port) + + server_thread = threading.Thread(target=server.serve_forever, daemon=True) + server_thread.start() + + proc = subprocess.Popen(_build_claude_argv(binary, tool_args, relayed=True)) + try: + returncode = proc.wait() + except KeyboardInterrupt: + proc.send_signal(signal.SIGINT) + returncode = proc.wait() + finally: + cache.stop() + server.shutdown() + raise SystemExit(returncode) def launch(state: dict, tool_args: list[str]) -> None: binary = SPEC["binary"] workspace = state.get("workspace") + if state.get("claude_relayed"): + _launch_relayed(state, binary, tool_args) + return if workspace: os.environ["OAUTH_TOKEN"] = get_databricks_token(workspace, state.get("profile")) exec_or_spawn(_build_claude_argv(binary, tool_args)) @@ -677,3 +820,10 @@ def validate_cmd(binary: str) -> list[str]: "--max-turns", "1", ] + + +def skip_validation(state: dict) -> bool: + """Relayed configs can't be probed with a live message: the loopback proxy + and subscription login are only established at launch, so a validation-time + request has nothing listening and would hang (and burn subscription quota).""" + return bool(state.get("claude_relayed")) diff --git a/src/ucode/cli.py b/src/ucode/cli.py index d072020..ae45971 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -64,6 +64,7 @@ load_full_state, load_state, save_state, + set_current_workspace, set_provider_service, ) from ucode.tracing import configure_tracing_command @@ -913,9 +914,15 @@ def _launch_tool( ctx: typer.Context, provider: str | None = None, skip_preflight: bool = False, + workspace: str | None = None, ) -> None: try: tool = normalize_tool(tool_name) + # An explicit --workspace targets that workspace for this launch (and + # auto-configures it if unseen), so `ucode claude --provider ... --workspace ...` + # works without a prior `ucode configure`. + if workspace: + set_current_workspace(normalize_workspace_url(workspace)) existing = load_state() # Workspaces configured with --use-pat export the profile's PAT as # DATABRICKS_BEARER up front so every auth check below (and the @@ -937,8 +944,9 @@ def _launch_tool( # Surfaces a clear error up front instead of a cryptic gateway failure # mid-session. For a Bedrock service this also returns the model ids. provider_models = None + relayed = False if provider: - provider_models, error = resolve_provider_models(tool, state, provider) + provider_models, error, relayed = resolve_provider_models(tool, state, provider) if error: raise RuntimeError(error) # Re-fetch model lists on every launch so newly-added Databricks @@ -961,9 +969,7 @@ def _launch_tool( resolved_model = None else: state, resolved_model = resolve_launch_model(tool, state, None) - state = configure_tool( - tool, state, resolved_model, provider=provider, provider_models=provider_models - ) + state = configure_tool(tool, state, resolved_model, provider=provider, provider_models=provider_models, relayed=relayed) print_section(f"ucode with {TOOL_SPECS[tool]['display']}") if provider: print_kv("Provider", provider) @@ -997,6 +1003,17 @@ def _launch_tool( ), ] +# Target this launch at a specific workspace, auto-configuring (and logging in) +# if it hasn't been set up yet — so a launch needs no prior `ucode configure`. +WorkspaceOption = Annotated[ + str | None, + typer.Option( + "--workspace", + help="Databricks workspace URL to launch against; sets up and authenticates it " + "if not already configured.", + ), +] + @app.command("codex", context_settings={"allow_extra_args": True, "ignore_unknown_options": True}) def codex_cmd( @@ -1011,9 +1028,10 @@ def codex_cmd( ), ] = None, skip_preflight: SkipPreflightOption = False, + workspace: WorkspaceOption = None, ) -> None: """Launch Codex via Databricks.""" - _launch_tool("codex", ctx, provider=provider, skip_preflight=skip_preflight) + _launch_tool("codex", ctx, provider=provider, skip_preflight=skip_preflight, workspace=workspace) @app.command("claude", context_settings={"allow_extra_args": True, "ignore_unknown_options": True}) @@ -1029,9 +1047,10 @@ def claude_cmd( ), ] = None, skip_preflight: SkipPreflightOption = False, + workspace: WorkspaceOption = None, ) -> None: """Launch Claude Code via Databricks.""" - _launch_tool("claude", ctx, provider=provider, skip_preflight=skip_preflight) + _launch_tool("claude", ctx, provider=provider, skip_preflight=skip_preflight, workspace=workspace) @app.command("gemini", context_settings={"allow_extra_args": True, "ignore_unknown_options": True}) diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index 897f4ac..80fa68f 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -1400,9 +1400,11 @@ def list_model_provider_services(workspace: str, token: str) -> tuple[list[dict] Returns ``(services, reason)`` where each service is ``{"name": "..", "provider_type": "anthropic"|..., - "targets": [model_id, ...], "allow_all_targets": bool}``. ``targets`` is the - provider-side model ids the service exposes (used to pin Bedrock model - names). A non-None ``reason`` means the listing call itself failed. + "targets": [model_id, ...], "allow_all_targets": bool, "relayed": bool}``. + ``targets`` is the provider-side model ids the service exposes (used to pin + Bedrock model names). ``relayed`` is True for a credential-less Anthropic + service (Claude Max/Team/Enterprise subscription relay). A non-None + ``reason`` means the listing call itself failed. """ hostname = workspace_hostname(workspace) url = f"https://{hostname}/api/2.1/unity-catalog/model-provider-services" @@ -1425,12 +1427,18 @@ def list_model_provider_services(workspace: str, token: str) -> tuple[list[dict] model_id = target.get("model") if isinstance(target, dict) else None if isinstance(model_id, str) and model_id: targets.append(model_id) + # Relayed = credential-less Anthropic (subscription relay). Only whether + # it's relayed matters here; the tier (Max vs Team/Enterprise) is governed + # server-side, so both launch identically. + anthropic_cfg = config.get("anthropic") + relayed = isinstance(anthropic_cfg, dict) and "relayed" in anthropic_cfg services.append( { "name": full_name, "provider_type": _provider_type_tag(config.get("provider_type")), "targets": targets, "allow_all_targets": bool(config.get("allow_all_targets")), + "relayed": relayed, } ) services.sort(key=lambda s: s["name"]) diff --git a/src/ucode/gateway_proxy.py b/src/ucode/gateway_proxy.py new file mode 100644 index 0000000..29bedff --- /dev/null +++ b/src/ucode/gateway_proxy.py @@ -0,0 +1,174 @@ +"""Loopback refresh proxy for relayed Anthropic (Claude Max/Team/Enterprise). + +A relayed Model Provider Service authenticates the caller's own Anthropic +subscription OAuth (which Claude Code owns in the `Authorization` header) and +carries a Databricks credential in the `X-Databricks-AI-Gateway-Token` swap +header. That Databricks token is short-lived and a static settings.json header +can't be refreshed, so `ucode claude` points `ANTHROPIC_BASE_URL` at this proxy +instead: it forwards every request to the workspace gateway unchanged except for +adding a freshly-minted swap header, and streams the response back verbatim. + +Security invariants (mirroring `databricks.py` token handling): + - Binds 127.0.0.1 only; never exposed off-host. + - Never logs header values or bodies. The Databricks token lives in memory, + refreshed off the request path; the Anthropic OAuth in `Authorization` is + passed through untouched and never read, stored, or logged. +""" + +from __future__ import annotations + +import threading +from email.message import Message +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import IO +from urllib import error as urllib_error +from urllib import request as urllib_request +from urllib.parse import urljoin + +from ucode.databricks import TOKEN_REFRESH_INTERVAL_SECONDS, get_databricks_token + +# Header we overwrite with the freshly-minted Databricks credential. Any +# client-supplied value is replaced, so a stale settings.json value can't leak. +_SWAP_HEADER = "X-Databricks-AI-Gateway-Token" +# Hop-by-hop headers must not be forwarded across the proxy. +_HOP_BY_HOP = frozenset( + h.lower() + for h in ("connection", "keep-alive", "proxy-authenticate", "proxy-authorization", + "te", "trailers", "transfer-encoding", "upgrade", "host", "content-length") +) +_STREAM_CHUNK = 8192 + + +class _TokenCache: + """Holds the current Databricks token, refreshed by a background thread so + minting never blocks a request.""" + + def __init__(self, workspace: str, profile: str | None) -> None: + self._workspace = workspace + self._profile = profile + self._lock = threading.Lock() + self._token = get_databricks_token(workspace, profile) + self._stop = threading.Event() + + @property + def token(self) -> str: + with self._lock: + return self._token + + def refresh(self) -> None: + token = get_databricks_token(self._workspace, self._profile, force_refresh=True) + with self._lock: + self._token = token + + def run_refresher(self) -> None: + while not self._stop.wait(TOKEN_REFRESH_INTERVAL_SECONDS): + try: + self.refresh() + except RuntimeError: + continue + + def stop(self) -> None: + self._stop.set() + + +def _forwarded_request_headers(handler: BaseHTTPRequestHandler, token: str) -> dict[str, str]: + headers = { + key: value + for key, value in handler.headers.items() + if key.lower() not in _HOP_BY_HOP and key.lower() != _SWAP_HEADER.lower() + } + headers[_SWAP_HEADER] = f"Bearer {token}" + return headers + + +class _ProxyHandler(BaseHTTPRequestHandler): + # Set by the server factory. + cache: _TokenCache + upstream_base: str + + def log_message(self, format: str, *args: object) -> None: + return + + def _handle(self) -> None: + length = int(self.headers.get("Content-Length", 0) or 0) + body = self.rfile.read(length) if length else None + target = urljoin(self.upstream_base, self.path.lstrip("/")) + req = urllib_request.Request( + target, + data=body, + method=self.command, + headers=_forwarded_request_headers(self, self.cache.token), + ) + try: + with urllib_request.urlopen(req, timeout=600) as resp: + self._relay_response(resp.status, resp.headers, resp) + except urllib_error.HTTPError as exc: + # Upstream (gateway/Anthropic) error — relay status + body verbatim so + # the agent sees the real error (e.g. 429 rate_limit_error). + self._relay_response(exc.code, exc.headers, exc) + except (urllib_error.URLError, OSError): + # The client (Claude Code) may already have disconnected, in which case + # reporting the error writes to a dead socket and raises again; swallow it. + try: + self.send_error(502, "gateway proxy upstream error") + except OSError: + pass + + # Streaming passthrough: forward chunks as they arrive so SSE token streaming + # is not buffered (buffering would add full-response latency to first token). + def _relay_response(self, status: int, headers: Message, stream: IO[bytes]) -> None: + try: + self.send_response(status) + for key, value in headers.items(): + if key.lower() not in _HOP_BY_HOP: + self.send_header(key, value) + self.end_headers() + while True: + chunk = stream.read(_STREAM_CHUNK) + if not chunk: + break + self.wfile.write(chunk) + self.wfile.flush() + except (BrokenPipeError, ConnectionResetError): + # Client (Claude Code) closed the connection mid-response — routine on + # cancelled turns / SSE teardown. There is nothing left to relay to, so + # stop quietly rather than crashing the handler thread. + return + + # Forward every method: this is a transparent pass-through, so routing any + # `do_` lookup to `_handle` lets the gateway reject unsupported methods. + def __getattr__(self, name: str): + if name.startswith("do_"): + return self._handle + raise AttributeError(name) + + +def start_proxy( + workspace: str, profile: str | None, port: int +) -> tuple[ThreadingHTTPServer, _TokenCache]: + """Start the loopback refresh proxy + its background token refresher. + + Binds ``port``, falling back to a fresh OS-assigned port when it is already + in use (e.g. a prior session's proxy that was killed before its teardown ran + still holds the socket). The caller reads ``server.server_address[1]`` for the + actual port and points Claude Code at it. Returns (server, cache); the caller + runs the server (e.g. in a thread) and calls shutdown()/cache.stop() on exit. + """ + upstream_base = f"{workspace.rstrip('/')}/ai-gateway/anthropic/" + cache = _TokenCache(workspace, profile) + + handler = type( + "BoundProxyHandler", + (_ProxyHandler,), + {"cache": cache, "upstream_base": upstream_base}, + ) + try: + server = ThreadingHTTPServer(("127.0.0.1", port), handler) + except OSError: + # Cached port is occupied (stale proxy from a killed session). Port 0 lets + # the OS pick any free port; the caller reconciles the base URL to it. + server = ThreadingHTTPServer(("127.0.0.1", 0), handler) + + refresher = threading.Thread(target=cache.run_refresher, daemon=True) + refresher.start() + return server, cache diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index b6952ee..60b4a6a 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -102,6 +102,42 @@ def test_sets_api_key_helper(self): assert "apiKeyHelper" in overlay assert WS in overlay["apiKeyHelper"] + def test_relayed_omits_api_key_helper(self): + # Claude Code's own subscription OAuth must own Authorization; an + # apiKeyHelper would outrank it. + overlay, _ = claude.render_overlay( + WS, + None, + provider="c.s.mps", + relayed=True, + relayed_base_url="http://127.0.0.1:9", + ) + assert "apiKeyHelper" not in overlay + + def test_relayed_points_base_url_at_proxy(self): + overlay, _ = claude.render_overlay( + WS, + None, + provider="c.s.mps", + relayed=True, + relayed_base_url="http://127.0.0.1:9", + ) + assert overlay["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:9" + + def test_relayed_sends_mps_header_but_not_swap_token(self): + # The MPS header selects the service; the swap token is injected by the + # proxy, never written into settings. + overlay, _ = claude.render_overlay( + WS, + None, + provider="c.s.mps", + relayed=True, + relayed_base_url="http://127.0.0.1:9", + ) + headers = overlay["env"]["ANTHROPIC_CUSTOM_HEADERS"] + assert "Databricks-Model-Provider-Service: c.s.mps" in headers + assert "X-Databricks-AI-Gateway-Token" not in headers + def test_model_overrides_when_all_provided(self): models = { "sonnet": "databricks-claude-sonnet-4-6", @@ -604,6 +640,34 @@ def test_no_caller_settings_uses_ucode_file(self, monkeypatch): argv = claude._build_claude_argv("claude", ["-p", "hi"]) assert argv == ["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "-p", "hi"] + def test_non_relayed_does_not_set_setting_sources(self, monkeypatch): + # Normal launches must keep loading user settings (hooks/permissions) — + # no --setting-sources so nothing changes for the stored-key path. + monkeypatch.setattr(claude, "read_json_safe", lambda p: {"apiKeyHelper": "u"}) + argv = claude._build_claude_argv("claude", ["-p", "hi"], relayed=False) + assert "--setting-sources" not in argv + + def test_relayed_excludes_user_scope_via_setting_sources(self, monkeypatch): + # Relayed must drop the user scope so a stale ~/.claude/settings.json + # apiKeyHelper can't merge through and shadow the subscription OAuth. + monkeypatch.setattr(claude, "read_json_safe", lambda p: {"env": {}}) + argv = claude._build_claude_argv("claude", ["-p", "hi"], relayed=True) + assert "--setting-sources" in argv + src = argv[argv.index("--setting-sources") + 1] + assert src == claude._RELAYED_SETTING_SOURCES + assert "user" not in src + # ucode's own settings file is still passed. + assert "--settings" in argv + assert str(claude.CLAUDE_SETTINGS_PATH) in argv + + def test_relayed_with_caller_settings_keeps_setting_sources(self, monkeypatch): + # Even when composing a caller --settings, relayed still excludes user scope. + monkeypatch.setattr(claude, "read_json_safe", lambda p: {"env": {}}) + caller = json.dumps({"statusLine": {"type": "command", "command": "sl"}}) + argv = claude._build_claude_argv("claude", ["--settings", caller], relayed=True) + assert argv[:3] == ["claude", "--setting-sources", claude._RELAYED_SETTING_SOURCES] + assert argv.count("--settings") == 1 + def test_inline_caller_settings_merged_into_single_flag(self, monkeypatch): ucode_settings = { "apiKeyHelper": "ucode-helper", diff --git a/tests/test_agents_init.py b/tests/test_agents_init.py index 3443aad..9b0a600 100644 --- a/tests/test_agents_init.py +++ b/tests/test_agents_init.py @@ -277,14 +277,22 @@ def _patch(self, monkeypatch, service, error): ) def test_none_provider_returns_none(self): - models, error = agents_mod.resolve_provider_models("claude", self._STATE, None) - assert (models, error) == (None, None) + models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, None) + assert (models, error, relayed) == (None, None, False) def test_anthropic_returns_no_models(self, monkeypatch): self._patch(monkeypatch, {"provider_type": "anthropic", "targets": []}, None) - models, error = agents_mod.resolve_provider_models("claude", self._STATE, "main.a.svc") + models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.a.svc") assert error is None assert models is None + assert relayed is False + + def test_relayed_anthropic_flagged(self, monkeypatch): + self._patch(monkeypatch, {"provider_type": "anthropic", "targets": [], "relayed": True}, None) + models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.a.relayed") + assert error is None + assert models is None + assert relayed is True def test_bedrock_returns_pinned_models(self, monkeypatch): service = { @@ -292,8 +300,9 @@ def test_bedrock_returns_pinned_models(self, monkeypatch): "targets": ["us.anthropic.claude-sonnet-4-6", "global.anthropic.claude-opus-4-8"], } self._patch(monkeypatch, service, None) - models, error = agents_mod.resolve_provider_models("claude", self._STATE, "main.b.svc") + models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.b.svc") assert error is None + assert relayed is False assert models == { "sonnet": "us.anthropic.claude-sonnet-4-6", "opus": "global.anthropic.claude-opus-4-8", @@ -301,9 +310,10 @@ def test_bedrock_returns_pinned_models(self, monkeypatch): def test_invalid_provider_returns_error(self, monkeypatch): self._patch(monkeypatch, None, "boom") - models, error = agents_mod.resolve_provider_models("claude", self._STATE, "main.x.svc") + models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.x.svc") assert models is None assert error == "boom" + assert relayed is False class TestInstallToolBinary: @@ -666,3 +676,17 @@ def fake_run(cmd, **kwargs): assert ok is False assert err == "timed out" + + def test_relayed_claude_skips_live_probe(self, monkeypatch): + # Relayed configs have no proxy/login at validation time; probing them + # with a live message would hang, so validation must trust the config. + def fail_run(cmd, **kwargs): + raise AssertionError("relayed validation must not run a subprocess") + + monkeypatch.setattr("ucode.agents.subprocess.run", fail_run) + monkeypatch.setattr(agents_mod, "load_state", lambda: {"claude_relayed": True}) + + ok, err = agents_mod.validate_tool("claude") + + assert ok is True + assert err == "" diff --git a/tests/test_cli.py b/tests/test_cli.py index 59cc755..9f52ed3 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -153,6 +153,31 @@ def test_no_agent_flag(self): result = runner.invoke(app, ["--agent", "claude"]) assert result.exit_code != 0 + def test_workspace_flag_sets_current_workspace(self): + """--workspace targets that workspace (normalized) before launch.""" + patches = _patch_launch("claude") + with ( + patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6], + patch("ucode.cli.set_current_workspace") as mock_set, + ): + result = runner.invoke( + app, + ["claude", "--workspace", "https://eng-ml-inference.staging.cloud.databricks.com/"], + ) + assert result.exit_code == 0, result.output + mock_set.assert_called_once_with("https://eng-ml-inference.staging.cloud.databricks.com") + + def test_no_workspace_flag_leaves_current_workspace(self): + """Without --workspace, launch never reassigns the current workspace.""" + patches = _patch_launch("claude") + with ( + patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6], + patch("ucode.cli.set_current_workspace") as mock_set, + ): + result = runner.invoke(app, ["claude"]) + assert result.exit_code == 0, result.output + mock_set.assert_not_called() + class TestMcpSubcommands: def test_web_search_subcommand_help(self): diff --git a/tests/test_databricks.py b/tests/test_databricks.py index 38649f2..f6386a0 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -281,7 +281,7 @@ def test_ignores_non_system_ai_schemas(self, monkeypatch): payload = { "model_services": [ _model_service("system.ai.gpt-5"), - _model_service("main.svenwb.gpt-5-5"), + _model_service("main.schema3.gpt-5-5"), _model_service("temp.erni.kimi-k2-7-code"), _model_service("temp.erni.claude-opus-4-8"), _model_service("dnasi_agent_cuj.default.dnasi-gpt55-test"), @@ -339,15 +339,22 @@ class TestListModelProviderServices: _PAYLOAD = { "model_provider_services": [ { - "name": "model-provider-services/main.aarushi.anthropic-svc", + "name": "model-provider-services/main.schema1.anthropic-svc", "config": {"provider_type": "EXTERNAL_MODEL_PROVIDER_TYPE_ANTHROPIC"}, }, { - "name": "model-provider-services/main.aarushi.openai-svc", + "name": "model-provider-services/main.schema1.claude-max-svc", + "config": { + "provider_type": "EXTERNAL_MODEL_PROVIDER_TYPE_ANTHROPIC", + "anthropic": {"relayed": {}}, + }, + }, + { + "name": "model-provider-services/main.schema1.openai-svc", "config": {"provider_type": "EXTERNAL_MODEL_PROVIDER_TYPE_OPENAI"}, }, { - "name": "model-provider-services/main.bob.bedrock-svc", + "name": "model-provider-services/main.schema2.bedrock-svc", "config": { "provider_type": "EXTERNAL_MODEL_PROVIDER_TYPE_AMAZON_BEDROCK", "allow_all_targets": False, @@ -361,7 +368,7 @@ class TestListModelProviderServices: }, }, { - "name": "model-provider-services/main.bob.bedrock-titan-svc", + "name": "model-provider-services/main.schema2.bedrock-titan-svc", "config": { "provider_type": "EXTERNAL_MODEL_PROVIDER_TYPE_AMAZON_BEDROCK", "targets": [{"model": "amazon.titan-text-express-v1"}], @@ -377,10 +384,11 @@ def test_strips_prefix_and_tags_provider_type(self, monkeypatch): services, reason = db_mod.list_model_provider_services(WS, "token") assert reason is None assert services[0] == { - "name": "main.aarushi.anthropic-svc", + "name": "main.schema1.anthropic-svc", "provider_type": "anthropic", "targets": [], "allow_all_targets": False, + "relayed": False, } assert {s["provider_type"] for s in services} == { "anthropic", @@ -388,12 +396,21 @@ def test_strips_prefix_and_tags_provider_type(self, monkeypatch): "amazon_bedrock", } + def test_flags_relayed_anthropic(self, monkeypatch): + monkeypatch.setattr( + db_mod, "_http_get_json", lambda url, token, timeout=30: (self._PAYLOAD, None) + ) + services, _ = db_mod.list_model_provider_services(WS, "token") + by_name = {s["name"]: s for s in services} + assert by_name["main.schema1.claude-max-svc"]["relayed"] is True + assert by_name["main.schema1.anthropic-svc"]["relayed"] is False + def test_extracts_targets(self, monkeypatch): monkeypatch.setattr( db_mod, "_http_get_json", lambda url, token, timeout=30: (self._PAYLOAD, None) ) services, _ = db_mod.list_model_provider_services(WS, "token") - bedrock = next(s for s in services if s["name"] == "main.bob.bedrock-svc") + bedrock = next(s for s in services if s["name"] == "main.schema2.bedrock-svc") assert bedrock["targets"] == [ "us.anthropic.claude-sonnet-4-6", "global.anthropic.claude-opus-4-8", @@ -413,16 +430,21 @@ def test_claude_includes_anthropic_and_usable_bedrock(self, monkeypatch): ) names, reason = db_mod.list_tool_provider_services("claude", WS, "token") assert reason is None - # Anthropic + the Bedrock service with Claude targets; the Bedrock service - # exposing only Titan is hidden (no Claude models to pin). - assert names == ["main.aarushi.anthropic-svc", "main.bob.bedrock-svc"] + # Anthropic (stored-key + relayed) + the Bedrock service with Claude + # targets; the Bedrock service exposing only Titan is hidden (no Claude + # models to pin). + assert names == [ + "main.schema1.anthropic-svc", + "main.schema1.claude-max-svc", + "main.schema2.bedrock-svc", + ] def test_codex_filters_to_openai(self, monkeypatch): monkeypatch.setattr( db_mod, "_http_get_json", lambda url, token, timeout=30: (self._PAYLOAD, None) ) names, _ = db_mod.list_tool_provider_services("codex", WS, "token") - assert names == ["main.aarushi.openai-svc"] + assert names == ["main.schema1.openai-svc"] class TestMapBedrockClaudeModels: @@ -472,7 +494,7 @@ def _patch(self, monkeypatch): def test_anthropic_ok(self, monkeypatch): self._patch(monkeypatch) service, error = db_mod.resolve_provider_service( - "claude", "main.aarushi.anthropic-svc", WS, "token" + "claude", "main.schema1.anthropic-svc", WS, "token" ) assert error is None assert service["provider_type"] == "anthropic" @@ -480,7 +502,7 @@ def test_anthropic_ok(self, monkeypatch): def test_bedrock_with_claude_ok(self, monkeypatch): self._patch(monkeypatch) service, error = db_mod.resolve_provider_service( - "claude", "main.bob.bedrock-svc", WS, "token" + "claude", "main.schema2.bedrock-svc", WS, "token" ) assert error is None assert service["provider_type"] == "amazon_bedrock" @@ -488,7 +510,7 @@ def test_bedrock_with_claude_ok(self, monkeypatch): def test_wrong_type_rejected(self, monkeypatch): self._patch(monkeypatch) service, error = db_mod.resolve_provider_service( - "claude", "main.aarushi.openai-svc", WS, "token" + "claude", "main.schema1.openai-svc", WS, "token" ) assert service is None assert "can't route to" in error @@ -496,7 +518,7 @@ def test_wrong_type_rejected(self, monkeypatch): def test_bedrock_without_claude_rejected(self, monkeypatch): self._patch(monkeypatch) service, error = db_mod.resolve_provider_service( - "claude", "main.bob.bedrock-titan-svc", WS, "token" + "claude", "main.schema2.bedrock-titan-svc", WS, "token" ) assert service is None assert "no Claude models" in error @@ -506,7 +528,7 @@ def test_not_found_lists_usable(self, monkeypatch): service, error = db_mod.resolve_provider_service("claude", "main.x.missing", WS, "token") assert service is None assert "was not found" in error - assert "main.aarushi.anthropic-svc" in error + assert "main.schema1.anthropic-svc" in error def test_feature_unavailable(self, monkeypatch): reason = "HTTP 400 Bad Request: ModelProviderService feature is not available" @@ -600,7 +622,7 @@ def test_ignores_non_system_ai_entries(self, monkeypatch): payload = { "mcp_services": [ {"name": "mcp-services/system.ai.github"}, - {"name": "mcp-services/main.svenwb.github_mcp"}, + {"name": "mcp-services/main.schema3.github_mcp"}, {"name": "mcp-services/temp.erni.github_mcp"}, ] } @@ -643,15 +665,15 @@ def fake_get(url, token, timeout=30): monkeypatch.setattr(db_mod, "_http_get_json", fake_get) - db_mod.list_mcp_services(WS, "token", parent="main.svenwb") + db_mod.list_mcp_services(WS, "token", parent="main.schema3") - assert "parent=schemas%2Fmain.svenwb" in captured["url"] + assert "parent=schemas%2Fmain.schema3" in captured["url"] def test_custom_parent_filters_to_namespace(self, monkeypatch): payload = { "mcp_services": [ - {"name": "mcp-services/main.svenwb.github"}, - {"name": "mcp-services/main.svenwb.slack"}, + {"name": "mcp-services/main.schema3.github"}, + {"name": "mcp-services/main.schema3.slack"}, {"name": "mcp-services/system.ai.github"}, ] } @@ -659,10 +681,10 @@ def test_custom_parent_filters_to_namespace(self, monkeypatch): db_mod, "_http_get_json", lambda url, token, timeout=30: (payload, None) ) - names, reason = db_mod.list_mcp_services(WS, "token", parent="main.svenwb") + names, reason = db_mod.list_mcp_services(WS, "token", parent="main.schema3") assert reason is None - assert names == ["main.svenwb.github", "main.svenwb.slack"] + assert names == ["main.schema3.github", "main.schema3.slack"] def test_http_404_reason_surfaces_for_invalid_parent(self, monkeypatch): monkeypatch.setattr( @@ -1625,7 +1647,7 @@ def test_schema_level_use_schema_denial_matches(self): def test_unrelated_catalog_denial_falls_through(self): msg = ( "[INSUFFICIENT_PERMISSIONS] Insufficient privileges: " - "User does not have USE CATALOG on Catalog 'aarushi'. " + "User does not have USE CATALOG on Catalog 'schema1'. " "SQLSTATE: 42501" ) assert db_mod._is_usage_table_access_error(self._err(msg)) is False @@ -1692,11 +1714,11 @@ def test_falls_through_for_unrelated_permission_error(self, monkeypatch): original = ServerOperationError( "[INSUFFICIENT_PERMISSIONS] Insufficient privileges: " - "User does not have USE CATALOG on Catalog 'aarushi'. SQLSTATE: 42501" + "User does not have USE CATALOG on Catalog 'schema1'. SQLSTATE: 42501" ) self._patch_connect_to_raise(monkeypatch, original) - with pytest.raises(RuntimeError, match="aarushi") as exc_info: + with pytest.raises(RuntimeError, match="schema1") as exc_info: db_mod.run_usage_query(WS, "/sql/1.0/warehouses/abc", "tok", "SELECT 1") assert "Ask your workspace admin" not in str(exc_info.value) assert str(exc_info.value).startswith("Usage query failed:") diff --git a/tests/test_gateway_proxy.py b/tests/test_gateway_proxy.py new file mode 100644 index 0000000..68c8eeb --- /dev/null +++ b/tests/test_gateway_proxy.py @@ -0,0 +1,120 @@ +"""Tests for the relayed refresh proxy header handling.""" + +from __future__ import annotations + +import io +import socket +from email.message import Message + +from ucode import gateway_proxy + + +class _FakeHandler: + """Minimal stand-in exposing a `.headers` mapping like BaseHTTPRequestHandler.""" + + def __init__(self, headers: dict[str, str]) -> None: + self.headers = headers + + +class TestForwardedRequestHeaders: + def test_injects_swap_header_with_bearer(self): + handler = _FakeHandler({"Authorization": "Bearer anthropic-oauth"}) + out = gateway_proxy._forwarded_request_headers(handler, "dbx-token") + assert out["X-Databricks-AI-Gateway-Token"] == "Bearer dbx-token" + + def test_passes_authorization_through_untouched(self): + # The caller's Anthropic OAuth must survive verbatim — the proxy never + # reads or rewrites it. + handler = _FakeHandler({"Authorization": "Bearer anthropic-oauth"}) + out = gateway_proxy._forwarded_request_headers(handler, "dbx-token") + assert out["Authorization"] == "Bearer anthropic-oauth" + + def test_overwrites_client_supplied_swap_header(self): + # A stale settings.json value must not survive; the proxy replaces it. + handler = _FakeHandler({"X-Databricks-AI-Gateway-Token": "Bearer stale"}) + out = gateway_proxy._forwarded_request_headers(handler, "fresh") + assert out["X-Databricks-AI-Gateway-Token"] == "Bearer fresh" + + def test_strips_hop_by_hop_headers(self): + handler = _FakeHandler( + {"Host": "localhost:9", "Content-Length": "5", "Connection": "keep-alive"} + ) + out = gateway_proxy._forwarded_request_headers(handler, "t") + assert "Host" not in out + assert "Content-Length" not in out + assert "Connection" not in out + + +class _BrokenPipeWriter(io.RawIOBase): + """A wfile stand-in that raises BrokenPipeError on write, mimicking a client + (Claude Code) that closed the connection mid-response.""" + + def write(self, _data): # type: ignore[override] + raise BrokenPipeError(32, "Broken pipe") + + +class TestRelayResponseClientDisconnect: + def _handler(self, wfile) -> gateway_proxy._ProxyHandler: + # Bypass BaseHTTPRequestHandler.__init__ (which would service a socket); + # we only exercise _relay_response's write path. Set the few attributes the + # send_response/send_header machinery reads (normally populated by __init__). + handler = object.__new__(gateway_proxy._ProxyHandler) + handler.wfile = wfile + handler.request_version = "HTTP/1.1" + handler.requestline = "POST /v1/messages HTTP/1.1" + handler.command = "POST" + handler._headers_buffer = [] + return handler + + def test_relay_swallows_broken_pipe_on_headers(self): + # Client gone before headers flush: end_headers write raises BrokenPipe. + handler = self._handler(_BrokenPipeWriter()) + stream = io.BytesIO(b'{"ok":true}') + # Must not raise — a dead client is a routine teardown, not an error. + handler._relay_response(200, Message(), stream) + + def test_relay_swallows_connection_reset_mid_stream(self): + # Headers flush ok, then the client resets while streaming body chunks. + writes: list[bytes] = [] + + class _ResetAfterHeaders(io.RawIOBase): + def write(self, data): # type: ignore[override] + writes.append(bytes(data)) + if b"chunk" in bytes(data): + raise ConnectionResetError(54, "Connection reset by peer") + return len(data) + + def flush(self): + return None + + handler = self._handler(_ResetAfterHeaders()) + handler._relay_response(200, Message(), io.BytesIO(b"chunk-of-sse-data")) + + +class TestStartProxyPortFallback: + def test_falls_back_to_free_port_when_cached_port_busy(self, monkeypatch): + # A stale proxy from a killed session can still hold the cached port; the + # bind must fall back to an OS-assigned free port rather than crash. + class _StubCache: + def run_refresher(self): + return None + + monkeypatch.setattr(gateway_proxy, "_TokenCache", lambda workspace, profile: _StubCache()) + # Occupy a port to simulate the leftover proxy holding it. + occupied = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + occupied.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + occupied.bind(("127.0.0.1", 0)) + occupied.listen(1) + busy_port = occupied.getsockname()[1] + try: + server, _cache = gateway_proxy.start_proxy( + "https://x.staging.cloud.databricks.com", None, busy_port + ) + try: + bound = server.server_address[1] + assert bound != busy_port # fell back to a different, free port + assert bound != 0 + finally: + server.server_close() + finally: + occupied.close() From 55cd67c22cae3215efb8830bbbacb3d1549df6f8 Mon Sep 17 00:00:00 2001 From: Anjali Sujithan Date: Mon, 27 Jul 2026 02:44:35 +0000 Subject: [PATCH 2/2] fix e2e tests to mock interactive prompts --- src/ucode/agents/claude.py | 8 +++++++- src/ucode/cli.py | 17 ++++++++++++++--- src/ucode/gateway_proxy.py | 14 ++++++++++++-- tests/test_agents_init.py | 20 +++++++++++++++----- tests/test_cli.py | 16 ++++++++++++++-- tests/test_e2e.py | 8 ++++++++ 6 files changed, 70 insertions(+), 13 deletions(-) diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index 0190578..194210a 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -716,7 +716,13 @@ def _build_claude_argv(binary: str, tool_args: list[str], relayed: bool = False) # ucode wins over the caller for conflicting keys (protects gateway auth); # hooks from both sides survive. merged = _merge_claude_settings(caller_settings, read_json_safe(CLAUDE_SETTINGS_PATH)) - return [binary, *source_args, "--settings", json.dumps(merged, separators=(",", ":")), *remaining] + return [ + binary, + *source_args, + "--settings", + json.dumps(merged, separators=(",", ":")), + *remaining, + ] def _has_subscription_login() -> bool: diff --git a/src/ucode/cli.py b/src/ucode/cli.py index ae45971..eec5590 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -969,7 +969,14 @@ def _launch_tool( resolved_model = None else: state, resolved_model = resolve_launch_model(tool, state, None) - state = configure_tool(tool, state, resolved_model, provider=provider, provider_models=provider_models, relayed=relayed) + state = configure_tool( + tool, + state, + resolved_model, + provider=provider, + provider_models=provider_models, + relayed=relayed, + ) print_section(f"ucode with {TOOL_SPECS[tool]['display']}") if provider: print_kv("Provider", provider) @@ -1031,7 +1038,9 @@ def codex_cmd( workspace: WorkspaceOption = None, ) -> None: """Launch Codex via Databricks.""" - _launch_tool("codex", ctx, provider=provider, skip_preflight=skip_preflight, workspace=workspace) + _launch_tool( + "codex", ctx, provider=provider, skip_preflight=skip_preflight, workspace=workspace + ) @app.command("claude", context_settings={"allow_extra_args": True, "ignore_unknown_options": True}) @@ -1050,7 +1059,9 @@ def claude_cmd( workspace: WorkspaceOption = None, ) -> None: """Launch Claude Code via Databricks.""" - _launch_tool("claude", ctx, provider=provider, skip_preflight=skip_preflight, workspace=workspace) + _launch_tool( + "claude", ctx, provider=provider, skip_preflight=skip_preflight, workspace=workspace + ) @app.command("gemini", context_settings={"allow_extra_args": True, "ignore_unknown_options": True}) diff --git a/src/ucode/gateway_proxy.py b/src/ucode/gateway_proxy.py index 29bedff..07c2e13 100644 --- a/src/ucode/gateway_proxy.py +++ b/src/ucode/gateway_proxy.py @@ -33,8 +33,18 @@ # Hop-by-hop headers must not be forwarded across the proxy. _HOP_BY_HOP = frozenset( h.lower() - for h in ("connection", "keep-alive", "proxy-authenticate", "proxy-authorization", - "te", "trailers", "transfer-encoding", "upgrade", "host", "content-length") + for h in ( + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailers", + "transfer-encoding", + "upgrade", + "host", + "content-length", + ) ) _STREAM_CHUNK = 8192 diff --git a/tests/test_agents_init.py b/tests/test_agents_init.py index 9b0a600..97acf30 100644 --- a/tests/test_agents_init.py +++ b/tests/test_agents_init.py @@ -282,14 +282,20 @@ def test_none_provider_returns_none(self): def test_anthropic_returns_no_models(self, monkeypatch): self._patch(monkeypatch, {"provider_type": "anthropic", "targets": []}, None) - models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.a.svc") + models, error, relayed = agents_mod.resolve_provider_models( + "claude", self._STATE, "main.a.svc" + ) assert error is None assert models is None assert relayed is False def test_relayed_anthropic_flagged(self, monkeypatch): - self._patch(monkeypatch, {"provider_type": "anthropic", "targets": [], "relayed": True}, None) - models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.a.relayed") + self._patch( + monkeypatch, {"provider_type": "anthropic", "targets": [], "relayed": True}, None + ) + models, error, relayed = agents_mod.resolve_provider_models( + "claude", self._STATE, "main.a.relayed" + ) assert error is None assert models is None assert relayed is True @@ -300,7 +306,9 @@ def test_bedrock_returns_pinned_models(self, monkeypatch): "targets": ["us.anthropic.claude-sonnet-4-6", "global.anthropic.claude-opus-4-8"], } self._patch(monkeypatch, service, None) - models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.b.svc") + models, error, relayed = agents_mod.resolve_provider_models( + "claude", self._STATE, "main.b.svc" + ) assert error is None assert relayed is False assert models == { @@ -310,7 +318,9 @@ def test_bedrock_returns_pinned_models(self, monkeypatch): def test_invalid_provider_returns_error(self, monkeypatch): self._patch(monkeypatch, None, "boom") - models, error, relayed = agents_mod.resolve_provider_models("claude", self._STATE, "main.x.svc") + models, error, relayed = agents_mod.resolve_provider_models( + "claude", self._STATE, "main.x.svc" + ) assert models is None assert error == "boom" assert relayed is False diff --git a/tests/test_cli.py b/tests/test_cli.py index 9f52ed3..58bbc18 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -157,7 +157,13 @@ def test_workspace_flag_sets_current_workspace(self): """--workspace targets that workspace (normalized) before launch.""" patches = _patch_launch("claude") with ( - patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6], + patches[0], + patches[1], + patches[2], + patches[3], + patches[4], + patches[5], + patches[6], patch("ucode.cli.set_current_workspace") as mock_set, ): result = runner.invoke( @@ -171,7 +177,13 @@ def test_no_workspace_flag_leaves_current_workspace(self): """Without --workspace, launch never reassigns the current workspace.""" patches = _patch_launch("claude") with ( - patches[0], patches[1], patches[2], patches[3], patches[4], patches[5], patches[6], + patches[0], + patches[1], + patches[2], + patches[3], + patches[4], + patches[5], + patches[6], patch("ucode.cli.set_current_workspace") as mock_set, ): result = runner.invoke(app, ["claude"]) diff --git a/tests/test_e2e.py b/tests/test_e2e.py index 30ce287..025ced4 100644 --- a/tests/test_e2e.py +++ b/tests/test_e2e.py @@ -301,6 +301,10 @@ def test_only_picks_codex_writes_only_codex_config(self, tmp_path, monkeypatch, # selection plumbing, not the agent binaries themselves. monkeypatch.setattr(cli_mod, "install_tool_binary", lambda tool, **kwargs: True) monkeypatch.setattr(cli_mod, "validate_all_tools", lambda state: None) + # Answer the interactive prompts (provider picker + AI Tools opt-in) so no + # prompt reads stdin under capture; "databricks" keeps the Databricks path. + monkeypatch.setattr(cli_mod, "prompt_for_selection", lambda prompt, options: "databricks") + monkeypatch.setattr(cli_mod, "prompt_yes_no_default", lambda prompt, *, default: default) rc = cli_mod.configure_workspace_command() assert rc == 0 @@ -332,6 +336,10 @@ def test_rerun_with_different_pick_preserves_previous( ) monkeypatch.setattr(cli_mod, "install_tool_binary", lambda tool, **kwargs: True) monkeypatch.setattr(cli_mod, "validate_all_tools", lambda state: None) + # Answer the interactive prompts (provider picker + AI Tools opt-in) so no + # prompt reads stdin under capture; "databricks" keeps the Databricks path. + monkeypatch.setattr(cli_mod, "prompt_for_selection", lambda prompt, options: "databricks") + monkeypatch.setattr(cli_mod, "prompt_yes_no_default", lambda prompt, *, default: default) # First run: pick codex. monkeypatch.setattr(cli_mod, "prompt_for_tools", lambda available: ["codex"])