Skip to content
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
35 changes: 30 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,39 @@ 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)
# 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 +204,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
Loading