diff --git a/src/ucode/agents/__init__.py b/src/ucode/agents/__init__.py index 8ea3e20d..fd630147 100644 --- a/src/ucode/agents/__init__.py +++ b/src/ucode/agents/__init__.py @@ -426,6 +426,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": @@ -444,6 +445,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 a84a2911..c83b161e 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, @@ -44,6 +48,7 @@ reconcile_managed_file, revert_managed_file, ) +from ucode.model_service_headers import model_service_routing_headers from ucode.smart_routing import v2 as smart_routing_v2 from ucode.smart_routing.claude_hooks import ( remove_smart_routing_hooks, @@ -166,7 +171,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, + MODEL_SERVICE_PARENT_SCHEMA_HEADER, } ) CLAUDE_TRACING_STOP_HOOK_SUFFIX = " autolog claude stop-hook" @@ -321,6 +327,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. @@ -356,8 +363,10 @@ def render_overlay( "x-databricks-use-coding-agent-mode: true", f"User-Agent: ucode/{ucode_version()} claude/{agent_version('claude')}", ] - if provider: - header_lines.append(f"Databricks-Model-Provider-Service: {provider}") + header_lines.extend( + f"{name}: {value}" + for name, value in model_service_routing_headers(provider, parent_schema).items() + ) # 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) @@ -572,6 +581,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) @@ -593,6 +603,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 ef4e2231..34260ef0 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -128,6 +128,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, @@ -1965,10 +1966,13 @@ 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 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 @@ -2063,6 +2067,7 @@ def _launch_tool( ) if managed_provider: provider = managed_provider + effective_parent_schema = None if provider else parent_schema # 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: @@ -2157,6 +2162,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=effective_parent_schema, ) # Relayed = a Claude subscription: forward --model to Claude Code's own flag, like `-- --model X`. if tool == "claude" and provider and relayed and model and not forwarded_model: @@ -2500,6 +2506,13 @@ def claude_cmd( "before any `--` separator.", ), ] = None, + parent: Annotated[ + str | None, + typer.Option( + "--parent", + help="Discover model services in `.`.", + ), + ] = None, model: Annotated[ str | None, typer.Option( @@ -2571,7 +2584,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): _launch_tool( @@ -2582,6 +2595,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..0c7cd80c 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/model_service_headers.py b/src/ucode/model_service_headers.py new file mode 100644 index 00000000..7b54ec59 --- /dev/null +++ b/src/ucode/model_service_headers.py @@ -0,0 +1,17 @@ +"""Headers used to route model service requests.""" + +from __future__ import annotations + +from ucode.constants import MODEL_PROVIDER_SERVICE_HEADER, MODEL_SERVICE_PARENT_SCHEMA_HEADER + + +def model_service_routing_headers( + provider: str | None, + parent_schema: str | None, +) -> dict[str, str]: + # A provider selects one MPS; a parent discovers Model Services. + if provider: + return {MODEL_PROVIDER_SERVICE_HEADER: provider} + if parent_schema: + return {MODEL_SERVICE_PARENT_SCHEMA_HEADER: parent_schema} + return {} 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 11c50978..2a96aa81 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -288,7 +288,7 @@ def test_relayed_sends_mps_header_but_not_swap_token(self): relayed_base_url="http://127.0.0.1:9", ) headers = overlay["env"]["ANTHROPIC_CUSTOM_HEADERS"] - assert "Databricks-Model-Provider-Service: c.s.mps" in headers + assert "databricks-model-provider-service: c.s.mps" in headers assert "X-Databricks-AI-Gateway-Token" not in headers def test_model_overrides_when_all_provided(self): @@ -349,7 +349,7 @@ def test_fable_not_pinned_under_provider(self): def test_provider_adds_routing_header(self): overlay, _ = claude.render_overlay(WS, "s4", provider="main.aarushi.aarushi-claude") assert ( - "Databricks-Model-Provider-Service: main.aarushi.aarushi-claude" + "databricks-model-provider-service: main.aarushi.aarushi-claude" in overlay["env"]["ANTHROPIC_CUSTOM_HEADERS"] ) @@ -369,7 +369,26 @@ def test_provider_skips_model_pinning(self): 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"] + 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_provider_suppresses_discovery_header(self): + overlay, _ = claude.render_overlay( + WS, + "s4", + provider="main.default.anthropic", + parent_schema="main.default", + ) + assert ( + "databricks-model-service-parent-schema" + not in overlay["env"]["ANTHROPIC_CUSTOM_HEADERS"] + ) def test_bedrock_provider_pins_model_ids(self): provider_models = { @@ -390,7 +409,7 @@ def test_bedrock_provider_pins_model_ids(self): # Bedrock ids are pinned verbatim — no `[1m]` suffix mangling. assert "[1m]" not in env["ANTHROPIC_DEFAULT_OPUS_MODEL"] assert ( - "Databricks-Model-Provider-Service: main.bob.bedrock-svc" + "databricks-model-provider-service: main.bob.bedrock-svc" in env["ANTHROPIC_CUSTOM_HEADERS"] ) @@ -407,7 +426,7 @@ def test_non_relayed_provider_pins_tier_via_anthropic_model(self): env = overlay["env"] assert env["ANTHROPIC_MODEL"] == "claude-haiku-4-5" assert ( - "Databricks-Model-Provider-Service: main.mcao.anthropic-mps" + "databricks-model-provider-service: main.mcao.anthropic-mps" in (env["ANTHROPIC_CUSTOM_HEADERS"]) ) assert "apiKeyHelper" in overlay @@ -459,6 +478,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 9ff7d5ee..640e9c19 100644 --- a/tests/test_agents_init.py +++ b/tests/test_agents_init.py @@ -387,6 +387,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 00a01baf..25dfcff8 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -583,6 +583,14 @@ 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_enable_model_discovery_is_hidden_from_help(self): result = runner.invoke(app, ["claude", "--help"]) diff --git a/tests/test_model_service_headers.py b/tests/test_model_service_headers.py new file mode 100644 index 00000000..e5c2ee87 --- /dev/null +++ b/tests/test_model_service_headers.py @@ -0,0 +1,22 @@ +"""Tests for model service routing headers.""" + +from ucode.constants import MODEL_PROVIDER_SERVICE_HEADER, MODEL_SERVICE_PARENT_SCHEMA_HEADER +from ucode.model_service_headers import model_service_routing_headers + + +def test_provider_header(): + assert model_service_routing_headers("main.default.provider", None) == { + MODEL_PROVIDER_SERVICE_HEADER: "main.default.provider" + } + + +def test_parent_schema_header(): + assert model_service_routing_headers(None, "main.default") == { + MODEL_SERVICE_PARENT_SCHEMA_HEADER: "main.default" + } + + +def test_provider_takes_precedence(): + assert model_service_routing_headers("main.default.provider", "main.default") == { + MODEL_PROVIDER_SERVICE_HEADER: "main.default.provider" + } 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)