Skip to content
Closed
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
49 changes: 44 additions & 5 deletions src/ucode/agents/claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,13 +10,15 @@
import socket
import subprocess
import threading
import time
from collections.abc import Callable
from pathlib import Path

from ucode import gateway_proxy
from ucode.config_io import (
APP_DIR,
ToolSpec,
atomic_write_json,
backup_existing_file,
deep_merge_dict,
read_json_safe,
Expand All @@ -34,13 +36,15 @@
CustomOAuthConfig,
build_custom_auth_shell_command,
custom_oauth_cli_enabled,
get_custom_client_token,
)
from ucode.databricks import (
AnthropicModelCatalog,
build_auth_shell_command,
build_otel_headers_shell_command,
build_otel_traces_endpoint,
build_tool_base_url,
fetch_anthropic_gateway_models,
get_databricks_token,
ug_binary,
)
Expand Down Expand Up @@ -1453,6 +1457,38 @@ def _rewrite_relayed_port(state: dict, port: int) -> None:
write_json_file(CLAUDE_SETTINGS_PATH, settings)


def _refresh_gateway_models_cache(state: dict) -> None:
if os.environ.get(GATEWAY_MODEL_DISCOVERY_ENV_VAR) != "1" or not state.get("workspace"):
return None
workspace = state["workspace"]
custom_oauth = state.get("custom_oauth")
token = (
get_custom_client_token(workspace, **custom_oauth)
if custom_oauth
else get_databricks_token(workspace, state.get("profile"))
)
env = read_json_safe(CLAUDE_SETTINGS_PATH).get("env", {})
headers = {}
for line in env.get("ANTHROPIC_CUSTOM_HEADERS", "").splitlines():
name, separator, value = line.partition(":")
if separator:
headers[name.strip()] = value.strip()
models, reason = fetch_anthropic_gateway_models(workspace, token, headers=headers)
if models is None:
raise RuntimeError(
f"Could not refresh Claude models: {reason}. Check your model source and retry."
)
config_dir = Path(os.environ.get("CLAUDE_CONFIG_DIR") or CLAUDE_CONFIG_DIR)
atomic_write_json(
config_dir / "cache" / "gateway-models.json",
{
"baseUrl": env.get("ANTHROPIC_BASE_URL") or build_tool_base_url("claude", workspace),
"fetchedAt": int(time.time() * 1000),
"models": models,
},
)


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
Expand All @@ -1479,12 +1515,14 @@ def token_provider(force_refresh: bool) -> str:
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()
_refresh_gateway_models_cache(state)
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()
Expand All @@ -1507,6 +1545,7 @@ def launch(
if state.get("claude_relayed"):
_launch_relayed(state, binary, tool_args)
return
_refresh_gateway_models_cache(state)
# Smart routing needs Unix PTY support, which Windows does not provide.
if options.launch_smart_routing and os.name == "nt":
raise RuntimeError(
Expand Down
48 changes: 42 additions & 6 deletions src/ucode/databricks.py
Original file line number Diff line number Diff line change
Expand Up @@ -2632,20 +2632,56 @@ def collect(result, schemas_done, schemas_total):


def _get_anthropic_models_json(
workspace: str, token: str, *, parent_schema: str | None = None
workspace: str,
token: str,
*,
parent_schema: str | None = None,
headers: dict[str, str] | None = None,
params: dict[str, str] | None = None,
) -> tuple[dict | list | None, str | None]:
hostname = workspace_hostname(workspace)
headers = (
{MODEL_SERVICE_PARENT_SCHEMA_HEADER: parent_schema} if parent_schema is not None else None
)
if parent_schema is not None:
headers = {**(headers or {}), MODEL_SERVICE_PARENT_SCHEMA_HEADER: parent_schema}
url = f"https://{workspace_hostname(workspace)}{ANTHROPIC_MODELS_PATH}"
if params:
url += f"?{urlencode(params)}"
return _http_get_json(
f"https://{hostname}{ANTHROPIC_MODELS_PATH}",
url,
token,
max_retries=_ANTHROPIC_MODEL_DISCOVERY_SETUP_MAX_RETRIES,
**({"headers": headers} if headers is not None else {}),
)


def fetch_anthropic_gateway_models(
workspace: str, token: str, *, headers: dict[str, str]
) -> tuple[list[dict] | None, str | None]:
"""Fetch every page of the launch-scoped Claude gateway catalog."""
models: list[dict] = []
cursors: set[str] = set()
params = {"limit": "1000"}
while True:
payload, reason = _get_anthropic_models_json(
workspace, token, headers=headers, params=params
)
if payload is None:
return None, reason
data = payload if isinstance(payload, dict) else {}
page = data.get("data")
if not isinstance(page, list) or any(
not isinstance(model, dict) or not isinstance(model.get("id"), str) or not model["id"]
for model in page
):
return None, "AI Gateway returned an invalid Anthropic model catalog"
models.extend(page)
if not data.get("has_more"):
return (models, None) if models else (None, "AI Gateway returned no Anthropic models")
cursor = data.get("last_id")
if not isinstance(cursor, str) or not cursor or cursor in cursors:
return None, "AI Gateway returned an invalid Anthropic pagination cursor"
cursors.add(cursor)
params["after_id"] = cursor


def list_anthropic_models(workspace: str, token: str) -> tuple[list[str], str | None]:
"""List every model id advertised by AI Gateway's Anthropic endpoint.

Expand Down
1 change: 1 addition & 0 deletions tests/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@ All tests live directly in `integration/`; shared mechanics live in `utils/`.
| `test_ug_configure_claude_rejects_invalid_credentials`, `test_ug_configure_codex_rejects_invalid_credentials` | Configure with a rejected bearer against the real workspace | Authentication failure; no successful saved setup |
| `test_ug_configure_managed_claude`, `test_ug_configure_managed_codex` | Configure against a workspace that publishes a managed CodingAgentConfig | No agent selector; each agent's generated config exposes exactly the admin's static model_services; real gateway prompt on launch |
| `test_case_01_*`, `test_case_03_*` | Launch managed Claude after configure and from fresh state, with personal discovery enabled and disabled | Claude receives the admin MPS header, caches native discovery results, and opens its real model picker |
| `test_case_01_managed_claude_uses_admin_discovery_after_configure` | Restart managed Claude in the same home | Cache timestamp advances past restart time despite an existing fresh cache; the real model picker opens again; no inference is tested |
| `test_case_05_*`, `test_case_07_*` | Pass a provider or model-location override to managed Claude after configure and from fresh state | ug rejects the override before Claude starts and preserves agent-owned state |
| `test_case_09_*`, `test_case_11_*` | Disable discovery and pass a provider or model-location override to managed Claude | ug still rejects both configured and fresh launches |
| `test_case_02_*`, `test_case_04_*` | Launch managed Codex after configure and from fresh state, with personal discovery enabled and disabled | Codex exposes exactly the admin MPS-scoped catalog |
Expand Down
3 changes: 3 additions & 0 deletions tests/integration/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -301,6 +301,9 @@ across all configured/fresh scenarios. The Codex module does the same with
model picker, Codex's exact app-server catalog, and both agents' rejection of personal source
overrides. Codex state comparisons exclude `.codex/tmp/arg0`, the disposable executable links
recreated by version checks, while continuing to compare persistent agent files.
The configured Claude Case 1 also restarts in the same home, requires the gateway cache timestamp
to advance past restart time, and opens the model picker again. This cache-refresh journey does
not exercise inference or require the provider's upstream inference credential to work.
In addition, `test_ug_configure_managed_codex_catalog_fallback` injects the intentionally nonexistent
`system.ai.gpt-99`, keeping it out of the real workspace while launching Codex through that
workspace on the valid default model `system.ai.gpt-5-6-sol`. With smart routing enabled, it opens
Expand Down
19 changes: 17 additions & 2 deletions tests/integration/test_ug_claude_managed_model_discovery.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import json
import os
import re
import time

import pytest
from utils.constants import MANAGED_CLAUDE_PROVIDER_SERVICE
Expand Down Expand Up @@ -88,9 +89,10 @@ def _assert_managed_provider_in_picker(session, workspace, screen):

@pytest.mark.tui
def test_case_01_managed_claude_uses_admin_discovery_after_configure(live_session, workspace):
"""Scenario: configure managed Claude, then launch its model picker.
"""Scenario: configure managed Claude, open its model picker, then restart.

Expected: the managed model catalog wins after configuration.
Expected: the managed model catalog wins, a fresh cache replaces the prior
session's cache on restart, and the model picker opens again. No inference is tested.
"""
session = live_session
result = session.run(
Expand All @@ -111,6 +113,19 @@ def test_case_01_managed_claude_uses_admin_discovery_after_configure(live_sessio

_assert_managed_provider_in_picker(session, workspace, screen)

cache_path = session.home / ".claude/cache/gateway-models.json"
previous = json.loads(cache_path.read_text())
restarted_at = time.time_ns() // 1_000_000
with AgentTerminal(session, "claude", command, "case-01-restarted") as tui:
tui.boot()
refreshed = json.loads(cache_path.read_text())
assert refreshed["fetchedAt"] >= restarted_at > previous["fetchedAt"]
assert refreshed["baseUrl"] == previous["baseUrl"]
assert refreshed["models"]
screen = tui.open_model_picker()
tui.exit_normally()
_assert_managed_provider_in_picker(session, workspace, screen)


@pytest.mark.tui
def test_case_01_fresh_managed_claude_uses_admin_discovery(live_session, workspace):
Expand Down
103 changes: 103 additions & 0 deletions tests/test_agent_claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -1900,6 +1900,7 @@ def test_gateway_discovery_uses_direct_gateway(self, monkeypatch):
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.delenv("OAUTH_TOKEN", raising=False)
monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token")
monkeypatch.setattr(claude, "_refresh_gateway_models_cache", Mock(return_value=None))
monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv))

claude.launch({"workspace": WS, "profile": "test"}, ["--debug"], options=LaunchOptions())
Expand All @@ -1914,6 +1915,7 @@ def test_gateway_discovery_enabled_under_provider(self, monkeypatch):
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.delenv("OAUTH_TOKEN", raising=False)
monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token")
monkeypatch.setattr(claude, "_refresh_gateway_models_cache", Mock(return_value=None))
monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv))

claude.launch(
Expand All @@ -1930,6 +1932,107 @@ def test_gateway_discovery_enabled_under_provider(self, monkeypatch):
assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]]


class TestGatewayModelsCache:
@pytest.fixture(autouse=True)
def setup(self, monkeypatch, tmp_path):
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1")
monkeypatch.setenv("CLAUDE_CONFIG_DIR", str(tmp_path / "custom-claude"))
monkeypatch.setattr(claude, "CLAUDE_SETTINGS_PATH", tmp_path / "settings.json")
monkeypatch.setattr(claude, "get_databricks_token", Mock(return_value="token"))
monkeypatch.setattr(claude, "exec_or_spawn", Mock())
self.cache_path = tmp_path / "custom-claude/cache/gateway-models.json"
self.models = [{"id": "claude-new", "display_name": "New Claude"}]
self.fetch = Mock(return_value=(self.models, None))
monkeypatch.setattr(claude, "fetch_anthropic_gateway_models", self.fetch)

@pytest.mark.parametrize(
"headers",
[
{},
{"Databricks-Model-Provider-Service": "main.default.mps"},
{"x-databricks-model-service-parent-schema": "main.models"},
],
)
def test_refreshes_before_each_launch_with_configured_scope(self, monkeypatch, headers):
claude.write_json_file(
claude.CLAUDE_SETTINGS_PATH,
{
"env": {
"ANTHROPIC_BASE_URL": f"{WS}/ai-gateway/anthropic",
"ANTHROPIC_CUSTOM_HEADERS": "\n".join(
f"{name}: {value}" for name, value in headers.items()
),
},
},
)
snapshots = []
monkeypatch.setattr(
claude,
"exec_or_spawn",
lambda argv: snapshots.append(json.loads(self.cache_path.read_text())),
)
state = {"workspace": WS, "profile": "test"}
claude.launch(state, [], options=LaunchOptions())
self.fetch.return_value = ([{"id": "claude-replacement"}], None)
claude.launch(state, [], options=LaunchOptions())
self.fetch.assert_called_with(WS, "token", headers=headers)
assert self.fetch.call_count == 2
assert os.environ["OAUTH_TOKEN"] == "token"
assert snapshots[0]["models"] == self.models
assert snapshots[1]["models"] == [{"id": "claude-replacement"}]
assert snapshots[1]["baseUrl"] == f"{WS}/ai-gateway/anthropic"
assert snapshots[1]["fetchedAt"] > 0

def test_discovery_disabled_leaves_cache_untouched(self, monkeypatch):
monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "0")
claude.write_json_file(self.cache_path, {"models": ["existing"]})
claude.launch({"workspace": WS}, [], options=LaunchOptions())
self.fetch.assert_not_called()
assert json.loads(self.cache_path.read_text()) == {"models": ["existing"]}

@pytest.mark.parametrize("failure", ["fetch", "write"])
def test_failure_blocks_launch_without_replacing_cache(self, monkeypatch, failure):
claude.write_json_file(self.cache_path, {"models": ["existing"]})
if failure == "fetch":
self.fetch.return_value = (None, "HTTP 403")
else:
monkeypatch.setattr(os, "replace", Mock(side_effect=OSError("disk full")))
with pytest.raises(RuntimeError):
claude.launch({"workspace": WS}, [], options=LaunchOptions())
claude.exec_or_spawn.assert_not_called()
assert json.loads(self.cache_path.read_text()) == {"models": ["existing"]}

@pytest.mark.parametrize("failed", [False, True])
def test_relay_refreshes_after_port_fallback_and_cleans_up(self, monkeypatch, failed):
server = Mock(server_address=("127.0.0.1", 54321))
cache, client = Mock(), Mock()
monkeypatch.setattr(claude, "_ensure_subscription_login", Mock())
monkeypatch.setattr(
claude.gateway_proxy, "start_relay_proxy", Mock(return_value=(server, cache, client))
)
claude.write_json_file(
claude.CLAUDE_SETTINGS_PATH, {"env": {"ANTHROPIC_BASE_URL": "http://127.0.0.1:12345"}}
)
process = Mock(return_value=Mock(wait=Mock(return_value=0)))
monkeypatch.setattr(claude.subprocess, "Popen", process)
if failed:
self.fetch.return_value = (None, "HTTP 403")
with pytest.raises(RuntimeError if failed else SystemExit):
claude.launch(
{"workspace": WS, "claude_relayed": True, "relayed_proxy_port": 12345},
[],
options=LaunchOptions(),
)
if failed:
process.assert_not_called()
else:
assert json.loads(self.cache_path.read_text())["baseUrl"] == "http://127.0.0.1:54321"
server.shutdown.assert_called_once()
cache.stop.assert_called_once()
client.close.assert_called_once()
self.fetch.assert_called_once_with(WS, "token", headers={})


class TestWriteToolConfigPrunesStaleModelEnv:
"""Stale ucode-managed model env keys (ANTHROPIC_MODEL, etc.) from earlier
ucode versions must be removed on every launch — otherwise they linger in
Expand Down
Loading
Loading