diff --git a/src/ucode/agents/__init__.py b/src/ucode/agents/__init__.py index 9870feaf..86a2caf6 100644 --- a/src/ucode/agents/__init__.py +++ b/src/ucode/agents/__init__.py @@ -17,6 +17,7 @@ from ucode.config_io import ToolSpec from ucode.databricks import ( + BEDROCK_PROVIDER_TYPES, get_databricks_token, install_ai_tools, install_databricks_cli, @@ -286,9 +287,19 @@ def resolve_launch_model( def resolve_provider_models( tool: str, state: dict, provider: str | None ) -> tuple[dict | None, str | None, bool]: + """Return provider family defaults while preserving the existing API.""" + models, error, relayed, _targets = resolve_provider_models_and_targets(tool, state, provider) + return models, error, relayed + + +def resolve_provider_models_and_targets( + tool: str, state: dict, provider: str | None +) -> tuple[dict | None, str | None, bool, list[str]]: """Validate ``provider`` for ``tool`` and return the model ids to pin. - Returns ``(provider_models, error, relayed)``. ``provider_models`` is a ``{family: model_id}`` + Returns ``(provider_models, error, relayed, provider_targets)``. + ``provider_targets`` contains explicit picker targets; allow-all services return an empty list. + ``provider_models`` is a ``{family: model_id}`` dict re-derived from the service's declared targets — Bedrock (provider-side slugs), API-key Anthropic, and relayed Anthropic that declares a curated allowlist alike — so the client uses the ids the MPS allows rather than Claude Code's defaults. It is None when ``provider`` is None, when @@ -303,11 +314,11 @@ def resolve_provider_models( versions win rather than being re-derived here. """ if not provider: - return None, None, False + return None, None, False, [] token = get_databricks_token(state["workspace"], state.get("profile")) service, error = resolve_provider_service(tool, provider, state["workspace"], token) if error or service is None: - return None, error, False + return None, error, False, [] relayed = bool(service.get("relayed")) # Relayed services enforce their declared targets too, so map them like any Anthropic service # (allow_all declares none). relayed gates auth, not model reconciliation. @@ -315,8 +326,12 @@ def resolve_provider_models( # its target through resolve_gemini_provider_model instead — so mapping their targets # through Claude-family logic would be meaningless (see docstring). if tool != "claude": - return None, None, relayed - return map_claude_family_models(service.get("targets") or []) or None, None, relayed + return None, None, relayed, [] + targets = list(dict.fromkeys(service.get("targets") or [])) + if service.get("provider_type") in BEDROCK_PROVIDER_TYPES: + targets = [target for target in targets if "claude" in target.casefold()] + picker_targets = [] if service.get("allow_all_targets") else targets + return map_claude_family_models(targets) or None, None, relayed, picker_targets def resolve_gemini_provider_model( @@ -379,6 +394,7 @@ def configure_tool( custom_model: str | None = None, coding_agent_config_defaults: dict[str, str] | None = None, parent_schema: str | None = None, + provider_targets: list[str] | None = None, ) -> dict: result: dict | tuple[dict, str] if tool == "codex": @@ -400,6 +416,7 @@ def configure_tool( custom_model=custom_model, coding_agent_config_defaults=coding_agent_config_defaults, parent_schema=parent_schema, + provider_targets=provider_targets, ) else: # Every tool in this branch needs a model — including gemini under a provider, @@ -523,11 +540,19 @@ def _configure_one(tool: str, state: dict, provider: str | None) -> dict: if error: raise RuntimeError(error) return configure_tool(tool, state, model, provider=provider) - provider_models, error, relayed = resolve_provider_models(tool, state, provider) + provider_models, error, relayed, provider_targets = resolve_provider_models_and_targets( + tool, state, provider + ) if error: raise RuntimeError(error) return configure_tool( - tool, state, None, provider=provider, provider_models=provider_models, relayed=relayed + tool, + state, + None, + provider=provider, + provider_models=provider_models, + relayed=relayed, + provider_targets=provider_targets, ) if tool == "codex": return configure_tool("codex", state) diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index e4b519bc..d3d19810 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -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 @@ -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, ) @@ -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" @@ -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",) @@ -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]: @@ -350,6 +353,7 @@ def render_overlay( parent_schema: str | None = None, static_models: list[str] | None = None, otel_tracing: bool = False, + provider_targets: list[str] | None = None, ) -> tuple[dict, list[list[str]]]: """Return (overlay, managed_key_paths) for Claude settings.json. @@ -478,6 +482,14 @@ def render_overlay( "options": [{"model": m, "label": _picker_label(m)} for m in static_models], } keys += [[key] for key in CLAUDE_MANAGED_PICKER_KEYS] + elif provider and provider_targets and not relayed: + overlay["modelPicker"] = { + "replaceBuiltInOptions": True, + "options": [ + {"model": target, "label": _picker_label(target)} for target in provider_targets + ], + } + keys.append(["modelPicker"]) if otel_tracing: otel_env = _otel_trace_env(workspace) @@ -710,6 +722,7 @@ def write_tool_config( custom_model: str | None = None, coding_agent_config_defaults: dict[str, str] | None = None, parent_schema: str | None = None, + provider_targets: list[str] | None = None, ) -> dict: # Back up only a file that predates ucode's management of the tool. A # re-configure would otherwise snapshot ucode's own generated file, and @@ -737,10 +750,16 @@ def write_tool_config( parent_schema=parent_schema, static_models=state.get("claude_static_models"), otel_tracing=bool(state.get("claude_otel_tracing")), + provider_targets=provider_targets, ) + previous_keys = (state.get("managed_configs") or {}).get("claude", {}).get("keys", []) + stale_picker_keys = [ + key for key in CLAUDE_MANAGED_PICKER_KEYS if [key] in previous_keys and key not in overlay + ] managed_file_keys = list(managed_keys) for path in ( - [["env", key] for key in CLAUDE_MANAGED_MODEL_ENV_KEYS] + [[key] for key in stale_picker_keys] + + [["env", key] for key in CLAUDE_MANAGED_MODEL_ENV_KEYS] + [["env", key] for key in CLAUDE_CONDITIONAL_ENV_KEYS] + [["env", key] for key in CLAUDE_REMOVED_ENV_KEYS] + [["env", key] for key in CLAUDE_OTEL_TRACE_ENV_KEYS] @@ -785,6 +804,8 @@ def _compose(base: dict, *, enforce_model_default_hierarchy: bool) -> dict: else: target_env[key] = selected_default_model merged = deep_merge_dict(base, overlay_for_merge) + for key in stale_picker_keys: + merged.pop(key, None) overlay_custom_headers = overlay_for_merge["env"][ANTHROPIC_CUSTOM_HEADERS_ENV_KEY] merged["env"][ANTHROPIC_CUSTOM_HEADERS_ENV_KEY] = _merge_anthropic_custom_headers( existing_custom_headers, overlay_custom_headers @@ -914,8 +935,7 @@ def _reconcile_managed_settings( configuration mirrors ucode's settings there. The same compose operation that produced the private file is applied to the existing managed file, preserving unrelated IT-authored keys. - `ug configure` updates gateway-owned fields in this file, but does not generate or modify - the `modelPicker` object; an existing picker is retained by the merge. + `ug configure` updates gateway-owned fields and manages the picker for explicit model lists. Relayed launches are skipped: they depend on a per-session loopback refresh proxy that only runs during `ucode claude`, so a bare `claude` could not reach the gateway anyway. @@ -1167,7 +1187,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, @@ -1267,6 +1288,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], @@ -1275,12 +1313,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": @@ -1289,17 +1331,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")) @@ -1312,7 +1355,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]: diff --git a/src/ucode/cli.py b/src/ucode/cli.py index e21b9a9e..13a59f14 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -31,7 +31,7 @@ normalize_tool, resolve_gemini_provider_model, resolve_launch_model, - resolve_provider_models, + resolve_provider_models_and_targets, ) from ucode.agents import claude as claude_agent from ucode.agents import codex as codex_agent @@ -2261,6 +2261,7 @@ def _launch_tool( # Gemini is exempt: it validates the service and resolves its target in a single # lookup via resolve_gemini_provider_model (below), and uses no family model map. provider_models = None + provider_targets = None relayed = False coding_agent_config_defaults = ( managed_claude_family_models(managed) or {} @@ -2268,7 +2269,9 @@ def _launch_tool( else {} ) if provider and tool != "gemini": - provider_models, error, relayed = resolve_provider_models(tool, state, provider) + provider_models, error, relayed, provider_targets = resolve_provider_models_and_targets( + tool, state, provider + ) if error: if managed is not None and provider == managed_provider_service(managed, tool): # Clear error if the admin has Unity Catalog grants the developer doesn't. @@ -2347,6 +2350,7 @@ def _launch_tool( resolved_model, provider=provider, provider_models=provider_models, + provider_targets=provider_targets, relayed=relayed, route_root_model=route_root_model, # Claude's explicit model is launch-scoped and is passed through LaunchOptions below. diff --git a/tests/README.md b/tests/README.md index 50e67c2d..f3f9d5bc 100644 --- a/tests/README.md +++ b/tests/README.md @@ -18,6 +18,11 @@ the distribution rename with mocked installer calls, including failure recovery Agent configuration tests also verify `ug` auth/MCP helper commands, including quoted executable paths and replacement of legacy `ucode` routing/web-search helpers. +Claude MPS picker component tests cover full target lists, reuse of provider validation, +managed defaults, and removal of stale ug-owned picker settings. Non-relayed MPSes with +explicit targets replace built-in rows; `allow_all_targets` keeps native discovery. +These checks do not establish live `/model` coverage. + Agent-picker regression coverage in `test_ui.py` and `test_cli.py` drives actual keyboard selection: nothing is selected by default, selecting Codex installs only Codex, and submitting an empty selection installs nothing. Rendering checks cover diff --git a/tests/integration/README.md b/tests/integration/README.md index 9ddbeb0e..068ccde2 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -138,6 +138,10 @@ re-routes to gateway auth (`route=databricks`) — one relayed session reaching It needs a subscription OAuth token (see below). Interactive model-picker selection remains uncovered. +The non-relayed MPS replacement picker has component coverage in +`../test_agent_claude.py`, `../test_agents_init.py`, and `../test_cli.py`. +Live `/model` replacement is not covered by this suite. + MPS CUJs select the existing services already used by e2e: - Claude: `main.ucode.ci_e2e_anthropic_nonrelay_mps`. diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index c843c889..a06a5805 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -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") @@ -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") @@ -506,6 +517,37 @@ def test_static_models_skipped_when_relayed(self): assert "availableModels" not in overlay assert "modelPicker" not in overlay + def test_provider_targets_replace_picker_without_enforcing_allowlist(self): + targets = ["claude-sonnet-4-6", "claude-sonnet-5", "claude-fable-5-1"] + overlay, keys = claude.render_overlay( + WS, + None, + provider="main.x.mps", + provider_models={"sonnet": "claude-sonnet-5"}, + provider_targets=targets, + static_models=["system.ai.claude-opus-4-8"], + ) + + assert overlay["modelPicker"] == { + "replaceBuiltInOptions": True, + "options": [{"model": target, "label": target} for target in targets], + } + assert ["modelPicker"] in keys + assert "availableModels" not in overlay + assert "enforceAvailableModels" not in overlay + + def test_relayed_provider_targets_do_not_replace_picker(self): + overlay, _ = claude.render_overlay( + WS, + None, + provider="main.x.relay", + provider_targets=["claude-sonnet-5"], + relayed=True, + relayed_base_url="http://localhost:8000", + ) + + assert "modelPicker" not in overlay + def test_static_models_label_strips_system_ai_prefix(self): # Picker labels show the model id without the ``system.ai.`` prefix. static = ["system.ai.claude-opus-4-8", "databricks-custom-model"] @@ -775,14 +817,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 = [] @@ -958,7 +1005,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, @@ -969,10 +1021,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): @@ -1228,6 +1280,76 @@ def test_static_models_not_written_when_absent(self, monkeypatch): assert "availableModels" not in managed_content assert "modelPicker" not in managed_content + def test_provider_picker_updates_both_settings_files(self, monkeypatch): + private_writes: list = [] + managed_writes: list = [] + existing_picker = { + "replaceBuiltInOptions": False, + "options": [{"model": "system.ai.claude-opus-4-8"}], + } + existing = { + str(path): { + "modelPicker": existing_picker, + "availableModels": ["system.ai.claude-opus-4-8"], + "enforceAvailableModels": True, + "theme": "light", + } + for path in (claude.CLAUDE_SETTINGS_PATH, FAKE_MANAGED_PATH) + } + self._patch(monkeypatch, private_writes, managed_writes, existing) + state = { + "workspace": WS, + "codex_models": [], + "managed_configs": { + "claude": {"keys": [[key] for key in claude.CLAUDE_MANAGED_PICKER_KEYS]} + }, + } + targets = ["claude-sonnet-4-6", "claude-sonnet-5"] + + updated = claude.write_tool_config( + state, None, provider="main.x.mps", provider_targets=targets + ) + + for written in (private_writes[0][1], json.loads(managed_writes[0][1])): + assert written["modelPicker"] == { + "replaceBuiltInOptions": True, + "options": [{"model": target, "label": target} for target in targets], + } + assert "availableModels" not in written + assert "enforceAvailableModels" not in written + assert written["theme"] == "light" + assert ["modelPicker"] in updated["managed_configs"]["claude"]["keys"] + + @pytest.mark.parametrize("provider", [None, "main.x.allow_all"]) + @pytest.mark.parametrize("owned", [False, True]) + def test_removes_only_ug_owned_picker_without_provider_targets( + self, monkeypatch, provider, owned + ): + private_writes: list = [] + managed_writes: list = [] + picker = { + "replaceBuiltInOptions": True, + "options": [{"model": "claude-sonnet-4-6"}], + } + existing = { + str(path): {"modelPicker": picker} + for path in (claude.CLAUDE_SETTINGS_PATH, FAKE_MANAGED_PATH) + } + self._patch(monkeypatch, private_writes, managed_writes, existing) + state = { + "workspace": WS, + "codex_models": [], + "managed_configs": {"claude": {"keys": [["modelPicker"]] if owned else []}}, + } + + claude.write_tool_config(state, "system.ai.claude-sonnet-4-6", provider=provider) + + for written in (private_writes[0][1], json.loads(managed_writes[0][1])): + if owned: + assert "modelPicker" not in written + else: + assert written["modelPicker"] == picker + class TestAddClaudeMcpServer: def test_registers_stdio_proxy_command(self, monkeypatch): @@ -1462,20 +1584,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] = [] @@ -1631,27 +1760,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( { @@ -1663,9 +1817,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 diff --git a/tests/test_agents_init.py b/tests/test_agents_init.py index 8c85d429..15183d6c 100644 --- a/tests/test_agents_init.py +++ b/tests/test_agents_init.py @@ -418,12 +418,13 @@ def test_bedrock_ignores_gpt_targets(self, monkeypatch): } self._patch(monkeypatch, service, None) - models, error, relayed = agents_mod.resolve_provider_models( + models, error, relayed, picker_targets = agents_mod.resolve_provider_models_and_targets( "claude", self._STATE, "main.b.mixed" ) assert error is None assert models == {"opus": "global.anthropic.claude-opus-4-8"} + assert picker_targets == ["global.anthropic.claude-opus-4-8"] assert relayed is False def test_invalid_provider_returns_error(self, monkeypatch): @@ -435,6 +436,43 @@ def test_invalid_provider_returns_error(self, monkeypatch): assert error == "boom" assert relayed is False + @pytest.mark.parametrize("allow_all_targets", [False, True]) + def test_picker_targets_reuse_provider_lookup(self, monkeypatch, allow_all_targets): + from unittest.mock import Mock + + targets = ["claude-sonnet-4-6", "claude-sonnet-5", "claude-fable-5-1"] + service = { + "provider_type": "anthropic", + "targets": [*targets, targets[0]], + "allow_all_targets": allow_all_targets, + } + self._patch(monkeypatch, service, None) + lookup = Mock(return_value=(service, None)) + monkeypatch.setattr(agents_mod, "resolve_provider_service", lookup) + + models, error, relayed, picker_targets = agents_mod.resolve_provider_models_and_targets( + "claude", self._STATE, "main.a.svc" + ) + + assert error is None + assert relayed is False + assert models == {"sonnet": "claude-sonnet-5"} + assert picker_targets == ([] if allow_all_targets else targets) + lookup.assert_called_once_with("claude", "main.a.svc", self._STATE["workspace"], "token") + + def test_configure_passes_full_targets_to_claude(self, monkeypatch): + from unittest.mock import Mock + + targets = ["claude-sonnet-4-6", "claude-sonnet-5"] + self._patch(monkeypatch, {"provider_type": "anthropic", "targets": targets}, None) + writer = Mock(return_value=dict(self._STATE)) + monkeypatch.setattr(agents_mod.claude, "write_tool_config", writer) + + agents_mod._configure_one("claude", dict(self._STATE), "main.a.svc") + + assert writer.call_args.kwargs["provider_targets"] == targets + assert writer.call_args.kwargs["provider_models"] == {"sonnet": "claude-sonnet-5"} + @pytest.mark.parametrize("tool", ["gemini", "codex"]) def test_non_claude_pins_no_family_map(self, monkeypatch, tool): # Only claude pins a per-family map; codex ignores it and gemini resolves its own diff --git a/tests/test_cli.py b/tests/test_cli.py index 7ffbe797..55cfc5a7 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -985,7 +985,9 @@ def _provider_launch(monkeypatch, argv, provider_models, relayed=False): mock_launch = MagicMock() monkeypatch.setattr(cli_mod, "launch_agent", mock_launch) monkeypatch.setattr( - cli_mod, "resolve_provider_models", lambda t, s, p: (provider_models, None, relayed) + cli_mod, + "resolve_provider_models_and_targets", + lambda t, s, p: (provider_models, None, relayed, []), ) mock_configure = MagicMock(return_value=MINIMAL_STATE) monkeypatch.setattr(cli_mod, "configure_tool", mock_configure) @@ -1100,7 +1102,10 @@ def test_provider_sets_transient_claude_launch_marker(self): patch("ucode.cli.load_state", return_value=state), patch("ucode.cli.ensure_provider_state", return_value=state), patch("ucode.cli.configure_shared_state", return_value=state), - patch("ucode.cli.resolve_provider_models", return_value=(None, None, False)), + patch( + "ucode.cli.resolve_provider_models_and_targets", + return_value=(None, None, False, []), + ), patch("ucode.cli.configure_tool", return_value=state), patch("ucode.cli._fetch_managed_config", return_value=(None, False)), patch("ucode.cli.launch_agent") as mock_launch, @@ -1117,7 +1122,10 @@ def test_provider_sets_transient_codex_launch_marker(self): patch("ucode.cli.load_state", return_value=state), patch("ucode.cli.ensure_provider_state", return_value=state), patch("ucode.cli.configure_shared_state", return_value=state), - patch("ucode.cli.resolve_provider_models", return_value=(None, None, False)), + patch( + "ucode.cli.resolve_provider_models_and_targets", + return_value=(None, None, False, []), + ), patch("ucode.cli.configure_tool", return_value=state), patch("ucode.cli._fetch_managed_config", return_value=(None, False)), patch("ucode.cli.launch_agent") as mock_launch, @@ -1154,7 +1162,9 @@ def _launch(monkeypatch, resolve_provider_models): monkeypatch.setattr("ucode.cli.ensure_provider_state", lambda t: state) monkeypatch.setattr("ucode.cli.configure_shared_state", lambda *a, **k: state) monkeypatch.setattr("ucode.cli._fetch_managed_config", lambda s: (None, False)) - monkeypatch.setattr("ucode.cli.resolve_provider_models", resolve_provider_models) + monkeypatch.setattr( + "ucode.cli.resolve_provider_models_and_targets", resolve_provider_models + ) monkeypatch.setattr("ucode.cli.configure_tool", lambda *a, **k: state) monkeypatch.setattr( "ucode.cli.resolve_gemini_provider_model", @@ -4063,6 +4073,41 @@ class TestManagedModelDiscoveryLaunch: "enabled_agents": {"claude": {"model_config": {"model_services": ["system.ai.claude"]}}} } + @pytest.mark.parametrize("authored_default", [None, "claude-sonnet-4-6"]) + def test_managed_mps_keeps_all_picker_targets(self, monkeypatch, authored_default): + state = dict(MINIMAL_STATE) + managed = json.loads(json.dumps(self.MPS)) + if authored_default: + managed["enabled_agents"]["claude"]["model_config"]["models"] = { + "default_sonnet_model": authored_default + } + targets = ["claude-sonnet-4-6", "claude-sonnet-5", "claude-fable-5-1"] + monkeypatch.setattr(cli_mod, "ensure_bootstrap_dependencies", lambda *args: None) + monkeypatch.setattr(cli_mod, "load_state", lambda: state) + monkeypatch.setattr(cli_mod, "ensure_provider_state", lambda tool: state) + monkeypatch.setattr(cli_mod, "configure_shared_state", lambda *args, **kwargs: state) + monkeypatch.setattr(cli_mod, "_fetch_budget_recommendation", lambda *args: None) + monkeypatch.setattr(cli_mod, "_download_managed_skills", lambda *args: None) + with ( + patch("ucode.cli._fetch_managed_config", return_value=(managed, False)) as fetch, + patch("ucode.agents.get_databricks_token", return_value="token"), + patch( + "ucode.agents.resolve_provider_service", + return_value=({"provider_type": "anthropic", "targets": targets}, None), + ) as lookup, + patch("ucode.cli.configure_tool", return_value=state) as configure, + patch("ucode.cli.launch_agent"), + ): + result = runner.invoke(app, ["claude"]) + + assert result.exit_code == 0, result.output + fetch.assert_called_once() + lookup.assert_called_once() + assert configure.call_args.kwargs["provider_targets"] == targets + assert configure.call_args.kwargs["provider_models"] == { + "sonnet": authored_default or "claude-sonnet-5" + } + @pytest.mark.parametrize(("managed", "expected"), [(MPS, "1"), (STATIC, "0")]) def test_sets_literal_value_and_restores_prior(self, monkeypatch, managed, expected): monkeypatch.setenv("UG_ENABLE_MODEL_DISCOVERY", "prior-value")