From 27a449521833b36c8b04c9629b2c39c9cfc97793 Mon Sep 17 00:00:00 2001 From: Andy Xu Date: Wed, 16 Sep 2026 20:24:17 +0000 Subject: [PATCH 1/6] cli: support launch-only Codex request headers --- README.md | 4 ++ src/ucode/agents/args.py | 1 + src/ucode/agents/codex.py | 76 ++++++++++++++++++++++++++-- src/ucode/cli.py | 66 ++++++++++++++++++++++++ tests/test_agent_codex.py | 54 ++++++++++++++++++++ tests/test_cli.py | 43 ++++++++++++++++ tests/test_codex_smart_routing_v2.py | 29 +++++++++++ 7 files changed, 270 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 94fcacf80..d4fda3c1c 100644 --- a/README.md +++ b/README.md @@ -156,6 +156,10 @@ ug skills remove --location main.default --via mcp | `ug revert` | Clear saved state and restore backed-up config files | | `ug upgrade` | Upgrade Unity Gateway | +`ug codex --header 'X-Development-Route: test-target'` adds a repeatable, launch-only +request header for Codex. Use it only for non-secret development routing values: headers are +passed through environment variables and are not written to persistent Codex configuration. + Databricks AI Tools are installed only by `ug configure`, never by agent launch commands. Use `--enable-databricks-ai-tools` or `--disable-databricks-ai-tools` with `ug configure` to control installation. diff --git a/src/ucode/agents/args.py b/src/ucode/agents/args.py index faca49dfe..a899e439a 100644 --- a/src/ucode/agents/args.py +++ b/src/ucode/agents/args.py @@ -11,6 +11,7 @@ class LaunchOptions: launch_smart_routing: bool = False user_pinned_model: str | None = None + custom_headers: tuple[tuple[str, str], ...] = () def explicit_model_arg_value(tool_args: list[str]) -> str | None: diff --git a/src/ucode/agents/codex.py b/src/ucode/agents/codex.py index 340fc5ee2..f65006658 100644 --- a/src/ucode/agents/codex.py +++ b/src/ucode/agents/codex.py @@ -93,6 +93,7 @@ LEGACY_CODEX_BACKUP_PATH = APP_DIR / "codex-config.backup.toml" CODEX_MODEL_PROVIDER_NAME = "Databricks" LEGACY_CODEX_MODEL_PROVIDER_NAME = "ucode-databricks" +CUSTOM_HEADER_ENV_PREFIX = "UCODE_CUSTOM_HEADER" # ug owns the whole provider http_headers table, so it is pruned before each merge and rewritten # from render_overlay — dropping stale routing and admin headers that deep_merge cannot delete. _PROVIDER_HTTP_HEADERS_KEY_PATHS = [ @@ -520,6 +521,61 @@ def _parse_managed_config(text: str) -> dict: raise RuntimeError(f"invalid TOML: {exc}") from exc +def _managed_header_names() -> set[str]: + """Return header names fixed by OS-managed config.""" + path = codex_managed_config_path() + if path is None: + return set() + text = read_managed_file(path) + if text is None: + return set() + try: + doc = _parse_managed_config(text) + except RuntimeError: + # Configuration validates the managed file before launch and reports the + # full parse error there. + return set() + providers = doc.get("model_providers") + provider = providers.get(CODEX_MODEL_PROVIDER_NAME) if isinstance(providers, dict) else None + if not isinstance(provider, dict): + return set() + names: set[str] = set() + for table_name in ("http_headers", "env_http_headers"): + headers = provider.get(table_name) + if isinstance(headers, dict): + names.update(str(name).casefold() for name in headers) + return names + + +def _with_custom_headers(doc: dict, custom_headers: dict[str, str]) -> dict: + """Return a launch-only config that reads custom header values from the environment.""" + if not custom_headers: + return doc + blocked_names = _managed_header_names() & {name.casefold() for name in custom_headers} + if blocked_names: + names = ", ".join(sorted(blocked_names)) + raise RuntimeError(f"--header cannot override OS-managed Codex header(s): {names}.") + + launch_doc = copy.deepcopy(doc) + providers = launch_doc.get("model_providers") + provider = providers.get(CODEX_MODEL_PROVIDER_NAME) if isinstance(providers, dict) else None + if not isinstance(provider, dict): + raise RuntimeError("Codex's Databricks model provider configuration is missing or invalid.") + env_headers = provider.setdefault("env_http_headers", {}) + if not isinstance(env_headers, dict): + raise RuntimeError("Codex's Databricks environment-header configuration is invalid.") + + requested_names = {name.casefold() for name in custom_headers} + for existing_name in list(env_headers): + if str(existing_name).casefold() in requested_names: + env_headers.pop(existing_name, None) + for index, (name, value) in enumerate(custom_headers.items()): + env_name = f"{CUSTOM_HEADER_ENV_PREFIX}_{os.getpid()}_{index}" + os.environ[env_name] = value + env_headers[name] = env_name + return launch_doc + + def managed_config_is_current(state: dict) -> bool: path = codex_managed_config_path() if path is None: @@ -917,8 +973,9 @@ def launch( *, options: LaunchOptions, ) -> None: + custom_headers = dict(options.custom_headers) if options.launch_smart_routing: - _launch_smart_routing(state, tool_args) + _launch_smart_routing(state, tool_args, custom_headers=custom_headers) return clear_model_preferences(state) binary = SPEC["binary"] @@ -948,6 +1005,8 @@ def launch( token = _launch_token(state, workspace) os.environ["OAUTH_TOKEN"] = token if _use_legacy_layout(): + if custom_headers: + raise RuntimeError(f"--header requires Codex {MINIMUM_CODEX_VERSION_TEXT} or newer.") print_warning_err( f"Codex {agent_version(binary)} is outdated. Upgrade Codex to " f"{LEGACY_LAYOUT_CODEX_VERSION_TEXT} or newer, then run `codex --version` to verify " @@ -1005,6 +1064,7 @@ def launch( slugs = catalog_slugs(catalog) if slugs: profile_doc["model"] = slugs[0] + profile_doc = _with_custom_headers(profile_doc, custom_headers) _run_codex( state, [binary, *codex_config_args(profile_doc)], @@ -1014,7 +1074,9 @@ def launch( ) -def _launch_smart_routing(state: dict, tool_args: list[str]) -> None: +def _launch_smart_routing( + state: dict, tool_args: list[str], *, custom_headers: dict[str, str] | None = None +) -> None: """Launch the Codex TUI through the smart-routing interposer.""" binary = SPEC["binary"] @@ -1026,12 +1088,20 @@ def _launch_smart_routing(state: dict, tool_args: list[str]) -> None: or (codex_model_id(models[0]) if models else None) or APP_SERVER_SMART_ROUTING_STARTING_MODEL ) + if custom_headers: + + def routed_render_overlay(*args, **kwargs) -> dict: + return _with_custom_headers(render_overlay(*args, **kwargs), custom_headers) + + else: + routed_render_overlay = render_overlay + smart_routing_v2.launch_codex( state, tool_args, binary=binary, start_model=start_model, - render_overlay=render_overlay, + render_overlay=routed_render_overlay, ) diff --git a/src/ucode/cli.py b/src/ucode/cli.py index 44d2b053f..961f59e36 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -4,6 +4,7 @@ from __future__ import annotations import os +import re import shutil import subprocess from collections.abc import Iterator @@ -2458,6 +2459,55 @@ def _smart_routing_launch_shape(tool: str, tool_args: list[str], explicit_prompt return tool == "claude" and tool_args[0].startswith("-") +_HTTP_HEADER_NAME_PATTERN = re.compile(r"[!#$%&'*+\-.^_`|~0-9A-Za-z]+") +_PROTECTED_CUSTOM_HEADER_NAMES = frozenset( + { + "authorization", + "connection", + "content-length", + "cookie", + "databricks-model-provider-service", + "databricks-model-service-parent-schema", + "databricks-smart-router-recipe", + "host", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailer", + "transfer-encoding", + "upgrade", + "user-agent", + "x-api-key", + "x-databricks-ai-gateway-token", + "x-databricks-use-coding-agent-mode", + } +) + + +def _parse_custom_headers(values: list[str] | None) -> dict[str, str]: + """Parse repeatable ``--header 'Name: value'`` options.""" + parsed: dict[str, tuple[str, str]] = {} + for item in values or []: + name, separator, value = item.partition(":") + name = name.strip() + if not separator or _HTTP_HEADER_NAME_PATTERN.fullmatch(name) is None: + raise RuntimeError("--header must use the format `Name: value` with a valid name.") + value = value.strip() + if any( + ord(character) < 32 or ord(character) == 127 or character in "\u0085\u2028\u2029" + for character in value + ): + raise RuntimeError( + "--header values cannot contain control characters or line separators." + ) + normalized_name = name.casefold() + if normalized_name in _PROTECTED_CUSTOM_HEADER_NAMES: + raise RuntimeError(f"--header cannot override protected header '{name}'.") + parsed[normalized_name] = (name, value) + return dict(parsed.values()) + + def _launch_options( tool: str, tool_args: list[str], @@ -2466,10 +2516,12 @@ def _launch_options( explicit_prompt: bool, user_pinned_model: str | None, provider: str | None, + custom_headers: dict[str, str] | None = None, ) -> LaunchOptions: return LaunchOptions( # Pinned models for providers are resolved above through the provider-specific launch path. user_pinned_model=user_pinned_model if provider is None else None, + custom_headers=tuple((custom_headers or {}).items()), launch_smart_routing=( # Smart routing is enabled globally. smart_routing_enabled @@ -2520,11 +2572,13 @@ def _launch_tool( model: str | None = None, parent_schema: str | None = None, custom_oauth: CustomOAuthConfig | None = None, + headers: list[str] | None = None, ) -> None: try: tool = normalize_tool(tool_name) if not custom_oauth_cli_enabled(custom_oauth): os.environ.pop(CUSTOM_OAUTH_CLI_ENV_VAR, None) + custom_headers = _parse_custom_headers(headers) # Before any status print: a stdio-protocol subcommand owns stdout, so # every ug line from here on must go to stderr instead. if _child_owns_stdout(tool, ctx.args): @@ -2836,6 +2890,7 @@ def _launch_tool( # initial/fallback model and still participates in a routed session. user_pinned_model=model or forwarded_model, provider=provider, + custom_headers=custom_headers, ) print_success(f"Starting {TOOL_SPECS[tool]['display']}") with _managed_smart_routing_environment(managed, tool): @@ -2877,6 +2932,15 @@ def _launch_tool( ), ] +CustomHeaderOption = Annotated[ + list[str] | None, + typer.Option( + "--header", + help="Add an HTTP header to AI Gateway requests as `Name: value`; repeatable. " + "Pass before any `--` separator. Credentials and transport headers are not allowed.", + ), +] + _PROMPT_SUFFIX_KEY = "ucode_explicit_prompt_suffix" @@ -3030,6 +3094,7 @@ def codex_cmd( help="Discover model services in `.`. Example: main.default", ), ] = None, + header: CustomHeaderOption = None, refresh: Annotated[ bool, typer.Option( @@ -3078,6 +3143,7 @@ def codex_cmd( workspace_url=workspace, parent_schema=model_location, custom_oauth=custom_oauth, + headers=header, ) diff --git a/tests/test_agent_codex.py b/tests/test_agent_codex.py index 550123231..228c41a93 100644 --- a/tests/test_agent_codex.py +++ b/tests/test_agent_codex.py @@ -834,6 +834,7 @@ def _patch(tmp_path, monkeypatch): lambda workspace, profile=None, force_refresh=False: "tok", ) monkeypatch.setattr(codex, "clear_model_preferences", lambda state: False) + monkeypatch.setattr(codex, "codex_managed_config_path", lambda: None) return launches def test_sets_oauth_token(self, tmp_path, monkeypatch): @@ -849,6 +850,46 @@ def test_sets_oauth_token(self, tmp_path, monkeypatch): assert os.environ["OAUTH_TOKEN"] == "fresh-token" assert launches[0][-1] == "--search" + def test_custom_header_is_launch_only(self, tmp_path, monkeypatch): + launches = self._patch(tmp_path, monkeypatch) + value = "route://development/test" + + codex.launch( + {"workspace": WS}, + [], + options=LaunchOptions(custom_headers=(("X-Development-Route", value),)), + ) + + env_name = f"{codex.CUSTOM_HEADER_ENV_PREFIX}_{os.getpid()}_0" + provider_arg = next( + arg for arg in launches[0] if arg.startswith("model_providers.Databricks=") + ) + assert "X-Development-Route" in provider_arg + assert env_name in provider_arg + assert value not in provider_arg + assert os.environ[env_name] == value + assert "X-Development-Route" not in codex.CODEX_CONFIG_PATH.read_text(encoding="utf-8") + monkeypatch.delenv(env_name) + + @pytest.mark.parametrize("header_table", ["http_headers", "env_http_headers"]) + def test_custom_header_rejects_managed_header(self, tmp_path, monkeypatch, header_table): + self._patch(tmp_path, monkeypatch) + managed_path = tmp_path / "managed_config.toml" + managed_path.write_text( + f'[model_providers.Databricks.{header_table}]\nX-Test = "ENTERPRISE_VALUE"\n', + encoding="utf-8", + ) + monkeypatch.setattr(codex, "codex_managed_config_path", lambda: managed_path) + + with pytest.raises(RuntimeError, match="OS-managed Codex header"): + codex.launch( + {"workspace": WS}, + [], + options=LaunchOptions(custom_headers=(("x-test", "temporary"),)), + ) + + assert "temporary" not in managed_path.read_text(encoding="utf-8") + def test_provider_discovery_uses_authoritative_catalog(self, tmp_path, monkeypatch): launches = self._patch(tmp_path, monkeypatch) catalog_path = tmp_path / "models.json" @@ -1098,6 +1139,19 @@ def test_provider_rejects_managed_model_catalog(self, tmp_path, monkeypatch): assert launches == [] + def test_custom_header_requires_modern_codex(self, tmp_path, monkeypatch): + launches = self._patch(tmp_path, monkeypatch) + monkeypatch.setattr(codex, "agent_version", lambda binary: "0.133.0") + + with pytest.raises(RuntimeError, match="--header requires Codex 0.145.0 or newer"): + codex.launch( + {"workspace": WS}, + [], + options=LaunchOptions(custom_headers=(("X-Test", "temporary"),)), + ) + + assert launches == [] + def test_non_provider_launch_removes_stale_provider_header(self, tmp_path, monkeypatch): launches = self._patch(tmp_path, monkeypatch) profile_path = tmp_path / "ucode.config.toml" diff --git a/tests/test_cli.py b/tests/test_cli.py index 35d2dec8e..a6a7f2b87 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -763,6 +763,49 @@ def test_codex_model_location_is_forwarded(self): assert mock_launch.call_args.kwargs["parent_schema"] == "main.default" assert mock_launch.call_args.args[1].args == [] + def test_codex_headers_are_forwarded(self): + with patch("ucode.cli._launch_tool") as mock_launch: + result = runner.invoke( + app, ["codex", "--header", "X-First: one", "--header", "X-Second: two:three"] + ) + + assert result.exit_code == 0, result.output + assert mock_launch.call_args.kwargs["headers"] == [ + "X-First: one", + "X-Second: two:three", + ] + + def test_headers_parse_values_with_colons_and_deduplicate_case_insensitively(self): + assert cli_mod._parse_custom_headers( + [ + "X-Test: first", + "x-test: second", + "X-Development-Route: route://development/test", + ] + ) == { + "x-test": "second", + "X-Development-Route": "route://development/test", + } + + @pytest.mark.parametrize( + ("value", "message"), + [ + ("missing-separator", "format `Name: value`"), + ("bad name: value", "format `Name: value`"), + ("X-Test: line\nbreak", "control characters"), + ("X-Test: safe\u0085Authorization: injected", "line separators"), + ("X-Test: safe\u2028Authorization: injected", "line separators"), + ("X-Test: safe\u2029Authorization: injected", "line separators"), + ("Authorization: secret", "protected header"), + ("Cookie: secret", "protected header"), + ("Databricks-Smart-Router-Recipe: custom", "protected header"), + ("databricks-smart-router-recipe: custom", "protected header"), + ], + ) + def test_invalid_codex_header_is_rejected(self, value, message): + with pytest.raises(RuntimeError, match=message): + cli_mod._parse_custom_headers([value]) + def test_codex_provider_and_model_location_are_mutually_exclusive(self): with _launch_policy_patches(None): result = runner.invoke( diff --git a/tests/test_codex_smart_routing_v2.py b/tests/test_codex_smart_routing_v2.py index c9ca054f5..4bd632199 100644 --- a/tests/test_codex_smart_routing_v2.py +++ b/tests/test_codex_smart_routing_v2.py @@ -76,6 +76,35 @@ def launch_v2(state, tool_args, **kwargs): ) ] + def test_codex_smart_routing_preserves_custom_header(self, monkeypatch): + captured = {} + monkeypatch.setenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, "1") + monkeypatch.setattr(codex, "_smart_routing_config_model", lambda state: "gpt-start") + monkeypatch.setattr(codex, "codex_managed_config_path", lambda: None) + + def launch_v2(state, tool_args, **kwargs): + captured.update(kwargs) + raise SystemExit(0) + + monkeypatch.setattr(v2, "launch_codex", launch_v2) + custom_headers = {"X-Development-Route": "route://development/test"} + + with pytest.raises(SystemExit): + codex.launch( + {"workspace": WS}, + [], + options=LaunchOptions( + launch_smart_routing=True, + custom_headers=tuple(custom_headers.items()), + ), + ) + + overlay = captured["render_overlay"](WS, "gpt-start") + headers = overlay["model_providers"]["Databricks"]["env_http_headers"] + env_name = headers["X-Development-Route"] + assert os.environ[env_name] == custom_headers["X-Development-Route"] + monkeypatch.delenv(env_name) + @pytest.mark.parametrize( "tool_args", [ From 3d5d196270b9bb68446571e3320a49d13f950000 Mon Sep 17 00:00:00 2001 From: Andy Xu Date: Tue, 6 Oct 2026 19:50:58 +0000 Subject: [PATCH 2/6] Protect Codex admin headers without an OS-managed file --- src/ucode/agents/codex.py | 6 ++++++ tests/README.md | 4 ++++ tests/integration/README.md | 3 +++ tests/test_agent_codex.py | 20 ++++++++++++++++++++ 4 files changed, 33 insertions(+) diff --git a/src/ucode/agents/codex.py b/src/ucode/agents/codex.py index 6d107f015..a99f4812a 100644 --- a/src/ucode/agents/codex.py +++ b/src/ucode/agents/codex.py @@ -1144,6 +1144,12 @@ def launch( options: LaunchOptions, ) -> None: custom_headers = dict(options.custom_headers) + blocked_names = {name.casefold() for name in custom_headers} & { + name.strip().casefold() for name in (state.get("codex_http_headers") or {}) + } + if blocked_names: + names = ", ".join(sorted(blocked_names)) + raise RuntimeError(f"--header cannot override managed Codex header(s): {names}.") if options.launch_smart_routing: _launch_smart_routing(state, tool_args, custom_headers=custom_headers) return diff --git a/tests/README.md b/tests/README.md index f9e9a8130..b2c80493d 100644 --- a/tests/README.md +++ b/tests/README.md @@ -42,6 +42,10 @@ These are component checks, not live Windows coverage for every agent. Agent configuration tests also verify `ug` auth/MCP helper commands, including quoted executable paths and replacement of legacy `ucode` routing/web-search helpers. +Custom request headers have component coverage in `test_cli.py`, `test_agent_codex.py`, +and `test_codex_smart_routing_v2.py`: parsing, launch-only values, and administrator +header collisions with or without an OS-managed file. Live custom-header journeys are not covered. + `test_mcp_web_search.py` and `test_agent_claude.py` cover custom OAuth search registration, stale registration repair, SDK cache reuse/refresh, CLI profile selection, and errors without browser consent through the MCP handler. These diff --git a/tests/integration/README.md b/tests/integration/README.md index 5f631a5eb..d2d6c8ead 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -18,6 +18,9 @@ The existing unit tests keep their fixtures. Integration has an independent pytest configuration and uses `--confcutdir` so those fixtures cannot leak in. It is not collected by the default `uv run pytest` command. +Custom `--header` parsing, delivery configuration, and collision rejection have component +coverage listed in `../README.md`; live custom-header journeys are not covered. + The `smart_defaults` wire schema, legacy `spend_tiers` cache reads, and recommendation request gating are covered by unit/component tests listed in `../README.md`. This suite does not yet assert live `recommendModel` request counts for configs with and without tiers. diff --git a/tests/test_agent_codex.py b/tests/test_agent_codex.py index cfd9f5b45..58c83493f 100644 --- a/tests/test_agent_codex.py +++ b/tests/test_agent_codex.py @@ -1035,6 +1035,26 @@ def test_custom_header_rejects_managed_header(self, tmp_path, monkeypatch, heade assert "temporary" not in managed_path.read_text(encoding="utf-8") + @pytest.mark.parametrize("smart_routing", [False, True]) + def test_custom_header_rejects_admin_header_without_os_config( + self, tmp_path, monkeypatch, smart_routing + ): + launches = self._patch(tmp_path, monkeypatch) + original = codex.CODEX_CONFIG_PATH.read_text(encoding="utf-8") + + with pytest.raises(RuntimeError, match="cannot override managed Codex header"): + codex.launch( + {"workspace": WS, "codex_http_headers": {"X-Test": "admin"}}, + [], + options=LaunchOptions( + launch_smart_routing=smart_routing, + custom_headers=(("x-test", "temporary"),), + ), + ) + + assert launches == [] + assert codex.CODEX_CONFIG_PATH.read_text(encoding="utf-8") == original + @pytest.mark.parametrize("custom_catalog", [None, "/user/isaac-app-model-catalog.json"]) def test_native_update_detaches_catalog_without_discovery( self, tmp_path, monkeypatch, custom_catalog From 6535ada754ee94f4403e2cff49aa27d362d9613a Mon Sep 17 00:00:00 2001 From: Andy Xu Date: Tue, 6 Oct 2026 21:29:06 +0000 Subject: [PATCH 3/6] Thread custom headers through gateway model discovery --- src/ucode/agents/__init__.py | 5 ++++ src/ucode/cli.py | 32 +++++++++++++++++++++---- src/ucode/databricks.py | 46 ++++++++++++++++++++++++++++++------ tests/test_cli.py | 28 ++++++++++++++++++++++ tests/test_databricks.py | 27 ++++++++++++++++++++- 5 files changed, 125 insertions(+), 13 deletions(-) diff --git a/src/ucode/agents/__init__.py b/src/ucode/agents/__init__.py index 242c43f57..2498e6662 100644 --- a/src/ucode/agents/__init__.py +++ b/src/ucode/agents/__init__.py @@ -392,6 +392,11 @@ def resolve_gemini_provider_model( ) +def validate_custom_headers(tool: str, state: dict, headers: dict[str, str]) -> None: + if tool == "codex" and headers: + codex.validate_custom_headers(state, headers) + + def configure_tool( tool: str, state: dict, diff --git a/src/ucode/cli.py b/src/ucode/cli.py index d89001be5..478dcbfae 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -37,6 +37,7 @@ resolve_gemini_provider_model, resolve_launch_model, resolve_provider_models, + validate_custom_headers, ) from ucode.agents import claude as claude_agent from ucode.agents import codex as codex_agent @@ -478,6 +479,7 @@ def configure_shared_state( databricks_ai_tools_enabled: bool | None = None, custom_oauth: CustomOAuthConfig | None = None, clear_custom_oauth: bool = False, + request_headers: dict[str, str] | None = None, ) -> dict: """Log into Databricks, verify AI Gateway, fetch model lists, persist state. @@ -597,6 +599,14 @@ def configure_shared_state( with spinner("Verifying Unity AI Gateway..."): if not cli_custom_oauth: token = get_databricks_token(workspace, profile) + if request_headers: + managed, _ = _fetch_managed_config(state) + for tool in tools or []: + validate_custom_headers( + tool, + resolve_state(managed, state, tool) if managed is not None else state, + request_headers, + ) model_service_probe = probe_unity_gateway_capabilities(workspace, token) if model_service_probe.resource_available: print_success("Unity Gateway connected") @@ -2260,7 +2270,11 @@ def claude_router_hook_cmd( sys.stdout.write(json.dumps(output)) -def _auto_configure_tool(tool: str, custom_oauth: CustomOAuthConfig | None = None) -> None: +def _auto_configure_tool( + tool: str, + custom_oauth: CustomOAuthConfig | None = None, + request_headers: dict[str, str] | None = None, +) -> None: """Configure a tool for launch without sending a separate validation prompt. The real agent session follows immediately; explicit configure retains the @@ -2272,6 +2286,8 @@ def _auto_configure_tool(tool: str, custom_oauth: CustomOAuthConfig | None = Non if not workspace: workspace, profile = _prompt_for_configuration(tool) configure_kwargs = {"custom_oauth": custom_oauth} if custom_oauth is not None else {} + if request_headers: + configure_kwargs["request_headers"] = request_headers state = configure_shared_state(workspace, profile=profile, tools=[tool], **configure_kwargs) state = configure_single_tool(tool, state) @@ -2683,10 +2699,11 @@ def _launch_tool( skip_cli_version_check=skip_preflight, ) if needs_auto_configure: - if custom_oauth is None: - _auto_configure_tool(tool) - else: - _auto_configure_tool(tool, custom_oauth=custom_oauth) + _auto_configure_tool( + tool, + **({"custom_oauth": custom_oauth} if custom_oauth is not None else {}), + **({"request_headers": custom_headers} if custom_headers else {}), + ) state = ensure_provider_state(tool) # Remembered before the fallback below collapses the two cases: a managed config may not # silently override a provider the user typed on the command line (it errors instead). @@ -2702,6 +2719,11 @@ def _launch_tool( coding_agent_config_feature_disabled = False if managed is None: managed, coding_agent_config_feature_disabled = _fetch_managed_config(state) + validate_custom_headers( + tool, + resolve_state(managed, state, tool) if managed is not None else state, + custom_headers, + ) _reject_managed_launch_source_options( managed, provider=explicit_provider, diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index 970f51ecc..e7f464054 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -2949,6 +2949,28 @@ def collect(result, schemas_done, schemas_total): return sorted(names), None +def _merge_model_discovery_headers( + request_headers: dict[str, str] | None, + mandatory_header: tuple[str, str] | None = None, +) -> dict[str, str] | None: + """Merge caller headers with an optional routing header for model discovery. + + Header names are case-insensitive. The routing header is added last so a + caller cannot override the provider or parent-schema selector used for a + scoped discovery request. + """ + merged = dict(request_headers or {}) + if mandatory_header is not None: + name, value = mandatory_header + merged = { + existing: existing_value + for existing, existing_value in merged.items() + if existing.casefold() != name.casefold() + } + merged[name] = value + return merged or None + + def _get_anthropic_models_json( workspace: str, token: str, @@ -3385,25 +3407,35 @@ def _fetch_codex_model_catalog( workspace: str, token: str, *, - source: CodexCatalogSource, - identifier: str, + source: CodexCatalogSource | None = None, + identifier: str | None = None, + request_headers: dict[str, str] | None = None, ) -> dict: - header_name, kind = source.value + mandatory_header = None + if source is not None: + assert identifier is not None + header_name, kind = source.value + mandatory_header = (header_name, identifier) + discovery_label = identifier + catalog_label = f"{kind} {identifier}" + else: + discovery_label = catalog_label = "workspace" + headers = _merge_model_discovery_headers(request_headers, mandatory_header) payload, reason = _http_get_json( f"{build_tool_base_url('codex', workspace)}/models", token, max_retries=2, - headers={header_name: identifier}, + **({"headers": headers} if headers is not None else {}), ) if reason: - message = f"Could not discover Codex models for {identifier}: {reason}" + message = f"Could not discover Codex models for {discovery_label}: {reason}" if "codex/v1/models is not enabled for this workspace" in reason.lower(): raise CodexMpsModelCatalogUnavailable(message) raise RuntimeError(message) if not isinstance(payload, dict) or not isinstance(payload.get("models"), list): - raise RuntimeError(f"{kind} {identifier} returned an invalid Codex model catalog.") + raise RuntimeError(f"{catalog_label} returned an invalid Codex model catalog.") if not payload["models"]: - raise RuntimeError(f"{kind} {identifier} returned no Codex models.") + raise RuntimeError(f"{catalog_label} returned no Codex models.") return payload diff --git a/tests/test_cli.py b/tests/test_cli.py index 81eead8cd..037c23e5f 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -904,6 +904,16 @@ def test_codex_headers_are_forwarded(self): "X-Second: two:three", ] + def test_codex_admin_header_collision_stops_before_discovery(self): + managed = {"enabled_agents": {"codex": {"http_headers": {"X-Test": "admin"}}}} + with _launch_policy_patches(managed) as calls: + result = runner.invoke(app, ["codex", "--header", "x-test: temporary"]) + + assert result.exit_code == 1 + assert "cannot override managed Codex header" in result.output + calls["shared"].assert_not_called() + calls["launch"].assert_not_called() + def test_headers_parse_values_with_colons_and_deduplicate_case_insensitively(self): assert cli_mod._parse_custom_headers( [ @@ -4901,6 +4911,24 @@ def _stub(monkeypatch): monkeypatch.setattr(cli_mod, "build_shared_base_urls", lambda w: {}) monkeypatch.setattr(cli_mod, "save_state", lambda s: None) + def test_custom_headers_validate_after_auth_before_discovery(self, monkeypatch): + self._stub(monkeypatch) + authenticated = MagicMock() + monkeypatch.setattr(cli_mod, "ensure_databricks_auth", authenticated) + discovery = MagicMock() + monkeypatch.setattr(cli_mod, "probe_unity_gateway_capabilities", discovery) + + def managed_after_auth(state): + authenticated.assert_called_once() + return {"enabled_agents": {"codex": {"http_headers": {"X-Route": "admin"}}}}, False + + monkeypatch.setattr(cli_mod, "_fetch_managed_config", managed_after_auth) + with pytest.raises(RuntimeError, match="cannot override managed Codex header"): + cli_mod.configure_shared_state( + MINIMAL_STATE["workspace"], tools=["codex"], request_headers={"x-route": "test"} + ) + discovery.assert_not_called() + def test_skips_family_discovery_and_fetches_web_search_model(self, monkeypatch): import ucode.cli as cli_mod diff --git a/tests/test_databricks.py b/tests/test_databricks.py index 40a47fcc0..9dcaac88d 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -86,11 +86,36 @@ def fake_get(url, token, **kwargs): "tok", source=db_mod.CodexCatalogSource.PROVIDER, identifier="main.default.openai", + request_headers={ + "databricks-model-provider-service": "wrong.provider", + "X-Trace-Id": "trace-1", + }, ) assert result["models"][0]["slug"] == "gpt-mps" assert seen["url"] == f"{WS}/ai-gateway/codex/v1/models" - assert seen["headers"] == {"Databricks-Model-Provider-Service": "main.default.openai"} + assert seen["headers"] == { + "X-Trace-Id": "trace-1", + "Databricks-Model-Provider-Service": "main.default.openai", + } + + def test_sends_custom_headers_without_selector(self, monkeypatch): + seen = {} + + def fake_get(url, token, **kwargs): + seen.update(url=url, token=token, **kwargs) + return {"models": [{"slug": "gpt-default"}]}, None + + monkeypatch.setattr(db_mod, "_http_get_json", fake_get) + + result = db_mod._fetch_codex_model_catalog( + WS, + "tok", + request_headers={"X-Trace-Id": "trace-1"}, + ) + + assert result["models"][0]["slug"] == "gpt-default" + assert seen["headers"] == {"X-Trace-Id": "trace-1"} def test_rejects_empty_catalog(self, monkeypatch): monkeypatch.setattr( From adb8bab4e2f7f199e1e127e11d2a5612da4bb648 Mon Sep 17 00:00:00 2001 From: Andy Xu Date: Tue, 6 Oct 2026 21:31:16 +0000 Subject: [PATCH 4/6] Keep header-routed Codex catalogs private to each launch --- src/ucode/agents/codex.py | 268 +++++++++++++++++++-------- src/ucode/launcher.py | 7 +- src/ucode/smart_routing/v2.py | 21 ++- tests/test_agent_codex.py | 100 +++++++++- tests/test_codex_smart_routing_v2.py | 33 +++- tests/test_launcher.py | 9 +- 6 files changed, 341 insertions(+), 97 deletions(-) diff --git a/src/ucode/agents/codex.py b/src/ucode/agents/codex.py index a99f4812a..f910dd077 100644 --- a/src/ucode/agents/codex.py +++ b/src/ucode/agents/codex.py @@ -559,15 +559,34 @@ def _managed_header_names() -> set[str]: return names -def _with_custom_headers(doc: dict, custom_headers: dict[str, str]) -> dict: - """Return a launch-only config that reads custom header values from the environment.""" +def validate_custom_headers(state: dict, custom_headers: dict[str, str]) -> None: + """Reject launch headers that collide with configured administrator values.""" if not custom_headers: - return doc - blocked_names = _managed_header_names() & {name.casefold() for name in custom_headers} + return + + requested_names = {name.casefold() for name in custom_headers} + blocked_names = requested_names & { + name.strip().casefold() for name in (state.get("codex_http_headers") or {}) + } + if blocked_names: + names = ", ".join(sorted(blocked_names)) + raise RuntimeError(f"--header cannot override managed Codex header(s): {names}.") + + blocked_names = _managed_header_names() & requested_names if blocked_names: names = ", ".join(sorted(blocked_names)) raise RuntimeError(f"--header cannot override OS-managed Codex header(s): {names}.") + +def _with_custom_headers(doc: dict, custom_headers: dict[str, str]) -> dict: + """Return a launch-only config that reads custom header values from the environment.""" + if not custom_headers: + return doc + # Keep this check at the final config boundary as well as before discovery. The latter avoids + # sending a request with a header that will later be rejected; this one protects callers that + # reach the launch renderer directly. + validate_custom_headers({}, custom_headers) + launch_doc = copy.deepcopy(doc) providers = launch_doc.get("model_providers") provider = providers.get(CODEX_MODEL_PROVIDER_NAME) if isinstance(providers, dict) else None @@ -921,6 +940,15 @@ def _model_catalog_path(workspace: str, scope: str) -> Path: return base.with_name(f"{base.stem}-{digest}{base.suffix}") +def _temporary_model_catalog_dir() -> tempfile.TemporaryDirectory: + """Allocate a private directory for a per-launch catalog.""" + parent = CODEX_MODEL_CATALOG_PATH.parent + parent.mkdir(parents=True, exist_ok=True) + return tempfile.TemporaryDirectory( + prefix=f".{CODEX_MODEL_CATALOG_PATH.stem}-launch-", dir=parent + ) + + def _write_model_catalog(path: Path, catalog: dict) -> None: if is_dry_run(): write_json_file(path, catalog) @@ -1126,6 +1154,7 @@ def _run_codex( *, otel_tracing: bool, workspace: str | None, + wait_for_exit: bool = False, ) -> None: """Launch Codex — via the loopback proxy when OTLP tracing is on, else exec-replace.""" if tool_args[:1] == ["update"]: @@ -1134,7 +1163,11 @@ def _run_codex( if otel_tracing and workspace: _launch_codex_with_otel_proxy(state, base_argv, tool_args, workspace) else: - exec_or_spawn([*base_argv, *tool_args]) + argv = [*base_argv, *tool_args] + if wait_for_exit: + exec_or_spawn(argv, wait_for_exit=True) + else: + exec_or_spawn(argv) def launch( @@ -1144,12 +1177,7 @@ def launch( options: LaunchOptions, ) -> None: custom_headers = dict(options.custom_headers) - blocked_names = {name.casefold() for name in custom_headers} & { - name.strip().casefold() for name in (state.get("codex_http_headers") or {}) - } - if blocked_names: - names = ", ".join(sorted(blocked_names)) - raise RuntimeError(f"--header cannot override managed Codex header(s): {names}.") + validate_custom_headers(state, custom_headers) if options.launch_smart_routing: _launch_smart_routing(state, tool_args, custom_headers=custom_headers) return @@ -1211,53 +1239,78 @@ def launch( updating = tool_args[:1] == ["update"] if updating and _is_ucode_catalog_reference(profile_doc.get("model_catalog_json")): profile_doc.pop("model_catalog_json") - if workspace and token and (provider or parent_schema) and not updating: - try: - if provider is not None: - catalog_source = CodexCatalogSource.PROVIDER - catalog_identifier = provider - catalog_scope = f"provider:{provider}" - elif parent_schema is not None: - catalog_source = CodexCatalogSource.PARENT_SCHEMA - catalog_identifier = parent_schema - catalog_scope = f"parent:{parent_schema}" - else: - raise RuntimeError("Codex model discovery requires a provider or parent schema.") - catalog = _fetch_codex_model_catalog( - workspace, - token, - source=catalog_source, - identifier=catalog_identifier, + temporary_catalog_dir: tempfile.TemporaryDirectory | None = None + temporary_catalog_path: Path | None = None + try: + if ( + workspace + and token + and not updating + and ( + provider + or parent_schema + or (custom_headers and not state.get("codex_static_models")) ) - validate_codex_catalog(binary, catalog) - except CodexMpsModelCatalogUnavailable: - detach_app_model_catalog() - except RuntimeError: - # A failed discovery/validation must not leave a previous workspace's - # catalog active in independently launched app servers. - _detach_app_catalog_after_failure() - raise - else: - catalog_path = _model_catalog_path(workspace, catalog_scope) - _write_model_catalog(catalog_path, catalog) - sync_app_model_catalog(catalog) - profile_doc["model_catalog_json"] = str(catalog_path) - # Codex otherwise boots on its bundled default model (e.g. gpt-5.6-sol), - # which an MPS's allowlist doesn't route, so the first request 403s. Pin - # the MPS's primary (first) target unless the user chose a model or a - # managed default already applies. - if not profile_doc.get("model") and not _tool_args_select_model(tool_args): - slugs = catalog_slugs(catalog) - if slugs: - profile_doc["model"] = slugs[0] - profile_doc = _with_custom_headers(profile_doc, custom_headers) - _run_codex( - state, - [binary, *codex_config_args(profile_doc)], - tool_args, - otel_tracing=otel_tracing, - workspace=workspace, - ) + ): + try: + fetch_kwargs: dict = {} + catalog_scope: str | None = None + if provider is not None: + fetch_kwargs["source"] = CodexCatalogSource.PROVIDER + fetch_kwargs["identifier"] = provider + catalog_scope = f"provider:{provider}" + elif parent_schema is not None: + fetch_kwargs["source"] = CodexCatalogSource.PARENT_SCHEMA + fetch_kwargs["identifier"] = parent_schema + catalog_scope = f"parent:{parent_schema}" + if custom_headers: + fetch_kwargs["request_headers"] = custom_headers + catalog = _fetch_codex_model_catalog(workspace, token, **fetch_kwargs) + validate_codex_catalog(binary, catalog) + except CodexMpsModelCatalogUnavailable: + if custom_headers: + raise + detach_app_model_catalog() + except RuntimeError: + # A failed discovery/validation must not leave a previous workspace's + # catalog active in independently launched app servers. + if not custom_headers: + _detach_app_catalog_after_failure() + raise + else: + if custom_headers: + temporary_catalog_dir = _temporary_model_catalog_dir() + temporary_catalog_path = ( + Path(temporary_catalog_dir.name) / CODEX_MODEL_CATALOG_PATH.name + ) + catalog_path = temporary_catalog_path + else: + assert catalog_scope is not None + catalog_path = _model_catalog_path(workspace, catalog_scope) + _write_model_catalog(catalog_path, catalog) + if not custom_headers: + sync_app_model_catalog(catalog) + profile_doc["model_catalog_json"] = str(catalog_path) + # Codex otherwise boots on its bundled default model (e.g. gpt-5.6-sol), + # which an MPS's allowlist doesn't route, so the first request 403s. Pin + # the MPS's primary (first) target unless the user chose a model or a + # managed default already applies. + if not profile_doc.get("model") and not _tool_args_select_model(tool_args): + slugs = catalog_slugs(catalog) + if slugs: + profile_doc["model"] = slugs[0] + profile_doc = _with_custom_headers(profile_doc, custom_headers) + _run_codex( + state, + [binary, *codex_config_args(profile_doc)], + tool_args, + otel_tracing=otel_tracing, + workspace=workspace, + wait_for_exit=temporary_catalog_path is not None, + ) + finally: + if temporary_catalog_dir is not None: + temporary_catalog_dir.cleanup() def _launch_smart_routing( @@ -1265,30 +1318,89 @@ def _launch_smart_routing( ) -> None: """Launch the Codex TUI through the smart-routing interposer.""" binary = SPEC["binary"] + temporary_catalog_dir: tempfile.TemporaryDirectory | None = None + temporary_catalog_path: Path | None = None + catalog_models: list[str] | None = None + try: + if custom_headers and not state.get("codex_static_models"): + workspace = state.get("workspace") + if isinstance(workspace, str) and workspace: + fetch_kwargs: dict = {"request_headers": custom_headers} + provider = state.get("_codex_launch_provider") + parent_schema = state.get("_codex_launch_parent_schema") + if isinstance(provider, str) and provider.strip(): + fetch_kwargs.update( + source=CodexCatalogSource.PROVIDER, + identifier=provider.strip(), + ) + elif isinstance(parent_schema, str) and parent_schema.strip(): + fetch_kwargs.update( + source=CodexCatalogSource.PARENT_SCHEMA, + identifier=parent_schema.strip(), + ) + catalog = _fetch_codex_model_catalog( + workspace, _launch_token(state, workspace), **fetch_kwargs + ) + validate_codex_catalog(binary, catalog) + catalog_models = catalog_slugs(catalog) + temporary_catalog_dir = _temporary_model_catalog_dir() + temporary_catalog_path = ( + Path(temporary_catalog_dir.name) / CODEX_MODEL_CATALOG_PATH.name + ) + _write_model_catalog(temporary_catalog_path, catalog) + + configured_model = _smart_routing_config_model(state) + if ( + custom_headers + and catalog_models + and not isinstance(state.get("codex_default_model"), str) + and configured_model + and not any( + configured_model == model + or codex_model_id(configured_model) == codex_model_id(model) + for model in catalog_models + ) + ): + # A profile model may have been pinned by an earlier ordinary launch. + # Let the header-specific catalog choose the bootstrap model unless an + # administrator explicitly supplied the startup model in state. + configured_model = None + # Prefer the launch-scoped catalog when one was discovered. Ordinary + # launches retain the configured static catalog and cached model list. + models = ( + catalog_models + if catalog_models is not None + else (custom_catalog_models() or routing_models(state)) + ) + start_model = ( + configured_model + or (codex_model_id(models[0]) if models else None) + or APP_SERVER_SMART_ROUTING_STARTING_MODEL + ) + if custom_headers: - configured_model = _smart_routing_config_model(state) - # Prefer the custom catalog if it exists. - models = custom_catalog_models() or routing_models(state) - start_model = ( - configured_model - or (codex_model_id(models[0]) if models else None) - or APP_SERVER_SMART_ROUTING_STARTING_MODEL - ) - if custom_headers: - - def routed_render_overlay(*args, **kwargs) -> dict: - return _with_custom_headers(render_overlay(*args, **kwargs), custom_headers) + def routed_render_overlay(*args, **kwargs) -> dict: + return _with_custom_headers(render_overlay(*args, **kwargs), custom_headers) - else: - routed_render_overlay = render_overlay + else: + routed_render_overlay = render_overlay - smart_routing_v2.launch_codex( - state, - tool_args, - binary=binary, - start_model=start_model, - render_overlay=routed_render_overlay, - ) + catalog_kwargs: dict = ( + {"catalog_models": catalog_models, "catalog_path": temporary_catalog_path} + if catalog_models is not None + else {} + ) + smart_routing_v2.launch_codex( + state, + tool_args, + binary=binary, + start_model=start_model, + render_overlay=routed_render_overlay, + **catalog_kwargs, + ) + finally: + if temporary_catalog_dir is not None: + temporary_catalog_dir.cleanup() def disable_smart_routing(state: dict) -> bool: diff --git a/src/ucode/launcher.py b/src/ucode/launcher.py index 382672a44..37bb7870a 100644 --- a/src/ucode/launcher.py +++ b/src/ucode/launcher.py @@ -9,12 +9,15 @@ from ucode.os_compatibility import subprocess_cross_os -def exec_or_spawn(argv: list[str]) -> None: +def exec_or_spawn(argv: list[str], *, wait_for_exit: bool = False) -> None: """Hand the terminal to ``argv``, then exit with its status. On POSIX we ``os.execvp`` — the agent process *replaces* ucode, inheriting the controlling terminal cleanly. + ``wait_for_exit`` keeps ucode alive on POSIX while the child runs. Launchers + use this only when they must clean up a launch-scoped resource afterwards. + On Windows there is no real ``exec``: ``os.execvp`` spawns a *new* process and immediately terminates the parent, so the launching shell resumes its prompt and fights the agent for the console. That produces the garbled, @@ -22,7 +25,7 @@ def exec_or_spawn(argv: list[str]) -> None: for it, and propagate its exit code — the same pattern the token-refreshing agents (gemini/opencode/copilot/pi) already use. """ - if os.name != "nt": + if os.name != "nt" and not wait_for_exit: os.execvp(argv[0], argv) return # unreachable on POSIX; keeps type-checkers happy diff --git a/src/ucode/smart_routing/v2.py b/src/ucode/smart_routing/v2.py index 50790dfac..748a2b388 100644 --- a/src/ucode/smart_routing/v2.py +++ b/src/ucode/smart_routing/v2.py @@ -628,6 +628,8 @@ def launch_codex( binary: str, start_model: str | None, render_overlay: Callable[..., dict], + catalog_models: list[str] | None = None, + catalog_path: Path | None = None, ) -> NoReturn: workspace = state.get("workspace") if not workspace: @@ -640,8 +642,13 @@ def launch_codex( ) os.environ[OAUTH_TOKEN_ENV_VAR] = _launch_token(state, workspace) - catalog_models = custom_catalog_models() - available_models = catalog_models or _cached_routing_models(state) + if catalog_models is not None: + # A launch-scoped header can select a different model-service view. Do not + # fall back to a persisted catalog from an ordinary launch in that case. + available_models = catalog_models + else: + configured_models = custom_catalog_models() + available_models = configured_models or _cached_routing_models(state) if not available_models: print_warning( "Smart routing model metadata is unavailable; automatic model switching is unavailable. " @@ -656,7 +663,9 @@ def launch_codex( custom_oauth=(custom_oauth if custom_oauth_cli_enabled(custom_oauth) else None), managed_http_headers=state.get("codex_http_headers"), ) - catalog_path = custom_catalog_path() + isolated_catalog = catalog_path is not None + if catalog_path is None: + catalog_path = custom_catalog_path() if catalog_path is not None: overlay["model_catalog_json"] = str(catalog_path) overlay["hooks"] = { @@ -673,7 +682,11 @@ def launch_codex( if not first_prompt_routing_enabled(): # Subagent-only routing needs neither the app-server nor the interposer: # the hooks ride in the CLI config, so launch the TUI directly. - exec_or_spawn([binary, *config_args, *tool_args]) + argv = [binary, *config_args, *tool_args] + if isolated_catalog: + exec_or_spawn(argv, wait_for_exit=True) + else: + exec_or_spawn(argv) app_port = _free_port() app_server_url = _loopback_websocket_url(app_port) diff --git a/tests/test_agent_codex.py b/tests/test_agent_codex.py index 58c83493f..86dce02e0 100644 --- a/tests/test_agent_codex.py +++ b/tests/test_agent_codex.py @@ -968,7 +968,13 @@ def _patch(tmp_path, monkeypatch): launches: list[list[str]] = [] monkeypatch.setattr(codex, "CODEX_CONFIG_PATH", profile_path) monkeypatch.setattr(codex, "agent_version", lambda binary: "0.134.0") - monkeypatch.setattr(codex, "exec_or_spawn", lambda argv: launches.append(argv)) + + def launch_process(argv, **kwargs): + launches.append(argv) + if kwargs.get("wait_for_exit"): + raise SystemExit(0) + + monkeypatch.setattr(codex, "exec_or_spawn", launch_process) monkeypatch.setattr( codex, "get_databricks_token", @@ -998,12 +1004,21 @@ def test_sets_oauth_token(self, tmp_path, monkeypatch): def test_custom_header_is_launch_only(self, tmp_path, monkeypatch): launches = self._patch(tmp_path, monkeypatch) value = "route://development/test" - - codex.launch( - {"workspace": WS}, - [], - options=LaunchOptions(custom_headers=(("X-Development-Route", value),)), + fetch_kwargs = {} + monkeypatch.setattr( + codex, + "_fetch_codex_model_catalog", + lambda workspace, token, **kwargs: ( + fetch_kwargs.update(kwargs) or {"models": [{"slug": "gpt-custom"}]} + ), ) + with pytest.raises(SystemExit) as exc: + codex.launch( + {"workspace": WS}, + [], + options=LaunchOptions(custom_headers=(("X-Development-Route", value),)), + ) + assert exc.value.code == 0 env_name = f"{codex.CUSTOM_HEADER_ENV_PREFIX}_{os.getpid()}_0" provider_arg = next( @@ -1014,8 +1029,81 @@ def test_custom_header_is_launch_only(self, tmp_path, monkeypatch): assert value not in provider_arg assert os.environ[env_name] == value assert "X-Development-Route" not in codex.CODEX_CONFIG_PATH.read_text(encoding="utf-8") + assert fetch_kwargs["request_headers"] == {"X-Development-Route": value} monkeypatch.delenv(env_name) + def test_custom_header_discovery_uses_private_catalog(self, tmp_path, monkeypatch): + launches = self._patch(tmp_path, monkeypatch) + value = "route://development/test" + old_catalog = {"models": [{"slug": "normal-route"}]} + custom_catalog = {"models": [{"slug": "development-route"}]} + codex.sync_app_model_catalog(old_catalog) + shared_path = tmp_path / "config.toml" + stable_path = tmp_path / "stable-provider-catalog.json" + stable_path.write_text(json.dumps(old_catalog), encoding="utf-8") + monkeypatch.setattr(codex, "_model_catalog_path", lambda workspace, scope: stable_path) + fetch_kwargs = {} + monkeypatch.setattr( + codex, + "_fetch_codex_model_catalog", + lambda workspace, token, **kwargs: fetch_kwargs.update(kwargs) or custom_catalog, + ) + observed_catalogs: list[Path] = [] + + def launch_process(argv, **kwargs): + assert kwargs == {"wait_for_exit": True} + catalog_arg = next(arg for arg in argv if arg.startswith("model_catalog_json=")) + catalog_path = Path(catalog_arg.partition("=")[2].strip('"')) + assert catalog_path.exists() + observed_catalogs.append(catalog_path) + launches.append(argv) + raise SystemExit(0) + + monkeypatch.setattr(codex, "exec_or_spawn", launch_process) + + with pytest.raises(SystemExit) as exc: + codex.launch( + {"workspace": WS, "_codex_launch_provider": "main.default.openai"}, + [], + options=LaunchOptions(custom_headers=(("X-Development-Route", value),)), + ) + + assert exc.value.code == 0 + assert fetch_kwargs["request_headers"] == {"X-Development-Route": value} + assert observed_catalogs + assert not observed_catalogs[0].exists() + assert not observed_catalogs[0].parent.exists() + assert json.loads(codex.CODEX_MODEL_CATALOG_PATH.read_text()) == old_catalog + assert read_toml_safe(shared_path)["model_catalog_json"] == str( + codex.CODEX_MODEL_CATALOG_PATH + ) + assert json.loads(stable_path.read_text()) == old_catalog + + def test_custom_header_discovery_failure_preserves_catalog(self, tmp_path, monkeypatch): + launches = self._patch(tmp_path, monkeypatch) + old_catalog = {"models": [{"slug": "normal-route"}]} + codex.sync_app_model_catalog(old_catalog) + shared_path = tmp_path / "config.toml" + app_catalog_before = codex.CODEX_MODEL_CATALOG_PATH.read_bytes() + shared_before = shared_path.read_text(encoding="utf-8") + monkeypatch.setattr( + codex, + "_fetch_codex_model_catalog", + lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("HTTP 403 Forbidden")), + ) + + with pytest.raises(RuntimeError, match="HTTP 403 Forbidden"): + codex.launch( + {"workspace": WS, "_codex_launch_provider": "main.default.openai"}, + [], + options=LaunchOptions(custom_headers=(("X-Development-Route", "route://test"),)), + ) + + assert launches == [] + assert codex.CODEX_MODEL_CATALOG_PATH.read_bytes() == app_catalog_before + assert shared_path.read_text(encoding="utf-8") == shared_before + assert not list(codex.CODEX_MODEL_CATALOG_PATH.parent.glob(".*-launch-*")) + @pytest.mark.parametrize("header_table", ["http_headers", "env_http_headers"]) def test_custom_header_rejects_managed_header(self, tmp_path, monkeypatch, header_table): self._patch(tmp_path, monkeypatch) diff --git a/tests/test_codex_smart_routing_v2.py b/tests/test_codex_smart_routing_v2.py index adb7b776f..4febc76a7 100644 --- a/tests/test_codex_smart_routing_v2.py +++ b/tests/test_codex_smart_routing_v2.py @@ -77,22 +77,39 @@ def launch_v2(state, tool_args, **kwargs): ) ] - def test_codex_smart_routing_preserves_custom_header(self, monkeypatch): + @pytest.mark.parametrize("static_models", [None, ["gpt-static"]]) + def test_codex_smart_routing_preserves_custom_header( + self, tmp_path, monkeypatch, static_models + ): captured = {} + requested = {} monkeypatch.setenv(v2.ENABLE_SMART_ROUTING_ENV_VAR, "1") - monkeypatch.setattr(codex, "_smart_routing_config_model", lambda state: "gpt-start") + monkeypatch.setattr(codex, "_smart_routing_config_model", lambda state: "gpt-old") monkeypatch.setattr(codex, "codex_managed_config_path", lambda: None) + monkeypatch.setattr(codex, "CODEX_MODEL_CATALOG_PATH", tmp_path / "app-catalog.json") + monkeypatch.setattr(codex, "get_databricks_token", lambda *_args, **_kwargs: "token") + monkeypatch.setattr( + codex, + "_fetch_codex_model_catalog", + lambda *_args, **kwargs: ( + requested.update(kwargs) or {"models": [{"slug": "gpt-start"}]} + ), + ) + monkeypatch.setattr(codex, "validate_codex_catalog", lambda *_args: None) def launch_v2(state, tool_args, **kwargs): captured.update(kwargs) + if not static_models: + assert json.loads(kwargs["catalog_path"].read_text())["models"] == [ + {"slug": "gpt-start"} + ] raise SystemExit(0) monkeypatch.setattr(v2, "launch_codex", launch_v2) custom_headers = {"X-Development-Route": "route://development/test"} - with pytest.raises(SystemExit): codex.launch( - {"workspace": WS}, + {"workspace": WS, "codex_static_models": static_models}, [], options=LaunchOptions( launch_smart_routing=True, @@ -105,6 +122,14 @@ def launch_v2(state, tool_args, **kwargs): env_name = headers["X-Development-Route"] assert os.environ[env_name] == custom_headers["X-Development-Route"] monkeypatch.delenv(env_name) + if static_models: + assert requested == {} + assert "catalog_models" not in captured + else: + assert requested["request_headers"] == custom_headers + assert captured["start_model"] == "gpt-start" + assert captured["catalog_models"] == ["gpt-start"] + assert not captured["catalog_path"].exists() @pytest.mark.parametrize( "tool_args", diff --git a/tests/test_launcher.py b/tests/test_launcher.py index 1c1caa5fe..57c5926c2 100644 --- a/tests/test_launcher.py +++ b/tests/test_launcher.py @@ -21,18 +21,21 @@ def test_posix_uses_execvp(self): execvp.assert_called_once_with("claude", ["claude", "--settings", "x"]) popen.assert_not_called() - def test_windows_spawns_and_waits(self): + @pytest.mark.parametrize("platform, wait_for_exit", [("nt", False), ("posix", True)]) + def test_spawns_and_waits(self, platform, wait_for_exit): # On Windows there is no real exec; we must spawn + wait so the parent # shell does not resume and corrupt the terminal (issue #173). proc = MagicMock() proc.wait.return_value = 0 with ( - patch.object(launcher.os, "name", "nt"), + patch.object(launcher.os, "name", platform), patch.object(launcher.os, "execvp") as execvp, patch.object(launcher.subprocess_cross_os, "popen", return_value=proc) as popen, ): with pytest.raises(SystemExit) as exc: - launcher.exec_or_spawn(["claude.exe", "--settings", "x"]) + launcher.exec_or_spawn( + ["claude.exe", "--settings", "x"], wait_for_exit=wait_for_exit + ) execvp.assert_not_called() popen.assert_called_once_with(["claude.exe", "--settings", "x"]) proc.wait.assert_called_once() From 42207218a3386d824af33a149885eff6d877b5d6 Mon Sep 17 00:00:00 2001 From: Andy Xu Date: Wed, 7 Oct 2026 21:17:19 +0000 Subject: [PATCH 5/6] Scope custom headers across UG workspace requests and routing --- README.md | 7 +- src/ucode/agents/__init__.py | 8 +- src/ucode/cli.py | 97 +++++++++---------- src/ucode/databricks.py | 82 +++++++++++++--- src/ucode/managed_config.py | 36 +++++-- src/ucode/request_headers.py | 145 +++++++++++++++++++++++++++++ src/ucode/smart_routing/routing.py | 33 ++++++- src/ucode/state.py | 25 +++-- tests/README.md | 2 + tests/integration/README.md | 3 +- tests/test_cli.py | 85 +++++++++++++++-- tests/test_codex_routing.py | 36 ++++++- tests/test_databricks.py | 127 +++++++++++++++++++++++++ tests/test_managed_config.py | 57 ++++++++++++ 14 files changed, 640 insertions(+), 103 deletions(-) create mode 100644 src/ucode/request_headers.py diff --git a/README.md b/README.md index 61f39fc4f..24830eb2f 100644 --- a/README.md +++ b/README.md @@ -229,9 +229,10 @@ ug skills remove --location main.default --via mcp | `ug revert` | Clear saved state and restore backed-up config files | | `ug upgrade` | Upgrade Unity Gateway | -`ug codex --header 'X-Development-Route: test-target'` adds a repeatable, launch-only -request header for Codex. Use it only for non-secret development routing values: headers are -passed through environment variables and are not written to persistent Codex configuration. +`ug --header 'X-Development-Route: test-target' usage` adds a repeatable header to workspace +requests for that invocation, including prelaunch discovery and configuration requests. +`ug codex --header 'X-Development-Route: test-target'` also passes it to Codex and its routing +helpers for that launch. Use non-secret values; headers are not saved in Codex configuration. Databricks AI Tools are installed only by `ug configure`, never by agent launch commands. Use `--enable-databricks-ai-tools` or `--disable-databricks-ai-tools` diff --git a/src/ucode/agents/__init__.py b/src/ucode/agents/__init__.py index 38e712ff0..eacd81adb 100644 --- a/src/ucode/agents/__init__.py +++ b/src/ucode/agents/__init__.py @@ -399,8 +399,12 @@ def resolve_gemini_provider_model( def validate_custom_headers(tool: str, state: dict, headers: dict[str, str]) -> None: - if tool == "codex" and headers: - codex.validate_custom_headers(state, headers) + if not headers: + return + validator = getattr(_MODULES[tool], "validate_custom_headers", None) + if validator is None: + raise RuntimeError(f"--header is not supported for {TOOL_SPECS[tool]['display']} launches.") + validator(state, headers) def configure_tool( diff --git a/src/ucode/cli.py b/src/ucode/cli.py index 45b9216be..f04881e09 100644 --- a/src/ucode/cli.py +++ b/src/ucode/cli.py @@ -4,11 +4,10 @@ from __future__ import annotations import os -import re import shutil import subprocess from collections.abc import Iterator -from contextlib import contextmanager +from contextlib import ExitStack, contextmanager from enum import StrEnum from importlib import metadata from typing import Annotated, Any @@ -127,6 +126,14 @@ revert_mcp_configs, ) from ucode.os_compatibility import subprocess_cross_os +from ucode.request_headers import ( + custom_header_environment, + custom_header_scope, + get_custom_headers, +) +from ucode.request_headers import ( + parse_custom_headers as _parse_custom_headers, +) from ucode.skills_download import ( configure_location_skills_download_command, configure_selected_skills_download_command, @@ -145,6 +152,7 @@ set_session_environment, ) from ucode.state import ( + LAUNCH_DISCOVERY_OVERLAY_KEY, clear_state, get_provider_service, load_state, @@ -499,6 +507,7 @@ def configure_shared_state( Only the local profile resolution and the shared state assembly still run; the saved model lists are preserved. """ + request_headers = request_headers or get_custom_headers() workspace = normalize_workspace_url(workspace) prior_state = load_state() previous_workspace = prior_state.get("workspace") @@ -675,6 +684,18 @@ def configure_shared_state( if oss_models: opencode_models["oss"] = oss_models + if request_headers: + state[LAUNCH_DISCOVERY_OVERLAY_KEY] = { + key: state.get(key) + for key in ( + "claude_models", + "codex_models", + "gemini_models", + "oss_models", + "opencode_models", + "web_search_model", + ) + } if skip_model_discovery: # Don't clobber any previously-discovered Databricks model lists; provider # mode just doesn't refresh or use them. Persist the web-search model so @@ -2556,55 +2577,6 @@ def _smart_routing_launch_shape(tool: str, tool_args: list[str], explicit_prompt return tool == "claude" and tool_args[0].startswith("-") -_HTTP_HEADER_NAME_PATTERN = re.compile(r"[!#$%&'*+\-.^_`|~0-9A-Za-z]+") -_PROTECTED_CUSTOM_HEADER_NAMES = frozenset( - { - "authorization", - "connection", - "content-length", - "cookie", - "databricks-model-provider-service", - "databricks-model-service-parent-schema", - "databricks-smart-router-recipe", - "host", - "keep-alive", - "proxy-authenticate", - "proxy-authorization", - "te", - "trailer", - "transfer-encoding", - "upgrade", - "user-agent", - "x-api-key", - "x-databricks-ai-gateway-token", - "x-databricks-use-coding-agent-mode", - } -) - - -def _parse_custom_headers(values: list[str] | None) -> dict[str, str]: - """Parse repeatable ``--header 'Name: value'`` options.""" - parsed: dict[str, tuple[str, str]] = {} - for item in values or []: - name, separator, value = item.partition(":") - name = name.strip() - if not separator or _HTTP_HEADER_NAME_PATTERN.fullmatch(name) is None: - raise RuntimeError("--header must use the format `Name: value` with a valid name.") - value = value.strip() - if any( - ord(character) < 32 or ord(character) == 127 or character in "\u0085\u2028\u2029" - for character in value - ): - raise RuntimeError( - "--header values cannot contain control characters or line separators." - ) - normalized_name = name.casefold() - if normalized_name in _PROTECTED_CUSTOM_HEADER_NAMES: - raise RuntimeError(f"--header cannot override protected header '{name}'.") - parsed[normalized_name] = (name, value) - return dict(parsed.values()) - - def _launch_options( tool: str, tool_args: list[str], @@ -2657,11 +2629,15 @@ def _launch_tool( custom_oauth: CustomOAuthConfig | None = None, headers: list[str] | None = None, ) -> None: + header_scope = ExitStack() try: tool = normalize_tool(tool_name) if not custom_oauth_cli_enabled(custom_oauth): os.environ.pop(CUSTOM_OAUTH_CLI_ENV_VAR, None) - custom_headers = _parse_custom_headers(headers) + custom_headers = _parse_custom_headers( + [f"{name}: {value}" for name, value in get_custom_headers().items()] + (headers or []) + ) + header_scope.enter_context(custom_header_scope(custom_headers)) # Before any status print: a stdio-protocol subcommand owns stdout, so # every ug line from here on must go to stderr instead. if _child_owns_stdout(tool, ctx.args): @@ -3014,8 +2990,11 @@ def _launch_tool( custom_headers=custom_headers, ) print_success(f"Starting {TOOL_SPECS[tool]['display']}") - with _smart_routing_v2_flag( - True if managed_smart_routing_enabled and smart_routing_enabled else None + with ( + _smart_routing_v2_flag( + True if managed_smart_routing_enabled and smart_routing_enabled else None + ), + custom_header_environment(state["workspace"], custom_headers), ): launch_agent(tool, state, ctx.args, options=launch_options) except RuntimeError as exc: @@ -3024,6 +3003,8 @@ def _launch_tool( except KeyboardInterrupt: print_err("Interrupted.") raise typer.Exit(130) from None + finally: + header_scope.close() # Launch-only escape hatch for managed/headless launchers (e.g. omnigent) that @@ -3061,7 +3042,7 @@ def _launch_tool( list[str] | None, typer.Option( "--header", - help="Add an HTTP header to AI Gateway requests as `Name: value`; repeatable. " + help="Add an HTTP header to workspace requests as `Name: value`; repeatable. " "Pass before any `--` separator. Credentials and transport headers are not allowed.", ), ] @@ -3114,8 +3095,14 @@ def default( ] = False, skip_preflight: SkipPreflightOption = False, workspace: WorkspaceOption = None, + header: CustomHeaderOption = None, ) -> None: """Configure and launch coding agents through Databricks AI Gateway.""" + try: + ctx.with_resource(custom_header_scope(_parse_custom_headers(header))) + except RuntimeError as exc: + print_err(str(exc)) + raise typer.Exit(1) from None if ctx.invoked_subcommand is not None: return set_dry_run(dry_run) diff --git a/src/ucode/databricks.py b/src/ucode/databricks.py index 21ac302a6..cfa978a63 100644 --- a/src/ucode/databricks.py +++ b/src/ucode/databricks.py @@ -37,6 +37,8 @@ MODEL_SERVICE_PARENT_SCHEMA_HEADER, ) from ucode.os_compatibility import subprocess_cross_os +from ucode.request_headers import get_custom_headers +from ucode.request_headers import urlopen as _open_http_request from ucode.telemetry import ug_version from ucode.ui import ( err_console, @@ -73,6 +75,16 @@ _HTTP_GET_RETRY_MAX_SECONDS = 5.0 _HTTP_GET_RETRY_AFTER_JITTER_SECONDS = 0.25 _ANTHROPIC_MODEL_DISCOVERY_SETUP_MAX_RETRIES = 2 +_TRANSPORT_HEADER_NAMES = frozenset( + { + "accept", + "authorization", + "content-type", + "user-agent", + MODEL_PROVIDER_SERVICE_HEADER.casefold(), + MODEL_SERVICE_PARENT_SCHEMA_HEADER.casefold(), + } +) @dataclass(frozen=True) @@ -286,6 +298,38 @@ def clear_workspace_org_id_cache() -> None: _WORKSPACE_ORG_IDS.clear() +def _merge_http_headers( + mandatory: dict[str, str], + explicit: dict[str, str] | None, + custom: dict[str, str], +) -> dict[str, str]: + """Merge request headers while treating names case-insensitively. + + Invocation-scoped headers are the lowest-precedence layer: an endpoint's explicit selector + (for example, a Model Provider Service header) wins over them, and transport-owned values such + as ``Authorization``/``Accept``/``Content-Type`` always win. Removing a case-insensitive + duplicate before adding the next layer also avoids sending both spellings of one header. + """ + merged: dict[str, str] = {} + for layer in (custom, explicit or {}, mandatory): + for name, value in layer.items(): + folded = name.casefold() + for existing in list(merged): + if existing.casefold() == folded: + del merged[existing] + merged[name] = value + return merged + + +def _has_custom_http_headers( + custom: dict[str, str], explicit: dict[str, str] | None = None +) -> bool: + """Whether a request carries invocation or caller-supplied routing headers.""" + return bool(custom) or any( + name.casefold() not in _TRANSPORT_HEADER_NAMES for name in (explicit or {}) + ) + + def _http_get_bytes( url: str, token: str, @@ -305,12 +349,17 @@ def _http_get_bytes( if max_retries < 0: raise ValueError("max_retries must be non-negative") - request_headers = {"Authorization": f"Bearer {token}"} - request_headers.update(headers or {}) + custom_headers = get_custom_headers() + request_headers = _merge_http_headers( + {"Authorization": f"Bearer {token}"}, headers, custom_headers + ) request = urllib_request.Request(url, headers=request_headers) + scoped_headers = _has_custom_http_headers(custom_headers, headers) for attempt in range(max_retries + 1): try: - with urllib_request.urlopen(request, timeout=timeout) as response: + with _open_http_request( + request, timeout=timeout, scoped_headers=scoped_headers + ) as response: body = response.read() _capture_org_id(url, getattr(response, "headers", None)) _debug(f"GET {url}", f"HTTP 200, {len(body)} bytes") @@ -408,12 +457,17 @@ def _http_send_json( empty body there is the expected result, not a decode failure. """ body_bytes = json.dumps(payload).encode("utf-8") if payload is not None else None - headers = {"Authorization": f"Bearer {token}", "Accept": "application/json"} + mandatory_headers = {"Authorization": f"Bearer {token}", "Accept": "application/json"} if body_bytes is not None: - headers["Content-Type"] = "application/json" + mandatory_headers["Content-Type"] = "application/json" + custom_headers = get_custom_headers() + headers = _merge_http_headers(mandatory_headers, None, custom_headers) request = urllib_request.Request(url, data=body_bytes, method=method, headers=headers) + scoped_headers = bool(custom_headers) try: - with urllib_request.urlopen(request, timeout=timeout) as response: + with _open_http_request( + request, timeout=timeout, scoped_headers=scoped_headers + ) as response: body = response.read().decode("utf-8") _debug(f"{method} {url}", f"HTTP {response.status}, {len(body)} bytes") if _debug_enabled(): @@ -1814,7 +1868,7 @@ def has_cached_model_provider_services(workspace: str, parent: str | None = None it deserves one, but repeating it per agent on an instant cache hit is just noise. Takes ``parent`` for the same reason the cache is keyed on it — a scoped listing is a separate entry. """ - return (workspace, parent or "") in _MODEL_PROVIDER_SERVICES_CACHE + return not get_custom_headers() and (workspace, parent or "") in _MODEL_PROVIDER_SERVICES_CACHE def list_model_services( @@ -1839,7 +1893,10 @@ def list_model_services( A successful result is memoized per workspace for the life of the process; pass ``use_cache=False`` to force a fresh walk. """ - if use_cache: + # A custom-header route can select a different catalog than the ordinary workspace request. + # Keep its result invocation-scoped: never read or overwrite the process-wide headerless cache. + use_process_cache = use_cache and not get_custom_headers() + if use_process_cache: cached = _MODEL_SERVICES_CACHE.get(workspace) if cached is not None: return list(cached), None @@ -1879,7 +1936,7 @@ def list_model_services( deduped = sorted(set(ids)) if deduped: - if use_cache: + if use_process_cache: _MODEL_SERVICES_CACHE[workspace] = list(deduped) return deduped, None return [], last_reason or "model-services listing returned no models" @@ -2403,7 +2460,10 @@ def list_model_provider_services( # vice versa) — a service that plainly exists would look absent, the same failure pagination was # added to fix. cache_key = (workspace, parent or "") - if use_cache: + # A custom-header route can select a different provider listing. Keep it out of the ordinary + # process cache so a later headerless launch cannot reuse the scoped response (or vice versa). + use_process_cache = use_cache and not get_custom_headers() + if use_process_cache: cached = _MODEL_PROVIDER_SERVICES_CACHE.get(cache_key) if cached is not None: # A fresh list of fresh dicts each time: callers treat the result as theirs (the wizard @@ -2445,7 +2505,7 @@ def list_model_provider_services( if not services and last_reason is not None: return [], last_reason services.sort(key=lambda s: s["name"]) - if use_cache: + if use_process_cache: _MODEL_PROVIDER_SERVICES_CACHE[cache_key] = [dict(service) for service in services] return services, None diff --git a/src/ucode/managed_config.py b/src/ucode/managed_config.py index 6e357e696..5937ef48d 100644 --- a/src/ucode/managed_config.py +++ b/src/ucode/managed_config.py @@ -12,7 +12,9 @@ re-fetching), falling back to the persisted copy when the read fails. There is deliberately one file: the workspace is the source of truth, so the pulled copy lives in -``managed-config.json`` and a launch re-reads it from there. +``managed-config.json`` and a launch re-reads it from there. A launch carrying invocation-scoped +custom headers is a separate workspace view: its read bypasses and never changes this ordinary +cache, and a failed scoped read cannot fall back to it. :func:`refresh_managed_config` is the launch path's entry point. It is called before model discovery, because the manifest decides whether that discovery is needed at all; the launch path then hands the @@ -37,6 +39,7 @@ fetch_model_recommendation, get_databricks_token, ) +from ucode.request_headers import get_custom_headers from ucode.time_utils import parse_update_time from ucode.ui import console, print_warning @@ -750,12 +753,15 @@ def refresh_managed_config(state: dict, *, force_refresh: bool = False) -> Manag applied watermark. The manifest is None when the workspace has no managed config, the normal case for a workspace whose admin hasn't published one. - A failed fetch never blocks the launch: an unreachable control plane shouldn't stop someone from - coding. Instead it falls back to the last config persisted for this workspace, so the admin's - most recent known policy still applies; only when there is no persisted config either does the - launch fall through to the developer's own settings. ``FEATURE_DISABLED`` is the exception — it - is an authoritative "off", not a transient failure, so it drops the cache rather than falling - back (see below). + A failed fetch ordinarily never blocks the launch: an unreachable control plane shouldn't stop + someone from coding. Instead it falls back to the last config persisted for this workspace, so + the admin's most recent known policy still applies; only when there is no persisted config either + does the launch fall through to the developer's own settings. ``FEATURE_DISABLED`` is the + exception — it is an authoritative "off", not a transient failure, so it drops the cache rather + than falling back (see below). When invocation-scoped custom headers are active, the response is + a separate in-memory view: the persistent cache is bypassed and untouched, and an auth/read + failure raises instead of falling back to that unrelated view; ``FEATURE_DISABLED`` remains an + authoritative scoped result without clearing the ordinary cache. ``coding_agent_config_feature_disabled`` is True whenever the gateway returned ``FEATURE_DISABLED`` — the coding-agent-configs feature isn't enabled server-side, so callers suppress the ``ucode @@ -766,21 +772,35 @@ def refresh_managed_config(state: dict, *, force_refresh: bool = False) -> Manag workspace = state.get("workspace") if not workspace: return ManagedConfigResult(None, False) - if not force_refresh: + custom_headers_active = bool(get_custom_headers()) + if not force_refresh and not custom_headers_active: cached = _cached_result_if_fresh(workspace) if cached is not None: return cached try: token = get_databricks_token(workspace, state.get("profile")) except RuntimeError as exc: + if custom_headers_active: + raise RuntimeError(f"Could not fetch managed configuration: {exc}") from exc return ManagedConfigResult(_persisted_fallback(workspace, str(exc)), False) raw, reason = get_managed_config(workspace, token) if reason is not None: + if custom_headers_active: + if _is_feature_disabled(reason): + return ManagedConfigResult(None, True) + raise RuntimeError(f"Could not fetch managed configuration: {reason}") if _is_feature_disabled(reason): save_managed_state(workspace, {}, outcome=_OUTCOME_FEATURE_DISABLED) return ManagedConfigResult(None, True) fallback = _persisted_fallback(workspace, reason, refused=_is_permission_denied(reason)) return ManagedConfigResult(fallback, False) + if custom_headers_active: + # A scoped route may expose a different admin policy. Keep it in memory for this + # invocation and leave the workspace's ordinary persistent cache untouched. + return ManagedConfigResult( + normalize_managed_config(raw) if raw is not None else None, + False, + ) if raw is None: # Record that this workspace has no config, rather than leaving an earlier one on disk: # the file doubles as the fallback above, so a removed policy would otherwise come back diff --git a/src/ucode/request_headers.py b/src/ucode/request_headers.py new file mode 100644 index 000000000..6cf4e3d7a --- /dev/null +++ b/src/ucode/request_headers.py @@ -0,0 +1,145 @@ +"""Validated, invocation-scoped custom headers for workspace requests.""" + +from __future__ import annotations + +import json +import os +import re +from collections.abc import Iterator +from contextlib import contextmanager +from contextvars import ContextVar +from urllib import error as urllib_error +from urllib import request as urllib_request +from urllib.parse import urljoin, urlsplit + +_CUSTOM_HEADERS: ContextVar[dict[str, str] | None] = ContextVar("custom_headers", default=None) +CUSTOM_HEADERS_ENV = "UCODE_CUSTOM_HEADERS" + + +_HTTP_HEADER_NAME_PATTERN = re.compile(r"[!#$%&'*+\-.^_`|~0-9A-Za-z]+") +_PROTECTED_CUSTOM_HEADER_NAMES = frozenset( + { + "accept", + "authorization", + "connection", + "content-length", + "content-type", + "cookie", + "databricks-model-provider-service", + "databricks-model-service-parent-schema", + "databricks-smart-router-recipe", + "host", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailer", + "transfer-encoding", + "upgrade", + "user-agent", + "x-api-key", + "x-databricks-ai-gateway-token", + "x-databricks-use-coding-agent-mode", + } +) + + +def parse_custom_headers(values: list[str] | None) -> dict[str, str]: + """Parse repeatable ``--header 'Name: value'`` options.""" + parsed: dict[str, tuple[str, str]] = {} + for item in values or []: + name, separator, value = item.partition(":") + name = name.strip() + if not separator or _HTTP_HEADER_NAME_PATTERN.fullmatch(name) is None: + raise RuntimeError("--header must use the format `Name: value` with a valid name.") + value = value.strip() + if any( + ord(character) < 32 or ord(character) == 127 or character in "\u0085\u2028\u2029" + for character in value + ): + raise RuntimeError( + "--header values cannot contain control characters or line separators." + ) + normalized_name = name.casefold() + if normalized_name in _PROTECTED_CUSTOM_HEADER_NAMES: + raise RuntimeError(f"--header cannot override protected header '{name}'.") + parsed[normalized_name] = (name, value) + return dict(parsed.values()) + + +def get_custom_headers() -> dict[str, str]: + return dict(_CUSTOM_HEADERS.get() or {}) + + +@contextmanager +def custom_header_scope(headers: dict[str, str]) -> Iterator[None]: + token = _CUSTOM_HEADERS.set(dict(headers)) + try: + yield + finally: + _CUSTOM_HEADERS.reset(token) + + +@contextmanager +def custom_header_environment(workspace: str, headers: dict[str, str]) -> Iterator[None]: + """Pass the launch's headers only to helpers that explicitly opt in.""" + previous = os.environ.get(CUSTOM_HEADERS_ENV) + if headers: + os.environ[CUSTOM_HEADERS_ENV] = json.dumps({"workspace": workspace, "headers": headers}) + else: + os.environ.pop(CUSTOM_HEADERS_ENV, None) + try: + yield + finally: + if previous is None: + os.environ.pop(CUSTOM_HEADERS_ENV, None) + else: + os.environ[CUSTOM_HEADERS_ENV] = previous + + +def inherited_custom_headers(workspace: str) -> dict[str, str]: + """Read launch headers in a routing helper for the same workspace only.""" + raw = os.environ.get(CUSTOM_HEADERS_ENV) + if not raw: + return {} + try: + payload = json.loads(raw) + if not isinstance(payload, dict) or not isinstance(payload.get("workspace"), str): + raise ValueError + if _origin(workspace) != _origin(payload["workspace"]): + return {} + headers = payload["headers"] + if not isinstance(headers, dict) or not all( + isinstance(name, str) and isinstance(value, str) for name, value in headers.items() + ): + raise ValueError + except (ValueError, TypeError, KeyError): + raise RuntimeError("Invalid custom-header context inherited from ug.") from None + return parse_custom_headers([f"{name}: {value}" for name, value in headers.items()]) + + +def _origin(url: str) -> tuple[str, str | None, int | None]: + parsed = urlsplit(url) + port = parsed.port + if port is None: + port = {"http": 80, "https": 443}.get(parsed.scheme) + return parsed.scheme, parsed.hostname, port + + +class _ScopedHeaderRedirectHandler(urllib_request.HTTPRedirectHandler): + def redirect_request(self, req, fp, code, msg, headers, newurl): + target = urljoin(req.full_url, newurl) + if _origin(req.full_url) != _origin(target): + raise urllib_error.HTTPError( + req.full_url, code, "Cross-origin redirect blocked for custom headers", headers, fp + ) + return super().redirect_request(req, fp, code, msg, headers, target) + + +def urlopen(request: urllib_request.Request, *, timeout: float, scoped_headers: bool = False): + """Preserve custom headers on same-origin redirects; reject cross-origin redirects.""" + if scoped_headers: + return urllib_request.build_opener(_ScopedHeaderRedirectHandler()).open( + request, timeout=timeout + ) + return urllib_request.urlopen(request, timeout=timeout) diff --git a/src/ucode/smart_routing/routing.py b/src/ucode/smart_routing/routing.py index 3d11f3d33..5f2aa7dd8 100644 --- a/src/ucode/smart_routing/routing.py +++ b/src/ucode/smart_routing/routing.py @@ -23,6 +23,9 @@ from pathlib import Path from typing import Any +from ucode import request_headers +from ucode.request_headers import get_custom_headers, inherited_custom_headers + ROUTER_NAME = "task_v3" ROUTER_NAME_ENV_VAR = "SMART_ROUTER_NAME" ROUTING_PATH = "/ai-gateway/routing/v1/routes:select" @@ -222,17 +225,37 @@ def select_route( "task": {"prompt": task}, "route_selector": {"router_name": router_name}, } + # Direct launches use the context-local value. Hook subprocesses and the + # PTY/interposer worker use the workspace-bound envelope explicitly here; + # unrelated commands do not inherit it because the launcher scopes that + # environment only around the child agent. + custom_headers = get_custom_headers() or inherited_custom_headers(workspace) + # Authentication and the request content type are owned by this client. + # The parser rejects protected names, but keep the boundary safe for any + # embedding caller that enters a context manually. + headers = { + name: value + for name, value in custom_headers.items() + if name.casefold() not in {"authorization", "content-type"} + } + headers.update( + { + "Authorization": f"Bearer {token}", + "Content-Type": "application/json", + } + ) request = urllib.request.Request( workspace.rstrip("/") + ROUTING_PATH, data=json.dumps(body).encode("utf-8"), - headers={ - "Authorization": f"Bearer {token}", - "Content-Type": "application/json", - }, + headers=headers, method="POST", ) try: - with urllib.request.urlopen(request, timeout=timeout) as response: + with request_headers.urlopen( + request, + timeout=timeout, + scoped_headers=bool(custom_headers), + ) as response: payload = json.loads(response.read().decode("utf-8")) except urllib.error.HTTPError as exc: detail = "" diff --git a/src/ucode/state.py b/src/ucode/state.py index 4dd5d5fa6..db47ffa86 100644 --- a/src/ucode/state.py +++ b/src/ucode/state.py @@ -24,6 +24,7 @@ # Present only in memory: the layered values render the agent settings files, while `save_state` # restores what's under it so `state.json` keeps recording the developer's own configuration. MANAGED_OVERLAY_KEY = "_managed_overlay" +LAUNCH_DISCOVERY_OVERLAY_KEY = "_launch_discovery_overlay" AUTH_COMMAND_TIMEOUT_MS = 5000 AUTH_REFRESH_INTERVAL_MS = 900_000 @@ -76,21 +77,25 @@ def save_state(state: dict) -> None: def _without_managed_overlay(state: dict) -> dict: - """Return ``state`` with managed-config values swapped back for the developer's own. + """Return ``state`` with transient overlays swapped back for the developer's own values. Returns a new dict and leaves ``state`` untouched, so the caller keeps the layered values it needs for rendering and repeated saves stay idempotent. """ - overlay = state.get(MANAGED_OVERLAY_KEY) - if not isinstance(overlay, dict): + overlay_keys = (MANAGED_OVERLAY_KEY, LAUNCH_DISCOVERY_OVERLAY_KEY) + if not any(isinstance(state.get(key), dict) for key in overlay_keys): return state - persisted = {key: value for key, value in state.items() if key != MANAGED_OVERLAY_KEY} - for key, value in overlay.items(): - # A key the developer never set is dropped rather than persisted as None. - if value is None: - persisted.pop(key, None) - else: - persisted[key] = value + persisted = {key: value for key, value in state.items() if key not in overlay_keys} + # Unwind managed settings first, then any launch-specific discovery beneath them. + for overlay_key in overlay_keys: + overlay = state.get(overlay_key) + if not isinstance(overlay, dict): + continue + for key, value in overlay.items(): + if value is None: + persisted.pop(key, None) + else: + persisted[key] = value return persisted diff --git a/tests/README.md b/tests/README.md index 339803a8f..51c9d98f0 100644 --- a/tests/README.md +++ b/tests/README.md @@ -45,6 +45,8 @@ quoted executable paths and replacement of legacy `ucode` routing/web-search hel Custom request headers have component coverage in `test_cli.py`, `test_databricks.py`, `test_agent_codex.py`, and `test_codex_smart_routing_v2.py`: parsing, launch-only values and catalogs, gateway discovery, and administrator-header collisions before discovery. +Global-option tests cover `usage`/`recommendModel`, local-option precedence, and cleanup on +success or failure; request/config tests cover redirect handling and cache isolation. `test_launcher.py` checks waiting for the child before cleaning up temporary catalogs. Live custom-header journeys are not covered. diff --git a/tests/integration/README.md b/tests/integration/README.md index b4cbc82f2..420f0a9c0 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -18,7 +18,8 @@ The existing unit tests keep their fixtures. Integration has an independent pytest configuration and uses `--confcutdir` so those fixtures cannot leak in. It is not collected by the default `uv run pytest` command. -Custom `--header` parsing, delivery, discovery isolation, and collision rejection have component +Custom `--header` parsing, invocation cleanup, usage delivery, discovery/cache isolation, and +collision rejection have component coverage listed in `../README.md`; live custom-header journeys are not covered. `TestChildStdoutLaunch` in `../test_cli.py` covers clean Claude print-mode and diff --git a/tests/test_cli.py b/tests/test_cli.py index 4b56619cf..c283338ff 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -910,6 +910,65 @@ def test_codex_headers_are_forwarded(self): "X-Second: two:three", ] + def test_global_headers_reach_prelaunch_requests_and_codex(self): + from ucode.request_headers import get_custom_headers, inherited_custom_headers + + with _launch_policy_patches(None) as calls: + calls["shared"].side_effect = lambda *args, **kwargs: ( + seen.append(get_custom_headers()) or calls["state"] + ) + calls["launch"].side_effect = lambda *args, **kwargs: seen.append( + inherited_custom_headers(MINIMAL_STATE["workspace"]) + ) + seen = [] + result = runner.invoke( + app, + ["--header", "X-Test: global", "codex", "--header", "x-test: local"], + ) + + assert result.exit_code == 0, result.output + assert seen == [{"x-test": "local"}, {"x-test": "local"}] + assert calls["launch"].call_args.kwargs["options"].custom_headers == (("x-test", "local"),) + assert get_custom_headers() == {} + assert inherited_custom_headers(MINIMAL_STATE["workspace"]) == {} + + def test_usage_header_reaches_recommend_model_and_is_cleared(self): + from ucode.request_headers import get_custom_headers + + response = MagicMock() + response.__enter__.return_value = response + response.read.return_value = b'{"current_spend": 12, "effective_threshold": 100}' + response.status = 200 + with ( + patch("ucode.cli.install_databricks_cli"), + patch("ucode.usage.load_state", return_value=MINIMAL_STATE), + patch("ucode.usage.apply_pat_environment"), + patch("ucode.usage.ensure_databricks_auth"), + patch("ucode.usage.get_databricks_token", return_value="token"), + patch("urllib.request.urlopen", return_value=response) as send, + patch("urllib.request.build_opener") as opener, + ): + opener.return_value.open = send + result = runner.invoke(app, ["--header", "X-Test: temporary", "usage"]) + assert result.exit_code == 0, result.output + request = send.call_args.args[0] + assert request.full_url.endswith("coding-agent-configs:recommendModel") + assert dict(request.header_items())["X-test"] == "temporary" + assert request.get_header("Authorization") == "Bearer token" + assert get_custom_headers() == {} + + result = runner.invoke(app, ["usage"]) + assert result.exit_code == 0, result.output + assert "X-test" not in dict(send.call_args.args[0].header_items()) + + def test_global_header_context_is_cleared_on_failure(self): + from ucode.request_headers import get_custom_headers + + with patch("ucode.cli.install_databricks_cli", side_effect=RuntimeError("test failure")): + result = runner.invoke(app, ["--header", "X-Test: temporary", "usage"]) + assert result.exit_code == 1 + assert get_custom_headers() == {} + def test_codex_admin_header_collision_stops_before_discovery(self): managed = {"enabled_agents": {"codex": {"http_headers": {"X-Test": "admin"}}}} with _launch_policy_patches(managed) as calls: @@ -942,6 +1001,8 @@ def test_headers_parse_values_with_colons_and_deduplicate_case_insensitively(sel ("X-Test: safe\u2028Authorization: injected", "line separators"), ("X-Test: safe\u2029Authorization: injected", "line separators"), ("Authorization: secret", "protected header"), + ("Content-Type: text/plain", "protected header"), + ("Accept: text/plain", "protected header"), ("Cookie: secret", "protected header"), ("Databricks-Smart-Router-Recipe: custom", "protected header"), ("databricks-smart-router-recipe: custom", "protected header"), @@ -4650,8 +4711,14 @@ def test_uc_models_used_without_legacy_fallback(self, monkeypatch): assert legacy_called == [] assert "uc_enabled" not in state - def test_codex_only_configure_persists_discovered_oss_models(self, monkeypatch): - cli_mod, *_ = self._stub_deps(monkeypatch, pat_token="dapi-pat") + @pytest.mark.parametrize("custom_headers", [{}, {"X-Test": "temporary"}]) + def test_codex_only_configure_persists_discovered_oss_models(self, monkeypatch, custom_headers): + from ucode.request_headers import custom_header_scope + from ucode.state import _without_managed_overlay + + prior = {"codex_models": ["saved-codex"], "oss_models": ["saved-oss"]} + cli_mod, *_ = self._stub_deps(monkeypatch, pat_token="dapi-pat", existing_state=prior) + monkeypatch.setattr(cli_mod, "_fetch_managed_config", lambda state: (None, False)) monkeypatch.setattr( cli_mod, "discover_model_services", @@ -4664,14 +4731,18 @@ def test_codex_only_configure_persists_discovered_oss_models(self, monkeypatch): ), ) - state = cli_mod.configure_shared_state( - self.WS, - profile="DEFAULT", - tools=["codex"], - ) + with custom_header_scope(custom_headers): + state = cli_mod.configure_shared_state( + self.WS, + profile="DEFAULT", + tools=["codex"], + ) assert state["codex_models"] == ["system.ai.gpt-5-6-sol"] assert state["oss_models"] == ["system.ai.glm-5-2"] + persisted = _without_managed_overlay(state) + for key in prior: + assert persisted[key] == (prior[key] if custom_headers else state[key]) def _stub_with_fable(self, monkeypatch): cli_mod, *_ = self._stub_deps(monkeypatch, pat_token="dapi-pat") diff --git a/tests/test_codex_routing.py b/tests/test_codex_routing.py index f18ea859f..d55dbff05 100644 --- a/tests/test_codex_routing.py +++ b/tests/test_codex_routing.py @@ -5,7 +5,10 @@ import json import urllib.error -from ucode.smart_routing import codex_routing +import pytest + +from ucode.request_headers import CUSTOM_HEADERS_ENV, custom_header_scope +from ucode.smart_routing import codex_routing, routing from ucode.smart_routing.codex_hooks import routing_models WS = "https://example.databricks.com" @@ -82,6 +85,37 @@ def fake_urlopen(request, timeout): } +@pytest.mark.parametrize("source", ["direct", "inherited", "other_workspace"]) +def test_route_request_includes_launch_scoped_custom_headers(monkeypatch, source): + captured = {} + headers = {"X-Development-Route": "route://development/test"} + active = source != "other_workspace" + + def fake_open(request, timeout, *, scoped_headers=False): + assert scoped_headers == active + captured.update({name.casefold(): value for name, value in request.header_items()}) + return _Response({"route_selection": [{"route_option": {"model": "gpt-5-6-sol"}}]}) + + monkeypatch.setattr(routing.request_headers, "urlopen", fake_open) + monkeypatch.setenv( + CUSTOM_HEADERS_ENV, + json.dumps({"workspace": WS if active else "https://other.example", "headers": headers}), + ) + with custom_header_scope(headers if source == "direct" else {}): + decision, error = codex_routing.request_routing_decision( + WS, + "token", + "route this task", + ["system.ai.gpt-5-6-sol"], + ) + + assert error is None + assert decision is not None + assert captured.get("x-development-route") == ("route://development/test" if active else None) + assert captured["authorization"] == "Bearer token" + assert captured["content-type"] == "application/json" + + def test_router_name_can_be_overridden_with_environment_variable(monkeypatch): captured = {} monkeypatch.setenv("SMART_ROUTER_NAME", " custom_router ") diff --git a/tests/test_databricks.py b/tests/test_databricks.py index 12325a0ca..af420204f 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -2711,6 +2711,107 @@ def fake_get(url, token): assert calls == [f"https://{WS_HOST}/api/2.1/unity-catalog/model-services?page_size=50"] +class TestHttpCustomHeaders: + def test_get_merges_custom_headers_without_overriding_auth_or_explicit(self, monkeypatch): + seen = [] + + class _FakeOpener: + def open(self, request, timeout=None): + seen.append(request) + return _FakeResponse({"ok": True}) + + monkeypatch.setattr( + db_mod, + "get_custom_headers", + lambda: { + "authorization": "wrong-token", + "X-Route": "scoped", + "x-selector": "scoped-selector", + }, + ) + monkeypatch.setattr(db_mod.urllib_request, "build_opener", lambda handler: _FakeOpener()) + + payload, reason = db_mod._http_get_json( + "https://x/y", + "tok", + headers={"X-Selector": "explicit-selector"}, + ) + + assert payload == {"ok": True} + assert reason is None + request_headers = {name.casefold(): value for name, value in seen[0].header_items()} + assert request_headers["authorization"] == "Bearer tok" + assert request_headers["x-route"] == "scoped" + assert request_headers["x-selector"] == "explicit-selector" + assert [name.casefold() for name, _ in seen[0].header_items()].count("authorization") == 1 + + def test_send_json_preserves_transport_headers_case_insensitively(self, monkeypatch): + seen = [] + + class _JsonResponse(_FakeResponse): + status = 200 + + class _FakeOpener: + def open(self, request, timeout=None): + seen.append(request) + return _JsonResponse({"ok": True}) + + monkeypatch.setattr( + db_mod, + "get_custom_headers", + lambda: { + "authorization": "wrong-token", + "accept": "text/plain", + "content-type": "text/plain", + "X-Route": "scoped", + }, + ) + monkeypatch.setattr(db_mod.urllib_request, "build_opener", lambda handler: _FakeOpener()) + + payload, reason = db_mod._http_post_json("https://x/y", "tok", {"value": 1}) + + assert payload == {"ok": True} + assert reason is None + request_headers = {name.casefold(): value for name, value in seen[0].header_items()} + assert request_headers["authorization"] == "Bearer tok" + assert request_headers["accept"] == "application/json" + assert request_headers["content-type"] == "application/json" + assert request_headers["x-route"] == "scoped" + + def test_custom_headers_are_absent_from_redirected_request(self, monkeypatch): + seen = [] + handlers = [] + + class _FakeOpener: + def open(self, request, timeout=None): + seen.append(request) + return _FakeResponse({"ok": True}) + + monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + monkeypatch.setattr( + db_mod.urllib_request, + "build_opener", + lambda handler: handlers.append(handler) or _FakeOpener(), + ) + db_mod._http_get_bytes("https://x/y", "tok") + + with pytest.raises(db_mod.urllib_error.HTTPError): + handlers[0].redirect_request(seen[0], None, 302, "Found", {}, "https://other/y") + + redirected = handlers[0].redirect_request(seen[0], None, 302, "Found", {}, "https://x/z") + assert redirected is not None + assert {name.casefold() for name, _ in redirected.header_items()} >= { + "x-route", + } + + seen.clear() + handlers.clear() + monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {}) + db_mod._http_get_bytes("https://x/y", "tok", headers={"X-Route": "explicit"}) + with pytest.raises(db_mod.urllib_error.HTTPError): + handlers[0].redirect_request(seen[0], None, 302, "Found", {}, "https://other/y") + + class TestHttpGetJsonReason: """The `reason` string returned by `_http_get_json` must include the response body so callers (e.g. the Unity Gateway capability probe) can route on it. Before issue #84's fix @@ -3590,6 +3691,19 @@ def test_use_cache_false_forces_a_fresh_walk(self, monkeypatch): db_mod.list_model_services(WS, "tok", use_cache=False) assert calls["n"] == 2 + def test_custom_headers_bypass_and_do_not_overwrite_cache(self, monkeypatch): + calls: dict = {} + db_mod.clear_model_services_cache() + monkeypatch.setattr(db_mod, "_get_model_services_page", self._counting_page(calls)) + + db_mod.list_model_services(WS, "tok") + monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + db_mod.list_model_services(WS, "tok") + monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {}) + db_mod.list_model_services(WS, "tok") + + assert calls["n"] == 2 + def test_each_workspace_is_cached_separately(self, monkeypatch): calls: dict = {} db_mod.clear_model_services_cache() @@ -3663,6 +3777,19 @@ def test_use_cache_false_forces_a_fresh_call(self, monkeypatch): db_mod.list_model_provider_services(WS, "tok", use_cache=False) assert calls["n"] == 2 + def test_custom_headers_bypass_and_do_not_overwrite_cache(self, monkeypatch): + calls: dict = {} + db_mod.clear_model_services_cache() + monkeypatch.setattr(db_mod, "_http_get_json", self._counting_listing(calls)) + + db_mod.list_model_provider_services(WS, "tok") + monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + db_mod.list_model_provider_services(WS, "tok") + monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {}) + db_mod.list_model_provider_services(WS, "tok") + + assert calls["n"] == 2 + def test_each_workspace_is_cached_separately(self, monkeypatch): calls: dict = {} db_mod.clear_model_services_cache() diff --git a/tests/test_managed_config.py b/tests/test_managed_config.py index 6ea15d435..314fac142 100644 --- a/tests/test_managed_config.py +++ b/tests/test_managed_config.py @@ -585,6 +585,49 @@ def test_read_failure_without_persisted_config_is_silent(self, monkeypatch): result, _ = refresh_managed_config(_state()) assert result is None + def test_custom_header_fetch_failure_raises_instead_of_using_cache(self, monkeypatch): + def boom(ws, profile): + raise RuntimeError("no token") + + monkeypatch.setattr(mc_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + monkeypatch.setattr(mc_mod, "get_databricks_token", boom) + monkeypatch.setattr( + mc_mod, + "_persisted_fallback", + lambda *args, **kwargs: pytest.fail("custom header fetch must not use the cache"), + ) + + with pytest.raises(RuntimeError, match="Could not fetch managed configuration: no token"): + refresh_managed_config(_state()) + + def test_custom_header_read_failure_raises_without_clearing_cache(self, monkeypatch): + saved: list[tuple] = [] + monkeypatch.setattr(mc_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + monkeypatch.setattr(mc_mod, "get_managed_config", lambda ws, tok: (None, "HTTP 403")) + monkeypatch.setattr( + mc_mod, + "save_managed_state", + lambda ws, cfg, **kwargs: saved.append((ws, cfg, kwargs)), + ) + + with pytest.raises(RuntimeError, match="Could not fetch managed configuration: HTTP 403"): + refresh_managed_config(_state()) + assert saved == [] + + def test_custom_feature_disabled_is_authoritative_without_clearing_cache(self, monkeypatch): + saved: list[tuple] = [] + reason = 'HTTP 400 Bad Request: {"error_code":"FEATURE_DISABLED"}' + monkeypatch.setattr(mc_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + monkeypatch.setattr(mc_mod, "get_managed_config", lambda ws, tok: (None, reason)) + monkeypatch.setattr( + mc_mod, + "save_managed_state", + lambda ws, cfg, **kwargs: saved.append((ws, cfg, kwargs)), + ) + + assert refresh_managed_config(_state()) == (None, True) + assert saved == [] + def test_auth_failure_falls_back_to_the_persisted_config(self, monkeypatch): warnings: list[str] = [] @@ -780,6 +823,20 @@ def test_fresh_published_cache_short_circuits(self, monkeypatch): self._no_fetch(monkeypatch) assert refresh_managed_config(_state()) == (normalize_managed_config(RAW_MANIFEST), False) + def test_custom_header_fetch_bypasses_and_preserves_persistent_cache(self, monkeypatch): + self._write_cache( + config=RAW_MANIFEST, outcome="published", retrieved_at=NOW - timedelta(minutes=1) + ) + calls = self._counting_fetch(monkeypatch) + monkeypatch.setattr(mc_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) + before = mc_mod.MANAGED_CONFIG_PATH.read_text(encoding="utf-8") + + result = refresh_managed_config(_state()) + + assert result == (normalize_managed_config(RAW_MANIFEST), False) + assert calls["n"] == 1 + assert mc_mod.MANAGED_CONFIG_PATH.read_text(encoding="utf-8") == before + def test_fresh_legacy_cache_normalizes_to_smart_defaults(self, monkeypatch): self._write_cache( config=LEGACY_RAW_MANIFEST, From 1773f562f58be602ab787495a8145e043712e8b4 Mon Sep 17 00:00:00 2001 From: Andy Xu Date: Wed, 7 Oct 2026 21:22:11 +0000 Subject: [PATCH 6/6] Verify scoped discovery preserves ordinary cached results --- tests/test_databricks.py | 46 ++++++++++++++++++++++++++++++++++------ 1 file changed, 40 insertions(+), 6 deletions(-) diff --git a/tests/test_databricks.py b/tests/test_databricks.py index af420204f..11bf7f86d 100644 --- a/tests/test_databricks.py +++ b/tests/test_databricks.py @@ -3652,6 +3652,12 @@ class TestModelServicesCache: def _counting_page(calls: dict): def page(url, token): calls["n"] = calls.get("n", 0) + 1 + if db_mod.get_custom_headers(): + return { + "model_services": [ + {"name": "model-services/system.ai.claude-opus-6"}, + ] + }, None return { "model_services": [ {"name": "model-services/system.ai.claude-opus-5"}, @@ -3696,12 +3702,15 @@ def test_custom_headers_bypass_and_do_not_overwrite_cache(self, monkeypatch): db_mod.clear_model_services_cache() monkeypatch.setattr(db_mod, "_get_model_services_page", self._counting_page(calls)) - db_mod.list_model_services(WS, "tok") + ordinary, _ = db_mod.list_model_services(WS, "tok") monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) - db_mod.list_model_services(WS, "tok") + scoped, _ = db_mod.list_model_services(WS, "tok") + assert scoped == ["system.ai.claude-opus-6"] monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {}) - db_mod.list_model_services(WS, "tok") + cached, _ = db_mod.list_model_services(WS, "tok") + assert ordinary == ["system.ai.claude-opus-4-8", "system.ai.claude-opus-5"] + assert cached == ordinary assert calls["n"] == 2 def test_each_workspace_is_cached_separately(self, monkeypatch): @@ -3740,6 +3749,25 @@ class TestModelProviderServicesCache: def _counting_listing(calls: dict): def get_json(url, token, timeout=10): calls["n"] = calls.get("n", 0) + 1 + if db_mod.get_custom_headers(): + return { + "model_provider_services": [ + { + "name": "model-provider-services/scoped.route.ant", + "config": { + "provider_type": "ANTHROPIC", + "targets": [{"model": "claude-opus-6"}], + }, + }, + { + "name": "model-provider-services/scoped.route.oai", + "config": { + "provider_type": "OPENAI", + "targets": [{"model": "gpt-6"}], + }, + }, + ] + }, None return { "model_provider_services": [ { @@ -3782,12 +3810,18 @@ def test_custom_headers_bypass_and_do_not_overwrite_cache(self, monkeypatch): db_mod.clear_model_services_cache() monkeypatch.setattr(db_mod, "_http_get_json", self._counting_listing(calls)) - db_mod.list_model_provider_services(WS, "tok") + ordinary, _ = db_mod.list_model_provider_services(WS, "tok") monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {"X-Route": "scoped"}) - db_mod.list_model_provider_services(WS, "tok") + scoped, _ = db_mod.list_model_provider_services(WS, "tok") + assert [service["name"] for service in scoped] == [ + "scoped.route.ant", + "scoped.route.oai", + ] monkeypatch.setattr(db_mod, "get_custom_headers", lambda: {}) - db_mod.list_model_provider_services(WS, "tok") + cached, _ = db_mod.list_model_provider_services(WS, "tok") + assert [service["name"] for service in ordinary] == ["main.j.ant", "main.j.oai"] + assert cached == ordinary assert calls["n"] == 2 def test_each_workspace_is_cached_separately(self, monkeypatch):