Skip to content
Open
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
82 changes: 77 additions & 5 deletions src/ucode/agents/pi.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@
build_pi_base_urls,
classify_model_family,
get_databricks_token,
model_token_limits,
)
from ucode.state import mark_tool_managed, save_state
from ucode.telemetry import agent_version, ucode_version
Expand All @@ -69,6 +70,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 @@ -97,13 +99,31 @@ def _resolve_model_selector(
return model


def _bedrock_model_entry(model_id: str) -> dict:
"""A Pi model entry for a Bedrock target, pinning known token limits.

Some Bedrock models cap output well below what Pi requests by default (e.g.
Nova rejects a `maxTokens` of 10k or more), so pin `maxTokens`/`contextWindow`
when the model has a known limit. Models with no known limit are left unbounded.
"""
entry: dict = {"id": model_id}
limits = model_token_limits(model_id)
if limits is not None:
entry["contextWindow"] = limits["context"]
entry["maxTokens"] = limits["output"]
return entry


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 +167,43 @@ 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,
# Pi's bedrock-converse-stream client (AWS SDK style) sets its own
# User-Agent; adding ours produces two `user-agent` values and the
# gateway rejects the request ("Header field ... must only have a
# single value"). Send only the MPS selector header here.
"headers": {"Databricks-Model-Provider-Service": provider},
"models": [_bedrock_model_entry(t) for t in bedrock_targets],
}
keys.append(["providers", "databricks-bedrock"])
resolved = _resolve_model_selector(model or "", claude_models, codex_models, gemini_models)
# Bedrock model IDs contain no `/` (e.g. `anthropic.claude-3-haiku-20240307-v1:0`), so
# _resolve_model_selector returns them unprefixed. _write_settings splits on `/` to get
# provider/model — without the prefix it gets an empty model_id and skips defaultProvider.
# Always force the `databricks-bedrock/` prefix when the Bedrock provider is active.
if "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 +224,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 Expand Up @@ -259,6 +304,33 @@ def default_model(state: dict) -> str | None:


def _refresh_token_once(state: dict, *, force_refresh: bool = False) -> str:
# Preserve a Bedrock provider block across token refreshes. The block is
# self-describing: its MPS header + model ids are enough to re-render it,
# so a refresh keeps routing through Bedrock instead of dropping to a
# system-hosted model. When the config has no Bedrock block (a non-Bedrock
# session, or after a non-Bedrock reconfigure overwrote it), fall through
# to the normal path.
existing = read_json_safe(PI_CONFIG_PATH)
bedrock = (existing.get("providers") or {}).get("databricks-bedrock")
provider: str | None = None
bedrock_targets: list[str] | None = None
if isinstance(bedrock, dict):
headers = bedrock.get("headers") or {}
provider = headers.get("Databricks-Model-Provider-Service")
bedrock_targets = [
m["id"]
for m in (bedrock.get("models") or [])
if isinstance(m, dict) and isinstance(m.get("id"), str)
] or None
if provider and bedrock_targets:
_, token = write_tool_config(
state,
bedrock_targets[0],
force_refresh=force_refresh,
provider=provider,
bedrock_targets=bedrock_targets,
)
return token
model = default_model(state)
if not model:
raise RuntimeError("No Pi model is available on this workspace.")
Expand Down
122 changes: 121 additions & 1 deletion src/ucode/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,15 +51,18 @@
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_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 +128,7 @@
from ucode.ui import (
console,
heading,
muted,
print_err,
print_heading,
print_kv,
Expand All @@ -133,9 +137,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 +1165,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 +2052,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 +2062,24 @@ 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 == "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 +2116,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 +2530,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 +3294,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