Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 49 additions & 2 deletions src/ucode/agents/claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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",
Expand Down Expand Up @@ -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.

Expand All @@ -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]
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Comment thread
andy-xu-db marked this conversation as resolved.
return
if workspace:
os.environ["OAUTH_TOKEN"] = get_databricks_token(workspace, state.get("profile"))
exec_or_spawn(_build_claude_argv(binary, tool_args))
Expand Down
71 changes: 42 additions & 29 deletions src/ucode/gateway_proxy.py

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I tested with relayed auth and it works without any issues.

Comment thread
andy-xu-db marked this conversation as resolved.
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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()
Expand All @@ -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.
Expand Down Expand Up @@ -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:
Expand All @@ -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).
Expand Down Expand Up @@ -181,18 +182,22 @@ 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


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
Expand All @@ -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",
Expand All @@ -234,17 +239,17 @@ 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.
try:
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",
Expand Down Expand Up @@ -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.

Expand All @@ -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.
Expand All @@ -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)
Expand Down
95 changes: 77 additions & 18 deletions tests/test_agent_claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",),
]


Expand Down
Loading
Loading