diff --git a/src/ucode/agents/__init__.py b/src/ucode/agents/__init__.py index b12d8d86..83fd9117 100644 --- a/src/ucode/agents/__init__.py +++ b/src/ucode/agents/__init__.py @@ -424,6 +424,7 @@ def configure_tool( route_root_model: str | None = None, custom_model: str | None = None, coding_agent_config_defaults: dict[str, str] | None = None, + parent_schema: str | None = None, ) -> dict: result: dict | tuple[dict, str] if tool == "codex": @@ -442,6 +443,7 @@ def configure_tool( route_root_model=route_root_model, custom_model=custom_model, coding_agent_config_defaults=coding_agent_config_defaults, + parent_schema=parent_schema, ) else: # Every tool in this branch needs a model — including gemini under a provider, diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index 21d41ab0..f44d49ae 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -24,7 +24,11 @@ read_json_safe, write_json_file, ) -from ucode.constants import LOOPBACK_HOST +from ucode.constants import ( + LOOPBACK_HOST, + MODEL_PROVIDER_SERVICE_HEADER, + MODEL_SERVICE_PARENT_SCHEMA_HEADER, +) from ucode.custom_oauth import CustomOAuthConfig, build_custom_auth_shell_command from ucode.databricks import ( build_auth_shell_command, @@ -169,7 +173,8 @@ def _resolve_web_search_model(state: dict) -> str | None: { "x-databricks-use-coding-agent-mode", "user-agent", - "databricks-model-provider-service", + MODEL_PROVIDER_SERVICE_HEADER.casefold(), + MODEL_SERVICE_PARENT_SCHEMA_HEADER.casefold(), } ) CLAUDE_TRACING_STOP_HOOK_SUFFIX = " autolog claude stop-hook" @@ -324,6 +329,7 @@ def render_overlay( relayed_base_url: str | None = None, route_root_model: str | None = None, custom_model: str | None = None, + parent_schema: str | None = None, ) -> tuple[dict, list[list[str]]]: """Return (overlay, managed_key_paths) for Claude settings.json. @@ -360,7 +366,9 @@ def render_overlay( f"User-Agent: ucode/{ucode_version()} claude/{agent_version('claude')}", ] if provider: - header_lines.append(f"Databricks-Model-Provider-Service: {provider}") + header_lines.append(f"{MODEL_PROVIDER_SERVICE_HEADER}: {provider}") + elif parent_schema: + header_lines.append(f"{MODEL_SERVICE_PARENT_SCHEMA_HEADER}: {parent_schema}") # Relayed: the X-Databricks-AI-Gateway-Token swap header is added per request # by the refresh proxy, not here — a static value would go stale mid-session. custom_headers = "\n".join(header_lines) @@ -575,6 +583,7 @@ def write_tool_config( route_root_model: str | None = None, custom_model: str | None = None, coding_agent_config_defaults: dict[str, str] | None = None, + parent_schema: str | None = None, ) -> dict: backup_existing_file(CLAUDE_SETTINGS_PATH, CLAUDE_BACKUP_PATH) web_search_model = _resolve_web_search_model(state) @@ -596,6 +605,7 @@ def write_tool_config( relayed_base_url=relayed_base_url, route_root_model=route_root_model, custom_model=custom_model, + parent_schema=parent_schema, ) tracing_env_vars = tracing_env(state, "claude") stop_hook_command = claude_tracing_stop_hook_command() if tracing_env_vars else None diff --git a/src/ucode/cli.py b/src/ucode/cli.py index 5466f2da..c76d6cec 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -124,6 +124,7 @@ set_current_workspace, set_provider_service, ) +from ucode.string_utils import is_valid_catalog_schema from ucode.tracing import configure_tracing_command from ucode.ui import ( console, @@ -2043,10 +2044,15 @@ def _launch_tool( managed: dict | None = None, recommendation: dict | None = None, model: str | None = None, + parent_schema: str | None = None, custom_oauth: CustomOAuthConfig | None = None, ) -> None: try: tool = normalize_tool(tool_name) + if provider is not None and parent_schema is not None: + raise RuntimeError("--provider and --parent cannot be used together.") + if parent_schema is not None and not is_valid_catalog_schema(parent_schema): + raise RuntimeError("--parent must be `.`.") explicit_prompt = _has_explicit_prompt(ctx) smart_routing_enabled = smart_routing_v2.enabled() # Launchers such as isaac put their harness arguments after `--`, so the harness's own @@ -2141,6 +2147,8 @@ def _launch_tool( ) if managed_provider: provider = managed_provider + if provider and parent_schema is not None: + raise RuntimeError("--provider and --parent cannot be used together.") # Checked after the managed config settles `provider`: an admin-set provider must trip this # guard too, or routing would be persisted as on while a provider is active. if tool in CAN_USE_CACHED_CONFIG_AGENTS and smart_routing_enabled and provider: @@ -2244,6 +2252,7 @@ def _launch_tool( # Claude's explicit model is launch-scoped and is passed through LaunchOptions below. custom_model=None, coding_agent_config_defaults=coding_agent_config_defaults, + parent_schema=parent_schema, ) # Relayed = a Claude subscription: forward the model to Claude Code's own flag, like `-- --model X`. should_forward_relayed_model = ( @@ -2585,6 +2594,13 @@ def claude_cmd( "before any `--` separator.", ), ] = None, + parent: Annotated[ + str | None, + typer.Option( + "--parent", + help="Discover model services in `.`. Example: main.default", + ), + ] = None, model: Annotated[ str | None, typer.Option( @@ -2656,7 +2672,7 @@ def claude_cmd( claude_agent.disable_smart_routing(load_state()) print_success("Claude Code smart routing disabled; ug routing hooks removed") return - if enable_model_discovery: + if enable_model_discovery or (parent is not None and provider is None): os.environ[claude_agent.GATEWAY_MODEL_DISCOVERY_ENV_VAR] = "1" with _smart_routing_v2_flag(enable_smart_routing_flag): with _disable_smart_routing_for_subcommand("claude", ctx): @@ -2668,6 +2684,7 @@ def claude_cmd( refresh=refresh, skip_preflight=skip_preflight, workspace_url=workspace, + parent_schema=parent, custom_oauth=custom_oauth, ) diff --git a/src/ucode/constants.py b/src/ucode/constants.py index f1664b17..8c08838a 100644 --- a/src/ucode/constants.py +++ b/src/ucode/constants.py @@ -2,3 +2,6 @@ LOCALHOST = "localhost" LOOPBACK_HOST = "127.0.0.1" + +MODEL_PROVIDER_SERVICE_HEADER = "Databricks-Model-Provider-Service" +MODEL_SERVICE_PARENT_SCHEMA_HEADER = "Databricks-Model-Service-Parent-Schema" diff --git a/src/ucode/string_utils.py b/src/ucode/string_utils.py new file mode 100644 index 00000000..7d518769 --- /dev/null +++ b/src/ucode/string_utils.py @@ -0,0 +1,14 @@ +"""Shared string validation helpers.""" + +from __future__ import annotations + + +def is_valid_catalog_schema(value: str) -> bool: + """Return whether value is a safe ``.`` reference.""" + parts = value.split(".") + return len(parts) == 2 and all( + part + and part.isprintable() + and not any(character.isspace() or character == "/" for character in part) + for part in parts + ) diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index 1ce87a8b..22219e87 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -372,6 +372,13 @@ def test_no_provider_header_without_flag(self): overlay, _ = claude.render_overlay(WS, "s4") assert "Databricks-Model-Provider-Service" not in overlay["env"]["ANTHROPIC_CUSTOM_HEADERS"] + def test_parent_adds_discovery_header(self): + overlay, _ = claude.render_overlay(WS, "s4", parent_schema="main.default") + assert ( + "Databricks-Model-Service-Parent-Schema: main.default" + in overlay["env"]["ANTHROPIC_CUSTOM_HEADERS"] + ) + def test_bedrock_provider_pins_model_ids(self): provider_models = { "opus": "global.anthropic.claude-opus-4-8", @@ -460,6 +467,15 @@ def test_headers_newline_delimited(self, monkeypatch): class TestMergeAnthropicCustomHeaders: + def test_removes_stale_parent_header(self): + existing = "X-User: keep\nDatabricks-Model-Service-Parent-Schema: main.default" + managed = "x-databricks-use-coding-agent-mode: true" + + merged = claude._merge_anthropic_custom_headers(existing, managed) + + assert "X-User: keep" in merged + assert "Databricks-Model-Service-Parent-Schema" not in merged + def test_merges_existing_settings_with_ucode_managed_headers(self): headers_from_existing_settings = "\n".join( [ diff --git a/tests/test_agents_init.py b/tests/test_agents_init.py index a59e40a1..2ee00493 100644 --- a/tests/test_agents_init.py +++ b/tests/test_agents_init.py @@ -406,6 +406,24 @@ def test_bedrock_returns_pinned_models(self, monkeypatch): "opus": "global.anthropic.claude-opus-4-8", } + def test_bedrock_ignores_gpt_targets(self, monkeypatch): + service = { + "provider_type": "amazon_bedrock", + "targets": [ + "global.anthropic.claude-opus-4-8", + "openai.gpt-oss-120b-1:0", + ], + } + self._patch(monkeypatch, service, None) + + models, error, relayed = agents_mod.resolve_provider_models( + "claude", self._STATE, "main.b.mixed" + ) + + assert error is None + assert models == {"opus": "global.anthropic.claude-opus-4-8"} + assert relayed is False + def test_invalid_provider_returns_error(self, monkeypatch): self._patch(monkeypatch, None, "boom") models, error, relayed = agents_mod.resolve_provider_models( diff --git a/tests/test_cli.py b/tests/test_cli.py index cd2d3f63..3a5c1ae0 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -618,6 +618,23 @@ def test_claude_enable_model_discovery_sets_ucode_env(self): assert os.environ["ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY"] == "1" assert mock_launch.call_args.args[1].args == [] + def test_claude_parent_is_forwarded(self): + with patch("ucode.cli._launch_tool") as mock_launch: + result = runner.invoke(app, ["claude", "--parent", "main.default"]) + + assert result.exit_code == 0, result.output + assert mock_launch.call_args.kwargs["parent_schema"] == "main.default" + assert os.environ["ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY"] == "1" + + def test_claude_provider_and_parent_are_mutually_exclusive(self): + result = runner.invoke( + app, + ["claude", "--provider", "main.default.provider", "--parent", "main.default"], + ) + + assert result.exit_code == 1 + assert "--provider and --parent cannot be used together" in result.output + def test_claude_enable_model_discovery_is_hidden_from_help(self): result = runner.invoke(app, ["claude", "--help"]) diff --git a/tests/test_string_utils.py b/tests/test_string_utils.py new file mode 100644 index 00000000..9c369df8 --- /dev/null +++ b/tests/test_string_utils.py @@ -0,0 +1,28 @@ +"""Tests for string validation helpers.""" + +import pytest + +from ucode.string_utils import is_valid_catalog_schema + + +@pytest.mark.parametrize("value", ["system.ai", "main.default", "my-catalog.my_schema"]) +def test_catalog_schema(value): + assert is_valid_catalog_schema(value) + + +@pytest.mark.parametrize( + "value", + [ + "", + "main", + "main.default.extra", + ".default", + "main.", + "main/development.models", + "main dev.models", + "main.\tmodels", + "main.\x7fmodels", + ], +) +def test_invalid_catalog_schema(value): + assert not is_valid_catalog_schema(value)