Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 10 additions & 3 deletions src/ucode/agents/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -350,6 +350,7 @@ def configure_tool(
relayed: bool = False,
route_root_model: str | None = None,
custom_model: str | None = None,
bedrock_targets: list[str] | None = None,
) -> dict:
result: dict | tuple[dict, str]
if tool == "codex":
Expand All @@ -370,16 +371,22 @@ def configure_tool(
custom_model=custom_model,
)
else:
# provider routing is claude/codex-only; every other tool needs a model.
if not model:
# provider routing is claude/codex-only; every other tool needs a model —
# except pi with a Bedrock provider, where targets replace the model list.
if not model and not (tool == "pi" and provider and bedrock_targets):
raise RuntimeError(f"A {tool} model must be selected before configuration.")
if tool == "gemini":
assert model is not None
result = gemini.write_tool_config(state, model)
elif tool == "copilot":
assert model is not None
result = copilot.write_tool_config(state, model)
elif tool == "pi":
result = pi.write_tool_config(state, model)
result = pi.write_tool_config(
state, model, provider=provider, bedrock_targets=bedrock_targets
)
else:
assert model is not None
result = opencode.write_tool_config(state, model)
# gemini/opencode/copilot/pi return (state, token); codex/claude return state
if isinstance(result, tuple):
Expand Down
9 changes: 5 additions & 4 deletions src/ucode/agents/codex.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,11 +300,12 @@ def revert_legacy_shared_config() -> bool:

def write_tool_config(state: dict, model: str | None = None, provider: str | None = None) -> dict:
workspace = state["workspace"]
# Leave model selection to Codex. The gateway still receives the configured
# provider and authentication settings, while Codex uses its own default.
# A managed default is the sole exception.
# Leave model selection to Codex — except when a provider is set and a target
# model was resolved from its MPS targets, or an admin managed default exists.
managed_model = state.get("codex_default_model")
chosen_model = managed_model if isinstance(managed_model, str) else None
chosen_model = (model if provider else None) or (
managed_model if isinstance(managed_model, str) else None
)
databricks_profile = state.get("profile")

if _use_legacy_layout():
Expand Down
32 changes: 27 additions & 5 deletions src/ucode/agents/pi.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,7 @@
"databricks-claude",
"databricks-openai",
"databricks-gemini",
"databricks-bedrock",
)

PROVIDER_KEYS: list[list[str]] = [["providers", name] for name in PROVIDER_NAMES]
Expand Down Expand Up @@ -98,12 +99,15 @@ def _resolve_model_selector(


def render_overlay(
model: str,
model: str | None,
token: str,
pi_base_urls: dict[str, str],
claude_models: dict[str, str],
codex_models: list[str],
gemini_models: list[str],
*,
provider: str | None = None,
bedrock_targets: list[str] | None = None,
) -> tuple[dict, list[list[str]]]:
"""Return (overlay, managed_key_paths) for Pi's private agent config."""
providers: dict = {}
Expand Down Expand Up @@ -147,20 +151,36 @@ def render_overlay(
"models": [{"id": m} for m in gemini_models],
}
keys.append(["providers", "databricks-gemini"])
overlay: dict = {
"model": _resolve_model_selector(model, claude_models, codex_models, gemini_models),
}
if provider and bedrock_targets:
providers["databricks-bedrock"] = {
"baseUrl": pi_base_urls.get(
"bedrock", f"{pi_base_urls['claude'].rsplit('/ai-gateway', 1)[0]}/ai-gateway"
),
"api": "bedrock-converse-stream",
"apiKey": token,
"authHeader": True,
"headers": {**ua_headers, "Databricks-Model-Provider-Service": provider},
"models": [{"id": t} for t in bedrock_targets],
}
keys.append(["providers", "databricks-bedrock"])
resolved = _resolve_model_selector(model or "", claude_models, codex_models, gemini_models)
# When launching with a Bedrock provider, default to the first target.
if not resolved and "databricks-bedrock" in providers and bedrock_targets:
resolved = f"databricks-bedrock/{bedrock_targets[0]}"
overlay: dict = {"model": resolved}
if providers:
overlay["providers"] = providers
return overlay, keys


def write_tool_config(
state: dict,
model: str,
model: str | None,
token: str | None = None,
*,
force_refresh: bool = False,
provider: str | None = None,
bedrock_targets: list[str] | None = None,
) -> tuple[dict, str]:
backup_existing_file(PI_CONFIG_PATH, PI_BACKUP_PATH)
if token is None:
Expand All @@ -181,6 +201,8 @@ def write_tool_config(
claude_models,
codex_models,
gemini_models,
provider=provider,
bedrock_targets=bedrock_targets,
)
existing = read_json_safe(PI_CONFIG_PATH)
providers = existing.get("providers")
Expand Down
187 changes: 185 additions & 2 deletions src/ucode/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@
)
from ucode.agents.codex import revert_legacy_shared_config
from ucode.agents.pi import PI_SETTINGS_BACKUP_PATH, PI_SETTINGS_PATH
from ucode.config_io import is_dry_run, restore_file, set_dry_run
from ucode.config_io import is_dry_run, read_toml_safe, restore_file, set_dry_run
from ucode.databricks import (
apply_pat_environment,
build_shared_base_urls,
Expand All @@ -51,15 +51,19 @@
find_profile_name_for_host,
get_databricks_profiles,
get_databricks_token,
get_model_provider_service,
install_databricks_cli,
is_model_provider_feature_unavailable,
is_workspace_admin,
list_model_provider_services,
list_mps_codex_models,
list_profile_entries,
list_tool_provider_services,
normalize_workspace_url,
resolve_pat_token,
resolve_provider_launch_model,
run_databricks_login,
service_usable_for_tool,
)
from ucode.managed_budget import (
budget_usage_percent,
Expand Down Expand Up @@ -125,6 +129,7 @@
from ucode.ui import (
console,
heading,
muted,
print_err,
print_heading,
print_kv,
Expand All @@ -133,9 +138,11 @@
print_success,
print_warning,
prompt_for_selection,
prompt_for_text,
prompt_for_tools,
prompt_for_workspace,
prompt_yes_no,
render_box_table,
set_verbosity,
spinner,
status_badge,
Expand Down Expand Up @@ -1159,6 +1166,10 @@ def revert() -> int:
app.add_typer(configure_app, name="configure", help="Configure workspace and tool settings.")
mcp_app = typer.Typer(add_completion=False, no_args_is_help=True)
app.add_typer(mcp_app, name="mcp", help="MCP servers exposed by ucode.")
providers_app = typer.Typer(add_completion=False, no_args_is_help=True)
app.add_typer(
providers_app, name="providers", help="Inspect Model Provider Services on the workspace."
)
setup_app = typer.Typer(add_completion=False, no_args_is_help=False)
app.add_typer(
setup_app,
Expand Down Expand Up @@ -2042,6 +2053,7 @@ def _launch_tool(
# The router's per-launch pick for the root session. Codex pins it as the
# resolved model; claude pins it via ANTHROPIC_MODEL (route_root_model).
route_root_model = None
bedrock_targets: list[str] | None = None
if provider:
# Routing through a Model Provider Service pins no Databricks model;
# the agent uses its own canonical model names (header selects the
Expand All @@ -2051,6 +2063,86 @@ def _launch_tool(
# Relayed services forward --model to Claude Code's own flag at launch (below), not env.
if tool == "claude" and not relayed and (model or provider_models):
route_root_model = resolve_provider_launch_model(model, provider_models or {})
elif tool == "codex":
# Codex's built-in model picker queries OpenAI, not the MPS, so it shows the
# wrong model list when routing through a Bedrock provider. Pin a target from
# the MPS so Codex never reaches its picker.
if model:
resolved_model = model
else:
_token = get_databricks_token(state["workspace"], state.get("profile"))
with spinner("Fetching provider model targets..."):
_svc, _ = get_model_provider_service(provider, state["workspace"], _token)
if _svc:
_targets: list[str] = _svc.get("targets") or []
if len(_targets) == 1:
resolved_model = _targets[0]
elif len(_targets) > 1:
_picked = prompt_for_selection(
"Select a model from the provider service:",
[(_t, _t) for _t in _targets],
)
if _picked is None:
raise KeyboardInterrupt
resolved_model = _picked
elif _svc.get("allow_all_targets"):
# No declared targets but the service allows any — query the
# provider's OpenAI-compatible /models endpoint to get the list.
# Reuse the previously-saved model only when the config was last
# written with this same provider; a workspace model from a
# non-MPS run must not bleed into a Bedrock session.
_prev_cfg = read_toml_safe(codex_agent.CODEX_CONFIG_PATH)
_stored_provider = (
_prev_cfg.get("model_providers", {})
.get(codex_agent.CODEX_MODEL_PROVIDER_NAME, {})
.get("http_headers", {})
.get("Databricks-Model-Provider-Service")
)
_prev_model: str | None = (
_prev_cfg.get("model") if _stored_provider == provider else None
)
with spinner("Querying available models from provider..."):
_mps_models, _mps_err = list_mps_codex_models(
provider, state["workspace"], _token
)
if _mps_err is None and _mps_models:
if len(_mps_models) == 1:
resolved_model = _mps_models[0]
else:
_mpicked = prompt_for_selection(
"Select a model from the provider service:",
[(_t, _t) for _t in _mps_models],
)
if _mpicked is None:
raise KeyboardInterrupt
resolved_model = _mpicked
else:
# Live query failed or returned nothing — fall back to a
# free-text prompt, defaulting to the previously-saved model
# so subsequent launches don't ask again.
resolved_model = prompt_for_text(
f"Enter the model ID to use with '{provider}'",
default=_prev_model,
required=not _prev_model,
)
elif tool == "pi":
# Pi receives the MPS targets as its databricks-bedrock model list;
# a single model is also set as the default for the session.
_pi_token = get_databricks_token(state["workspace"], state.get("profile"))
with spinner("Fetching provider model targets..."):
_pi_svc, _ = get_model_provider_service(provider, state["workspace"], _pi_token)
if _pi_svc:
bedrock_targets = _pi_svc.get("targets") or []
if bedrock_targets:
resolved_model = bedrock_targets[0]
elif _pi_svc.get("allow_all_targets"):
_pi_entered = prompt_for_text(
f"Enter a Bedrock model ID to use with '{provider}'",
required=True,
)
if _pi_entered:
bedrock_targets = [_pi_entered]
resolved_model = _pi_entered
else:
# A managed default_model is the model the admin wants sessions to start on, so it goes
# in as the explicit model rather than being applied afterwards: for codex the proto has
Expand Down Expand Up @@ -2087,6 +2179,7 @@ def _launch_tool(
# the latter pins a raw id into every family alias, which would clobber the service's
# per-family target pins.
custom_model=model if (tool == "claude" and not provider) else None,
bedrock_targets=bedrock_targets,
)
# Relayed = a Claude subscription: forward --model to Claude Code's own flag, like `-- --model X`.
if tool == "claude" and provider and relayed and model and not forwarded_model:
Expand Down Expand Up @@ -2500,12 +2593,20 @@ def copilot_cmd(
@app.command("pi", context_settings={"allow_extra_args": True, "ignore_unknown_options": True})
def pi_cmd(
ctx: typer.Context,
provider: Annotated[
str | None,
typer.Option(
"--provider",
help="Route through a Unity Catalog Model Provider Service "
"(<catalog>.<schema>.<name>). Pass before any `--` separator.",
),
] = None,
skip_preflight: SkipPreflightOption = False,
skip_managed_config: SkipManagedConfigOption = False,
) -> None:
"""Launch Pi coding agent via Databricks."""
_disable_managed_config_if_requested(skip_managed_config)
_launch_tool("pi", ctx, skip_preflight=skip_preflight)
_launch_tool("pi", ctx, provider=provider, skip_preflight=skip_preflight)


@app.command("cursor", context_settings={"allow_extra_args": True, "ignore_unknown_options": True})
Expand Down Expand Up @@ -3256,6 +3357,88 @@ def upgrade_cmd() -> None:
print_success("ucode upgraded")


@providers_app.command("list")
def providers_list_cmd(
tool: Annotated[
str | None,
typer.Option(
"--tool", help="Filter to services usable by a specific tool (claude, codex)."
),
] = None,
) -> None:
"""List Model Provider Services on the workspace."""
state = load_state()
workspace = state.get("workspace")
if not workspace:
print_err("No workspace configured. Run `ucode configure` first.")
raise typer.Exit(1) from None
token = get_databricks_token(workspace, state.get("profile"))
with spinner("Fetching model provider services..."):
services, reason = list_model_provider_services(workspace, token)
if reason is not None:
print_err(f"Could not list model provider services: {reason}")
raise typer.Exit(1) from None
if tool:
services = [s for s in services if service_usable_for_tool(tool, s)]
if not services:
msg = "No model provider services found" + (f" for {tool}" if tool else "") + "."
print_note(msg)
return
rows = [
[
s["name"],
s["provider_type"],
", ".join(s["targets"])
if s["targets"]
else ("(all)" if s["allow_all_targets"] else "—"),
]
for s in services
]
print_section("Model Provider Services")
console.print(
render_box_table(["Service", "Provider", "Targets"], rows, max_widths=[60, 20, 60])
)
if tool:
console.print(muted(f" Filtered to services usable by {tool}."))


@providers_app.command("show")
def providers_show_cmd(
service_name: Annotated[
str,
typer.Argument(help="Fully qualified service name (catalog.schema.service)."),
],
) -> None:
"""Show targets and configuration for a Model Provider Service."""
state = load_state()
workspace = state.get("workspace")
if not workspace:
print_err("No workspace configured. Run `ucode configure` first.")
raise typer.Exit(1) from None
token = get_databricks_token(workspace, state.get("profile"))
with spinner(f"Fetching {service_name}..."):
service, reason = get_model_provider_service(service_name, workspace, token)
if reason is not None:
print_err(f"Could not fetch '{service_name}': {reason}")
raise typer.Exit(1) from None
if service is None:
print_err(f"Model provider service '{service_name}' not found.")
raise typer.Exit(1) from None
print_section(service["name"])
print_kv("Provider type", service["provider_type"])
if service["relayed"]:
print_kv("Relay", "yes (subscription-backed, no credential stored)")
if service["allow_all_targets"]:
print_kv("Allow all targets", "yes")
targets = service["targets"]
if targets:
print_kv("Targets", targets[0])
for t in targets[1:]:
print_kv("", t)
else:
print_kv("Targets", "none declared")


def main() -> None:
app()

Expand Down
Loading