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
79 changes: 53 additions & 26 deletions src/ucode/agents/claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@
import socket
import subprocess
import threading
from collections.abc import Callable
from collections.abc import Callable, Iterator
from contextlib import contextmanager
from pathlib import Path

from ucode import gateway_proxy
Expand All @@ -26,6 +27,7 @@
LOOPBACK_HOST,
MCP_CLEANUP_SCOPES,
MCP_USER_SCOPE,
MODEL_DISCOVERY_ENV_VAR,
MODEL_PROVIDER_SERVICE_HEADER,
MODEL_SERVICE_PARENT_SCHEMA_HEADER,
)
Expand Down Expand Up @@ -69,6 +71,7 @@
from .args import LaunchOptions, has_explicit_model_arg

GATEWAY_MODEL_DISCOVERY_ENV_VAR = "ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY"
CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
# If set, Claude Code launches in headless mode instead of the interactive login flow.
CLAUDE_CODE_OAUTH_TOKEN_ENV_VAR = "CLAUDE_CODE_OAUTH_TOKEN"
CLAUDE_CONFIG_DIR = Path.home() / ".claude"
Expand Down Expand Up @@ -189,9 +192,12 @@ def _otel_trace_env(workspace: str) -> dict[str, str]:
"sonnet": "ANTHROPIC_DEFAULT_SONNET_MODEL",
"haiku": "ANTHROPIC_DEFAULT_HAIKU_MODEL",
}
# Launch-scoped feature flags that ucode may write into Claude settings. These
# must be removed again when the corresponding launch flag is absent.
CLAUDE_CONDITIONAL_ENV_KEYS = ("CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY",)
# Launch-scoped feature flags must never remain in Claude settings.
CLAUDE_CONDITIONAL_ENV_KEYS = (
MODEL_DISCOVERY_ENV_VAR,
GATEWAY_MODEL_DISCOVERY_ENV_VAR,
CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR,
)
# Env keys ucode used to write but no longer does; stripped from the managed
# settings file on every launch so stale values never linger.
CLAUDE_REMOVED_ENV_KEYS = ("CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS",)
Expand Down Expand Up @@ -258,10 +264,7 @@ def managed_settings_are_current(state: dict) -> bool:
def gateway_model_discovery_setting_is_absent() -> bool:
"""Return whether model discovery is absent from persistent Claude settings."""
env = read_json_safe(CLAUDE_SETTINGS_PATH).get("env")
actual = (
env.get("CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY") if isinstance(env, dict) else None
)
return actual is None
return not isinstance(env, dict) or not any(key in env for key in CLAUDE_CONDITIONAL_ENV_KEYS)


def managed_settings_status(state: dict) -> tuple[Path | None, str, str]:
Expand Down Expand Up @@ -1167,7 +1170,8 @@ def _build_claude_argv(
merged = _merge_claude_settings(merged, settings_override)
merged_env = merged.get("env")
if isinstance(merged_env, dict):
merged_env.pop("CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY", None)
for key in CLAUDE_CONDITIONAL_ENV_KEYS:
merged_env.pop(key, None)
return [
binary,
*source_args,
Expand Down Expand Up @@ -1267,6 +1271,23 @@ def _launch_relayed(state: dict, binary: str, tool_args: list[str]) -> None:
raise SystemExit(returncode)


@contextmanager
def _native_model_discovery_environment(enabled: bool) -> Iterator[None]:
"""Set native discovery for one launch and restore the caller's exact value."""
existed = CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR in os.environ
previous = os.environ.get(CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR)
if enabled:
os.environ[CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR] = "1"
try:
yield
finally:
if existed:
assert previous is not None
os.environ[CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR] = previous
else:
os.environ.pop(CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR, None)


def launch(
state: dict,
tool_args: list[str],
Expand All @@ -1275,12 +1296,16 @@ def launch(
) -> None:
binary = SPEC["binary"]
workspace = state.get("workspace")
if workspace and os.environ.get(GATEWAY_MODEL_DISCOVERY_ENV_VAR) == "1":
# Discovery is launch-scoped. Pass it in the process environment rather
# than persisting it in Claude's private or OS-managed settings.
os.environ["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] = "1"
discovery_enabled = bool(
workspace
and (
os.environ.get(MODEL_DISCOVERY_ENV_VAR) == "1"
or os.environ.get(GATEWAY_MODEL_DISCOVERY_ENV_VAR) == "1"
)
)
if state.get("claude_relayed"):
_launch_relayed(state, binary, tool_args)
with _native_model_discovery_environment(discovery_enabled):
_launch_relayed(state, binary, tool_args)
return
# Smart routing needs Unix PTY support, which Windows does not provide.
if options.launch_smart_routing and os.name == "nt":
Expand All @@ -1289,17 +1314,18 @@ def launch(
"Please use Codex or disable smart routing."
)
if options.launch_smart_routing:
smart_routing_v2.launch_claude(
state,
tool_args,
binary=binary,
user_settings_path=CLAUDE_USER_SETTINGS_PATH,
# With no user pin, let Claude resolve its starting model from its own settings.
launch_model=options.user_pinned_model,
compose_settings=_compose_v2_settings,
launch_model_args=_launch_model_args,
model_name=_maybe_add_1m_suffix,
)
with _native_model_discovery_environment(discovery_enabled):
smart_routing_v2.launch_claude(
state,
tool_args,
binary=binary,
user_settings_path=CLAUDE_USER_SETTINGS_PATH,
# With no user pin, let Claude resolve its starting model from its own settings.
launch_model=options.user_pinned_model,
compose_settings=_compose_v2_settings,
launch_model_args=_launch_model_args,
model_name=_maybe_add_1m_suffix,
)
return
if workspace and not custom_oauth_cli_enabled(state.get("custom_oauth")):
os.environ["OAUTH_TOKEN"] = get_databricks_token(workspace, state.get("profile"))
Expand All @@ -1312,7 +1338,8 @@ def launch(
*_launch_model_args(tool_args, options.user_pinned_model),
*tool_args,
]
exec_or_spawn(_build_claude_argv(binary, launch_args, settings_override=settings_override))
with _native_model_discovery_environment(discovery_enabled):
exec_or_spawn(_build_claude_argv(binary, launch_args, settings_override=settings_override))


def validate_cmd(binary: str) -> list[str]:
Expand Down
143 changes: 124 additions & 19 deletions tests/test_agent_claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,8 +203,10 @@ def test_gateway_model_discovery_disabled_unless_opted_in(self, monkeypatch, env

def test_does_not_persist_gateway_model_discovery(self, monkeypatch):
monkeypatch.setenv("ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY", "1")
monkeypatch.setenv("UG_ENABLE_MODEL_DISCOVERY", "1")
monkeypatch.setenv("CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY", "1")
overlay, _ = claude.render_overlay(WS, "s4")
assert "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" not in overlay["env"]
assert not set(claude.CLAUDE_CONDITIONAL_ENV_KEYS) & overlay["env"].keys()

def test_smart_routing_does_not_persist_gateway_model_discovery(self, monkeypatch):
monkeypatch.setenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, "1")
Expand All @@ -227,6 +229,15 @@ def test_gateway_model_discovery_setting_detects_stale_opt_in(self, monkeypatch)

assert claude.gateway_model_discovery_setting_is_absent() is False

def test_gateway_model_discovery_setting_detects_stale_shared_opt_in(self, monkeypatch):
monkeypatch.setattr(
claude,
"read_json_safe",
lambda path: {"env": {"UG_ENABLE_MODEL_DISCOVERY": "1"}},
)

assert claude.gateway_model_discovery_setting_is_absent() is False

def test_sets_api_key_helper(self, monkeypatch):
monkeypatch.setattr("ucode.databricks.shutil.which", lambda command: f"/my tools/{command}")
overlay, _ = claude.render_overlay(WS, "s4")
Expand Down Expand Up @@ -775,14 +786,19 @@ def test_strips_stale_disable_experimental_betas(self, monkeypatch):
assert written[0]["env"]["CLAUDE_CODE_USE_GATEWAY"] == "1"

def test_strips_stale_gateway_model_discovery(self, monkeypatch):
existing = {"env": {"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "1"}}
existing = {
"env": {
"UG_ENABLE_MODEL_DISCOVERY": "1",
"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "1",
}
}
written: list = []
self._patch(monkeypatch, existing, written)
state = {"workspace": WS, "codex_models": []}

claude.write_tool_config(state, "databricks-claude-sonnet-4")

assert "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" not in written[0]["env"]
assert not set(claude.CLAUDE_CONDITIONAL_ENV_KEYS) & written[0]["env"].keys()

def test_writes_otel_tracing_when_enabled(self, monkeypatch):
written: list = []
Expand Down Expand Up @@ -958,7 +974,12 @@ def test_managed_file_updates_gateway_settings_without_changing_model_picker(sel
def test_managed_file_strips_stale_gateway_model_discovery(self, monkeypatch):
private_writes: list = []
managed_writes: list = []
stale = {"env": {"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "1"}}
stale = {
"env": {
"UG_ENABLE_MODEL_DISCOVERY": "1",
"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "1",
}
}
existing = {
str(claude.CLAUDE_SETTINGS_PATH): stale,
str(FAKE_MANAGED_PATH): stale,
Expand All @@ -969,10 +990,10 @@ def test_managed_file_strips_stale_gateway_model_discovery(self, monkeypatch):

claude.write_tool_config(state, "databricks-claude-sonnet-4")

assert "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" not in private_writes[0][1]["env"]
assert not set(claude.CLAUDE_CONDITIONAL_ENV_KEYS) & private_writes[0][1]["env"].keys()
assert (
"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
not in json.loads(managed_writes[0][1])["env"]
not set(claude.CLAUDE_CONDITIONAL_ENV_KEYS)
& json.loads(managed_writes[0][1])["env"].keys()
)

def test_managed_file_merges_anthropic_custom_headers(self, monkeypatch):
Expand Down Expand Up @@ -1462,20 +1483,27 @@ def boom(name, entry, scope=claude.MCP_USER_SCOPE):

class TestClaudeLaunch:
def test_gateway_discovery_enabled_for_relayed_provider(self, monkeypatch):
calls: list[tuple[dict, str, list[str]]] = []
calls: list[tuple[dict, str, list[str], str | None]] = []
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.delenv("CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY", raising=False)
monkeypatch.delenv(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False)
monkeypatch.setattr(
claude,
"_launch_relayed",
lambda state, binary, tool_args: calls.append((state, binary, tool_args)),
lambda state, binary, tool_args: calls.append(
(
state,
binary,
tool_args,
os.environ.get(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR),
)
),
)
state = {"workspace": WS, "claude_relayed": True}

claude.launch(state, ["--debug"], options=LaunchOptions())

assert os.environ["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1"
assert calls == [(state, "claude", ["--debug"])]
assert claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR not in os.environ
assert calls == [(state, "claude", ["--debug"], "1")]

def test_relayed_launch_uses_refresh_proxy(self, monkeypatch):
calls: list[tuple] = []
Expand Down Expand Up @@ -1631,27 +1659,52 @@ def test_v2_positional_prompt_uses_first_prompt_routing(self, monkeypatch, tool_
model_name=claude._maybe_add_1m_suffix,
)

def test_gateway_discovery_uses_direct_gateway(self, monkeypatch):
def test_legacy_gateway_discovery_uses_direct_gateway(self, monkeypatch):
calls: list[list[str]] = []
child_discovery: list[str | None] = []
monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False)
monkeypatch.setenv(claude.MODEL_DISCOVERY_ENV_VAR, "0")
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.delenv(claude.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))
monkeypatch.setattr(
claude,
"exec_or_spawn",
lambda argv: (
calls.append(argv),
child_discovery.append(
os.environ.get(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR)
),
),
)

claude.launch({"workspace": WS, "profile": "test"}, ["--debug"], options=LaunchOptions())

assert os.environ["OAUTH_TOKEN"] == "token"
assert os.environ["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1"
assert child_discovery == ["1"]
assert claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR not in os.environ
assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]]

def test_gateway_discovery_enabled_under_provider(self, monkeypatch):
def test_shared_gateway_discovery_enabled_under_provider(self, monkeypatch):
calls: list[list[str]] = []
child_discovery: list[str | None] = []
monkeypatch.delenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, raising=False)
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.setenv(claude.MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.delenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False)
monkeypatch.delenv(claude.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))
monkeypatch.setattr(
claude,
"exec_or_spawn",
lambda argv: (
calls.append(argv),
child_discovery.append(
os.environ.get(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR)
),
),
)

claude.launch(
{
Expand All @@ -1663,9 +1716,61 @@ def test_gateway_discovery_enabled_under_provider(self, monkeypatch):
options=LaunchOptions(),
)

assert os.environ["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1"
assert child_discovery == ["1"]
assert claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR not in os.environ
assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]]

def test_shared_gateway_discovery_zero_does_not_enable_native(self, monkeypatch):
child_discovery: list[str | None] = []
monkeypatch.setenv(claude.MODEL_DISCOVERY_ENV_VAR, "0")
monkeypatch.delenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False)
monkeypatch.delenv(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False)
monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token")
monkeypatch.setattr(
claude,
"exec_or_spawn",
lambda _argv: child_discovery.append(
os.environ.get(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR)
),
)

claude.launch({"workspace": WS}, [], options=LaunchOptions())

assert child_discovery == [None]
assert claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR not in os.environ

def test_shared_gateway_discovery_restores_exact_native_value(self, monkeypatch):
child_discovery: list[str | None] = []
monkeypatch.setenv(claude.MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.delenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, raising=False)
monkeypatch.setenv(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR, "caller-value")
monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token")
monkeypatch.setattr(
claude,
"exec_or_spawn",
lambda _argv: child_discovery.append(
os.environ.get(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR)
),
)

claude.launch({"workspace": WS}, [], options=LaunchOptions())

assert child_discovery == ["1"]
assert os.environ[claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR] == "caller-value"

def test_launch_error_restores_exact_native_value(self, monkeypatch):
monkeypatch.setenv(claude.MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.setenv(claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR, "caller-value")
monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token")
monkeypatch.setattr(
claude, "exec_or_spawn", Mock(side_effect=RuntimeError("launch failed"))
)

with pytest.raises(RuntimeError, match="launch failed"):
claude.launch({"workspace": WS}, [], options=LaunchOptions())

assert os.environ[claude.CLAUDE_GATEWAY_MODEL_DISCOVERY_ENV_VAR] == "caller-value"


class TestWriteToolConfigPrunesStaleModelEnv:
"""Stale ucode-managed model env keys (ANTHROPIC_MODEL, etc.) from earlier
Expand Down
Loading