diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index 5fad522..a32329e 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -29,6 +29,7 @@ build_tool_base_url, get_databricks_token, ) +from ucode.gateway_proxy import AUTHORIZATION_HEADER, start_proxy from ucode.launcher import exec_or_spawn from ucode.managed_files import OS, current_os, write_managed_file from ucode.smart_routing.claude_hooks import ( @@ -43,6 +44,7 @@ CLAUDE_CONFIG_DIR = Path.home() / ".claude" CLAUDE_SETTINGS_PATH = CLAUDE_CONFIG_DIR / "ucode-settings.json" CLAUDE_BACKUP_PATH = APP_DIR / "claude-ucode-settings.backup.json" +GATEWAY_MODEL_DISCOVERY_ENV_VAR = "ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY" SPEC: ToolSpec = { "binary": "claude", @@ -868,7 +870,12 @@ def _merge_claude_settings(base: dict, overlay: dict) -> dict: return merged -def _build_claude_argv(binary: str, tool_args: list[str], relayed: bool = False) -> list[str]: +def _build_claude_argv( + binary: str, + tool_args: list[str], + relayed: bool = False, + settings_override: dict | None = None, +) -> list[str]: """Build the ``claude`` argv, composing any caller ``--settings`` with ucode's managed settings. @@ -891,7 +898,7 @@ def _build_claude_argv(binary: str, tool_args: list[str], relayed: bool = False) """ source_args = ["--setting-sources", _RELAYED_SETTING_SOURCES] if relayed else [] caller_values, remaining = _extract_caller_settings(tool_args) - if not caller_values: + if not caller_values and settings_override is None: # No caller --settings: hand Claude ucode's settings file directly (the # common path; behavior unchanged). return [binary, *source_args, "--settings", str(CLAUDE_SETTINGS_PATH), *tool_args] @@ -901,6 +908,8 @@ 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)) + if settings_override is not None: + merged = _merge_claude_settings(merged, settings_override) return [ binary, *source_args, @@ -1010,12 +1019,50 @@ def _launch_relayed(state: dict, binary: str, tool_args: list[str]) -> None: raise SystemExit(returncode) +def _launch_gateway(state: dict, binary: str, tool_args: list[str]) -> None: + workspace = state["workspace"] + server, cache, client = start_proxy( + workspace, + state.get("profile"), + 0, + token_header=AUTHORIZATION_HEADER, + force_refresh_near_expiry=True, + ) + token = cache.token + os.environ["OAUTH_TOKEN"] = token + os.environ["ANTHROPIC_AUTH_TOKEN"] = token + os.environ["ANTHROPIC_BASE_URL"] = f"http://127.0.0.1:{server.server_address[1]}" + os.environ["CLAUDE_CODE_USE_GATEWAY"] = "1" + + server_thread = threading.Thread(target=server.serve_forever, daemon=True) + server_thread.start() + settings_override = { + "env": {"ANTHROPIC_BASE_URL": os.environ["ANTHROPIC_BASE_URL"]}, + } + proc = subprocess.Popen( + _build_claude_argv(binary, tool_args, settings_override=settings_override) + ) + try: + returncode = proc.wait() + except KeyboardInterrupt: + proc.send_signal(signal.SIGINT) + returncode = proc.wait() + finally: + cache.stop() + server.shutdown() + client.close() + 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 and os.environ.get(GATEWAY_MODEL_DISCOVERY_ENV_VAR) == "1": + _launch_gateway(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)) diff --git a/src/ucode/gateway_proxy.py b/src/ucode/gateway_proxy.py index d11b5f1..aef1a67 100644 --- a/src/ucode/gateway_proxy.py +++ b/src/ucode/gateway_proxy.py @@ -1,18 +1,17 @@ -"""Loopback refresh proxy for relayed Anthropic (Claude Max/Team/Enterprise). +"""Loopback refresh proxy for Claude gateway requests. 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. +header. Native gateway discovery instead carries the Databricks credential in +`Authorization`. The proxy refreshes the applicable header and streams responses +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. + passed through untouched in relayed mode and never logged. """ from __future__ import annotations @@ -34,6 +33,7 @@ # 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" +AUTHORIZATION_HEADER = "Authorization" # Hop-by-hop headers must not be forwarded across the proxy. _HOP_BY_HOP = frozenset( h.lower() @@ -50,9 +50,6 @@ "content-length", ) ) -# Request headers the proxy manages itself and must never forward on: hop-by-hop -# plus the swap header (replaced with a freshly-minted value per request). -_STRIP_ON_FORWARD = _HOP_BY_HOP | {_SWAP_HEADER.lower()} # Per-operation upstream timeouts. `read` is generous because model turns stream # over a single response and Anthropic emits SSE pings, so inter-chunk gaps stay # small; `connect`/`pool` fail fast when the gateway is unreachable. @@ -117,23 +114,27 @@ class _TokenCache: boundary triggers exactly one CLI call, not a thundering herd on the shared token cache.""" - def __init__(self, workspace: str, profile: str | None) -> None: + def __init__( + self, + workspace: str, + profile: str | None, + *, + force_refresh_near_expiry: bool = False, + ) -> None: self._workspace = workspace self._profile = profile + self._force_refresh_near_expiry = force_refresh_near_expiry self._state_lock = threading.Lock() # guards _token / _expiry (brief) self._refresh_lock = threading.Lock() # single-flights the CLI refresh self._stop = threading.Event() self._token = "" self._expiry = 0.0 - # Force on start so we begin on a full-TTL token rather than inheriting a - # near-expiry one cached from an earlier CLI call. Raises if auth is dead - # (surfaced by the caller at launch, before Claude Code starts). - self._refresh(force=True) + # Preserve the existing non-forced relayed-auth fetch. Gateway discovery + # opts into a forced fetch so its static client token starts with a full TTL. + self._refresh(force=force_refresh_near_expiry) def _refresh(self, *, force: bool) -> None: - """Mint a token and record its expiry. Caller holds `_refresh_lock` (or is - __init__). Non-force lets a token another process just refreshed satisfy - this call from the shared cache with no write — shrinking lock contention.""" + """Mint a token and record its expiry.""" token = get_databricks_token(self._workspace, self._profile, force_refresh=force) expiry = _jwt_exp(token) or (time.time() + _DEFAULT_TTL_S) with self._state_lock: @@ -151,7 +152,7 @@ def _ensure_fresh(self) -> None: if self._fresh_enough(): # another thread refreshed while we waited return try: - self._refresh(force=False) + self._refresh(force=self._force_refresh_near_expiry) except RuntimeError as exc: # Keep serving the current token; a request that then 401s triggers # a forced refresh + retry (see _ProxyHandler._handle). @@ -181,11 +182,14 @@ def stop(self) -> None: self._stop.set() -def _forwarded_request_headers(handler: BaseHTTPRequestHandler, token: str) -> dict[str, str]: +def _forwarded_request_headers( + handler: BaseHTTPRequestHandler, token: str, token_header: str = _SWAP_HEADER +) -> dict[str, str]: + strip_on_forward = _HOP_BY_HOP | {token_header.lower()} headers = { - key: value for key, value in handler.headers.items() if key.lower() not in _STRIP_ON_FORWARD + key: value for key, value in handler.headers.items() if key.lower() not in strip_on_forward } - headers[_SWAP_HEADER] = f"Bearer {token}" + headers[token_header] = f"Bearer {token}" return headers @@ -193,6 +197,7 @@ class _ProxyHandler(BaseHTTPRequestHandler): # Set by the server factory. cache: _TokenCache client: httpx.Client + token_header = _SWAP_HEADER def log_message(self, format: str, *args: object) -> None: return @@ -219,7 +224,7 @@ def _handle(self) -> None: ) try: # First attempt with the current token. - headers = _forwarded_request_headers(self, self.cache.token) + headers = _forwarded_request_headers(self, self.cache.token, self.token_header) with self.client.stream(self.command, url, headers=headers, content=body) as resp: _diagnostic_log( "upstream_headers", @@ -234,9 +239,9 @@ def _handle(self) -> None: # Auth rejected. Drain the (small) error body so the pooled # connection can be reused, then fall through to one retry. resp.read() - # A 401/403 may be a stale Databricks swap token rather than a bad - # Anthropic OAuth — the two are indistinguishable from the status - # alone. Force-refresh the swap token and retry once. If it was the + # A relayed 401/403 may be a stale Databricks swap token rather than a + # bad Anthropic OAuth — the two are indistinguishable from the status + # alone. Force-refresh the Databricks token and retry once. If it was the # Anthropic layer, the retry still 401s and we relay it verbatim, so a # genuine re-auth is triggered; a stale-Databricks 401 self-heals here # instead of surfacing to Claude Code as a spurious Anthropic prompt. @@ -244,7 +249,7 @@ def _handle(self) -> None: self.cache.refresh() except RuntimeError: pass # refresh failed; retry with the existing token and relay whatever comes - headers = _forwarded_request_headers(self, self.cache.token) + headers = _forwarded_request_headers(self, self.cache.token, self.token_header) with self.client.stream(self.command, url, headers=headers, content=body) as resp: _diagnostic_log( "upstream_headers", @@ -357,7 +362,11 @@ def __getattr__(self, name: str): def start_proxy( - workspace: str, profile: str | None, port: int + workspace: str, + profile: str | None, + port: int, + token_header: str = _SWAP_HEADER, + force_refresh_near_expiry: bool = False, ) -> tuple[ThreadingHTTPServer, _TokenCache, httpx.Client]: """Start the loopback refresh proxy + its background token refresher. @@ -370,7 +379,11 @@ def start_proxy( thread) and calls shutdown()/cache.stop()/client.close() on exit. """ upstream_base = f"{workspace.rstrip('/')}/ai-gateway/anthropic/" - cache = _TokenCache(workspace, profile) + cache = _TokenCache( + workspace, + profile, + force_refresh_near_expiry=force_refresh_near_expiry, + ) # One pooled, keep-alive client shared across handler threads: reuses TCP+TLS # to the gateway instead of a fresh handshake per request. Don't follow # redirects — a proxy relays 3xx verbatim. @@ -379,7 +392,7 @@ def start_proxy( handler = type( "BoundProxyHandler", (_ProxyHandler,), - {"cache": cache, "client": client}, + {"cache": cache, "client": client, "token_header": token_header}, ) try: server = ThreadingHTTPServer(("127.0.0.1", port), handler) diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index b45b630..92ab663 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -634,30 +634,89 @@ def boom(name, entry, scope=mcp_mod.MCP_USER_SCOPE): class TestClaudeLaunch: - def test_sets_oauth_token_before_exec(self, monkeypatch): - exec_calls: list[tuple[str, list[str]]] = [] + def test_default_launch_keeps_existing_auth_path(self, monkeypatch): + calls: list[list[str]] = [] + monkeypatch.delenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False) + monkeypatch.delenv("OAUTH_TOKEN", raising=False) + monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") + monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) + + claude.launch({"workspace": WS, "profile": "test"}, ["--debug"]) + + assert os.environ["OAUTH_TOKEN"] == "token" + assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]] + + def test_runs_through_refresh_proxy(self, monkeypatch): + calls: list[tuple] = [] + + class Server: + server_address = ("127.0.0.1", 12345) + + def serve_forever(self): + calls.append(("serve",)) + + def shutdown(self): + calls.append(("shutdown",)) + + class Cache: + token = "fresh-token" + + def stop(self): + calls.append(("stop",)) - def fake_execvp(binary: str, args: list[str]) -> None: - exec_calls.append((binary, args)) - raise RuntimeError("stop") + class Client: + def close(self): + calls.append(("close",)) + class Process: + def __init__(self, argv): + calls.append(("popen", argv)) + + def wait(self): + return 0 + + def start_proxy(workspace, profile, port, token_header, force_refresh_near_expiry): + calls.append( + ( + "proxy", + workspace, + profile, + port, + token_header, + force_refresh_near_expiry, + ) + ) + return Server(), Cache(), Client() + + monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1") monkeypatch.delenv("OAUTH_TOKEN", raising=False) - monkeypatch.setattr( - claude, "get_databricks_token", lambda workspace, profile=None: "fresh-token" - ) - monkeypatch.setattr(os, "execvp", fake_execvp) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + monkeypatch.delenv("CLAUDE_CODE_USE_GATEWAY", raising=False) + monkeypatch.setattr(claude, "start_proxy", start_proxy) + monkeypatch.setattr(claude.subprocess, "Popen", Process) - try: - claude.launch({"workspace": WS}, ["--debug"]) - except RuntimeError as exc: - assert str(exc) == "stop" + with pytest.raises(SystemExit) as exc: + claude.launch({"workspace": WS, "profile": "test"}, ["--debug"]) + assert exc.value.code == 0 assert os.environ["OAUTH_TOKEN"] == "fresh-token" - assert exec_calls == [ - ( - "claude", - ["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"], - ) + assert os.environ["ANTHROPIC_AUTH_TOKEN"] == "fresh-token" + assert os.environ["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:12345" + assert os.environ["CLAUDE_CODE_USE_GATEWAY"] == "1" + assert calls[:2] == [ + ("proxy", WS, "test", 0, claude.AUTHORIZATION_HEADER, True), + ("serve",), + ] + assert calls[2][0] == "popen" + argv = calls[2][1] + assert argv[:2] == ["claude", "--settings"] + assert json.loads(argv[2])["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:12345" + assert argv[3:] == ["--debug"] + assert calls[3:] == [ + ("stop",), + ("shutdown",), + ("close",), ] diff --git a/tests/test_gateway_proxy.py b/tests/test_gateway_proxy.py index f103aca..a985174 100644 --- a/tests/test_gateway_proxy.py +++ b/tests/test_gateway_proxy.py @@ -47,6 +47,13 @@ def test_overwrites_client_supplied_swap_header(self): out = gateway_proxy._forwarded_request_headers(handler, "fresh") assert out["X-Databricks-AI-Gateway-Token"] == "Bearer fresh" + def test_overwrites_authorization_header(self): + handler = _FakeHandler({"Authorization": "Bearer stale"}) + out = gateway_proxy._forwarded_request_headers( + handler, "fresh", gateway_proxy.AUTHORIZATION_HEADER + ) + assert out["Authorization"] == "Bearer fresh" + def test_strips_hop_by_hop_headers(self): handler = _FakeHandler( {"Host": "localhost:9", "Content-Length": "5", "Connection": "keep-alive"} @@ -228,40 +235,44 @@ def fake(_ws, _profile, force_refresh=False): class TestTokenCache: - def test_initial_mint_is_forced(self, monkeypatch): + def test_initial_mint_preserves_default_nonforce_refresh(self, monkeypatch): state = _install_fake_token(monkeypatch, [5000]) gateway_proxy._TokenCache("ws", None) - assert state["forces"] == [True] # full-TTL start + assert state["forces"] == [False] def test_fresh_token_is_not_refreshed(self, monkeypatch): state = _install_fake_token(monkeypatch, [5000]) cache = gateway_proxy._TokenCache("ws", None) _ = cache.token _ = cache.token - assert state["forces"] == [True] # no extra mint while fresh + assert state["forces"] == [False] # no extra mint while fresh - def test_near_expiry_triggers_nonforce_refresh(self, monkeypatch): - # First mint expires within the buffer -> reading .token refreshes once, - # non-force (so a token another process just wrote can satisfy it). + def test_near_expiry_preserves_default_nonforce_refresh(self, monkeypatch): state = _install_fake_token(monkeypatch, [100, 5000]) cache = gateway_proxy._TokenCache("ws", None) _ = cache.token - assert state["forces"] == [True, False] + assert state["forces"] == [False, False] _ = cache.token # now fresh again - assert state["forces"] == [True, False] + assert state["forces"] == [False, False] + + def test_near_expiry_can_force_refresh(self, monkeypatch): + state = _install_fake_token(monkeypatch, [100, 5000]) + cache = gateway_proxy._TokenCache("ws", None, force_refresh_near_expiry=True) + _ = cache.token + assert state["forces"] == [True, True] def test_refresh_is_single_flighted(self, monkeypatch): # A burst of concurrent requests at the expiry boundary must trigger ONE # refresh, not a thundering herd on the shared token cache. state = _install_fake_token(monkeypatch, [100, 5000], delay=0.05) - cache = gateway_proxy._TokenCache("ws", None) + cache = gateway_proxy._TokenCache("ws", None, force_refresh_near_expiry=True) threads = [threading.Thread(target=lambda: cache.token) for _ in range(10)] for t in threads: t.start() for t in threads: t.join() - # 1 forced init + exactly 1 non-force refresh shared by all 10 readers. - assert state["forces"] == [True, False] + # 1 forced init + exactly 1 forced refresh shared by all 10 readers. + assert state["forces"] == [True, True] def test_ensure_fresh_keeps_token_when_refresh_fails(self, monkeypatch): _install_fake_token(monkeypatch, [5000]) @@ -405,7 +416,11 @@ class _StubCache: def run_refresher(self): return None - monkeypatch.setattr(gateway_proxy, "_TokenCache", lambda workspace, profile: _StubCache()) + monkeypatch.setattr( + gateway_proxy, + "_TokenCache", + lambda workspace, profile, **_kwargs: _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)