diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index ee7cc4e67..6d90d9233 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -10,6 +10,7 @@ import socket import subprocess import threading +import time from collections.abc import Callable from pathlib import Path @@ -17,6 +18,7 @@ from ucode.config_io import ( APP_DIR, ToolSpec, + atomic_write_json, backup_existing_file, deep_merge_dict, read_json_safe, @@ -34,6 +36,7 @@ CustomOAuthConfig, build_custom_auth_shell_command, custom_oauth_cli_enabled, + get_custom_client_token, ) from ucode.databricks import ( AnthropicModelCatalog, @@ -41,6 +44,7 @@ build_otel_headers_shell_command, build_otel_traces_endpoint, build_tool_base_url, + fetch_anthropic_gateway_models, get_databricks_token, ug_binary, ) @@ -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 @@ -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() @@ -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( diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index 40b191485..6ab279ed1 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -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. diff --git a/tests/README.md b/tests/README.md index c6be4ec3d..6027c99f4 100644 --- a/tests/README.md +++ b/tests/README.md @@ -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 | diff --git a/tests/integration/README.md b/tests/integration/README.md index a1b13314a..9727675ab 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -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 diff --git a/tests/integration/test_ug_claude_managed_model_discovery.py b/tests/integration/test_ug_claude_managed_model_discovery.py index fd00dfded..287f7ab93 100644 --- a/tests/integration/test_ug_claude_managed_model_discovery.py +++ b/tests/integration/test_ug_claude_managed_model_discovery.py @@ -8,6 +8,7 @@ import json import os import re +import time import pytest from utils.constants import MANAGED_CLAUDE_PROVIDER_SERVICE @@ -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( @@ -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): diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index 0b3fdb3b7..ee9085183 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -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()) @@ -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( @@ -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 diff --git a/tests/test_databricks.py b/tests/test_databricks.py index ba6c364fe..15a91f630 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -310,6 +310,56 @@ def test_multiple_locations_preserve_order(self): ) +class TestFetchAnthropicGatewayModels: + def test_preserves_metadata_and_scope_across_pages(self, monkeypatch): + first = {"id": "claude/first", "display_name": "First", "description": "Model"} + second = {"id": "claude-second"} + requests = [] + pages = iter( + [ + {"data": [first], "has_more": True, "last_id": first["id"]}, + {"data": [second], "has_more": False}, + ] + ) + + def fetch(url, token, **kwargs): + requests.append((url, token, kwargs)) + return next(pages), None + + monkeypatch.setattr(db_mod, "_http_get_json", fetch) + headers = {"Databricks-Model-Provider-Service": "main.default.mps"} + assert db_mod.fetch_anthropic_gateway_models(WS, "token", headers=headers) == ( + [first, second], + None, + ) + assert [request[0] for request in requests] == [ + f"{WS}/ai-gateway/anthropic/v1/models?limit=1000", + f"{WS}/ai-gateway/anthropic/v1/models?limit=1000&after_id=claude%2Ffirst", + ] + assert all( + request[1:] == ("token", {"headers": headers, "max_retries": 2}) for request in requests + ) + + @pytest.mark.parametrize( + "responses", + [ + [({"data": []}, None)], + [({"data": [{"id": ""}]}, None)], + [({"data": [{"id": "claude"}], "has_more": True, "last_id": "claude"}, None)] * 2, + [ + ({"data": [{"id": "claude"}], "has_more": True, "last_id": "claude"}, None), + (None, "HTTP 403"), + ], + ], + ) + def test_rejects_failed_or_invalid_catalog(self, monkeypatch, responses): + pages = iter(responses) + monkeypatch.setattr(db_mod, "_http_get_json", lambda *args, **kwargs: next(pages)) + models, reason = db_mod.fetch_anthropic_gateway_models(WS, "token", headers={}) + assert models is None + assert reason + + class TestDiscoverClaudeModels: def test_lists_all_anthropic_model_ids_without_legacy_validation(self, monkeypatch): captured = {}