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
5 changes: 3 additions & 2 deletions src/ucode/agents/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -363,17 +363,18 @@ def check_gateway_endpoint(state: dict, tool: str) -> bool:
bool(state.get("claude_models"))
or bool(state.get("codex_models"))
or bool(state.get("gemini_models"))
or bool(state.get("oss_models"))
)
return False


_TOOL_DISCOVERY_SOURCES: dict[str, tuple[str, ...]] = {
"claude": ("claude",),
"opencode": ("claude", "gemini", "oss"),
"opencode": ("claude", "codex", "gemini", "oss"),
"codex": ("codex",),
"gemini": ("gemini",),
"copilot": ("claude", "codex"),
"pi": ("claude", "codex", "gemini"),
"pi": ("claude", "codex", "gemini", "oss"),
}


Expand Down
29 changes: 28 additions & 1 deletion src/ucode/agents/opencode.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@

PROVIDER_KEYS: list[list[str]] = [
["provider", "databricks-anthropic"],
["provider", "databricks-openai"],
["provider", "databricks-google"],
["provider", "databricks-oss"],
]
Expand All @@ -51,13 +52,24 @@ def is_update_available() -> tuple[str, str] | None:

def _resolve_model_selector(model: str, opencode_models: dict[str, list[str]]) -> str:
"""Return an OpenCode model selector in provider/model form when possible."""
if model.startswith(("databricks-anthropic/", "databricks-google/", "databricks-oss/")):
if model.startswith(
(
"databricks-anthropic/",
"databricks-openai/",
"databricks-google/",
"databricks-oss/",
)
):
return model

anthropic_models = opencode_models.get("anthropic") or []
if model in anthropic_models:
return f"databricks-anthropic/{model}"

openai_models = opencode_models.get("openai") or []
if model in openai_models:
return f"databricks-openai/{model}"

gemini_models = opencode_models.get("gemini") or []
if model in gemini_models:
return f"databricks-google/{model}"
Expand Down Expand Up @@ -100,6 +112,7 @@ def render_overlay(
}

anthropic_models = opencode_models.get("anthropic") or []
openai_models = opencode_models.get("openai") or []
gemini_models = opencode_models.get("gemini") or []
oss_models = opencode_models.get("oss") or []

Expand All @@ -125,6 +138,17 @@ def render_overlay(
"models": dict.fromkeys(anthropic_models, anthropic_model_overlay),
}
keys.append(["provider", "databricks-anthropic"])
if openai_models:
providers["databricks-openai"] = {
"npm": "@ai-sdk/openai",
"options": {
"baseURL": opencode_base_urls["openai"],
"apiKey": token,
"headers": auth_headers,
},
"models": {m: {"headers": ua_header} for m in openai_models},
}
keys.append(["provider", "databricks-openai"])
if gemini_models:
providers["databricks-google"] = {
"npm": "@ai-sdk/google",
Expand Down Expand Up @@ -232,6 +256,9 @@ def default_model(state: dict) -> str | None:
anthropic = opencode_models.get("anthropic") or []
if anthropic:
return anthropic[0]
openai = opencode_models.get("openai") or []
if openai:
return openai[0]
gemini = opencode_models.get("gemini") or []
if gemini:
return gemini[0]
Expand Down
48 changes: 40 additions & 8 deletions src/ucode/agents/pi.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
- `databricks-claude` (api: anthropic-messages) → /ai-gateway/anthropic
- `databricks-openai` (api: openai-responses) → /ai-gateway/codex/v1
- `databricks-gemini` (api: google-generative-ai) → /ai-gateway/gemini/v1beta
- `databricks-oss` (api: openai-responses) → /ai-gateway/mlflow/v1

Per-provider `compat` flags work around fields the gateway translators reject:

Expand All @@ -16,10 +17,10 @@
sends the legacy `anthropic-beta: fine-grained-tool-streaming-...` header
instead, which the gateway accepts.

OSS / Databricks-foundation models (Llama, Qwen, etc.) are not exposed via
pi today — they live behind /ai-gateway/mlflow/v1 with per-model
`max_tokens` caps that pi has no global way to honor without per-model
config we don't currently maintain.
- OSS models use the MLflow Responses route because its chat-completions stream
can end without a `finish_reason`, which Pi treats as an error. They also
carry per-model `contextWindow` and `maxTokens` from the shared token-limits
table.

The bearer token is baked into the file and refreshed by a background thread
while the session runs (same pattern as OpenCode/Copilot).
Expand All @@ -45,6 +46,7 @@
TOKEN_REFRESH_INTERVAL_SECONDS,
build_pi_base_urls,
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 @@ -68,13 +70,14 @@
"databricks-claude",
"databricks-openai",
"databricks-gemini",
"databricks-oss",
)

PROVIDER_KEYS: list[list[str]] = [["providers", name] for name in PROVIDER_NAMES]

# Old provider names earlier ucode versions wrote; cleaned up on each write so
# users don't end up with stale entries pointing at routes that 400.
LEGACY_PROVIDER_NAMES = ("databricks-anthropic", "databricks-codex", "databricks-oss")
LEGACY_PROVIDER_NAMES = ("databricks-anthropic", "databricks-codex", "databricks-kimi")


def is_update_available() -> tuple[str, str] | None:
Expand All @@ -86,6 +89,7 @@ def _resolve_model_selector(
claude_models: dict[str, str],
codex_models: list[str],
gemini_models: list[str],
oss_models: list[str],
) -> str:
"""Return a Pi model selector in `<provider>/<model>` form when possible."""
for name in PROVIDER_NAMES:
Expand All @@ -97,16 +101,28 @@ def _resolve_model_selector(
return f"databricks-openai/{model}"
if model in gemini_models:
return f"databricks-gemini/{model}"
if model in oss_models:
return f"databricks-oss/{model}"
return model


def _oss_model_entry(model: str) -> dict:
entry: dict = {"id": model}
limits = model_token_limits(model)
if limits is not None:
entry["contextWindow"] = limits["context"]
entry["maxTokens"] = limits["output"]
return entry


def render_overlay(
model: str,
token: str,
pi_base_urls: dict[str, str],
claude_models: dict[str, str],
codex_models: list[str],
gemini_models: list[str],
oss_models: list[str],
) -> tuple[dict, list[list[str]]]:
"""Return (overlay, managed_key_paths) for ~/.pi/agent/models.json."""
providers: dict = {}
Expand Down Expand Up @@ -150,8 +166,20 @@ def render_overlay(
"models": [{"id": m} for m in gemini_models],
}
keys.append(["providers", "databricks-gemini"])
if oss_models:
providers["databricks-oss"] = {
"baseUrl": pi_base_urls["oss"],
"api": "openai-responses",
"apiKey": token,
"authHeader": True,
"headers": ua_headers,
"models": [_oss_model_entry(m) for m in oss_models],
}
keys.append(["providers", "databricks-oss"])
overlay: dict = {
"model": _resolve_model_selector(model, claude_models, codex_models, gemini_models),
"model": _resolve_model_selector(
model, claude_models, codex_models, gemini_models, oss_models
),
}
if providers:
overlay["providers"] = providers
Expand All @@ -178,6 +206,7 @@ def write_tool_config(
state.get("claude_models") or {},
state.get("codex_models") or [],
state.get("gemini_models") or [],
state.get("oss_models") or [],
)
existing = read_json_safe(PI_CONFIG_PATH)
providers = existing.get("providers")
Expand Down Expand Up @@ -206,7 +235,7 @@ def _write_settings(model_selector: str) -> None:


def default_model(state: dict) -> str | None:
"""Prefer Claude opus → sonnet → haiku; fall back to codex, gemini."""
"""Prefer Claude opus → sonnet → haiku; then codex, gemini, OSS."""
claude_models = state.get("claude_models") or {}
for family in ("opus", "sonnet", "haiku"):
if claude_models.get(family):
Expand All @@ -215,7 +244,10 @@ def default_model(state: dict) -> str | None:
if codex_models:
return codex_models[0]
gemini_models = state.get("gemini_models") or []
return gemini_models[0] if gemini_models else None
if gemini_models:
return gemini_models[0]
oss_models = state.get("oss_models") or []
return oss_models[0] if oss_models else None


def _refresh_token_once(state: dict, *, force_refresh: bool = False) -> str:
Expand Down
8 changes: 6 additions & 2 deletions src/ucode/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -370,8 +370,10 @@ def configure_shared_state(
fetch_all or "claude" in tools or "opencode" in tools or "copilot" in tools or "pi" in tools
)
want_gemini = fetch_all or "gemini" in tools or "opencode" in tools or "pi" in tools
want_codex = fetch_all or "codex" in tools or "copilot" in tools or "pi" in tools
want_oss = fetch_all or "opencode" in tools
want_codex = (
fetch_all or "codex" in tools or "opencode" in tools or "copilot" in tools or "pi" in tools
)
want_oss = fetch_all or "opencode" in tools or "pi" in tools

claude_reason: str | None = None
gemini_reason: str | None = None
Expand Down Expand Up @@ -425,6 +427,8 @@ def configure_shared_state(
oss_models, oss_reason = ms_oss, ms_reason
if claude_models:
opencode_models["anthropic"] = list(claude_models.values())
if codex_models:
opencode_models["openai"] = codex_models
if gemini_models:
opencode_models["gemini"] = gemini_models
if oss_models:
Expand Down
13 changes: 7 additions & 6 deletions src/ucode/databricks.py
Original file line number Diff line number Diff line change
Expand Up @@ -1333,7 +1333,9 @@ def discover_model_services(
- ``claude_models`` maps ``fable``/``opus``/``sonnet``/``haiku`` to the
newest matching ``system.ai.claude-*`` id (mirrors
``discover_claude_models``).
- ``codex_models`` is the list of ``system.ai.*gpt-*`` ids.
- ``codex_models`` is the list of Responses-compatible ``system.ai.*gpt-*``
ids; ``gpt-oss`` is excluded because its gateway rejects the session and
prompt-caching fields sent by Pi and OpenCode.
- ``gemini_models`` is the list of ``system.ai.*gemini-*`` ids, newest first.
- ``oss_models`` is the list of OSS-model ``system.ai.*`` ids.

Expand All @@ -1354,7 +1356,7 @@ def discover_model_services(
if candidates:
claude_models[family] = candidates[0]

codex_models = [m for m in ids if "gpt-" in m]
codex_models = [m for m in ids if "gpt-" in m and "gpt-oss-" not in m]
gemini_models = sorted([m for m in ids if "gemini-" in m], key=model_version_sort_key)

oss_models = [m for m in ids if any(family in m for family in _OSS_MODEL_FAMILIES)]
Expand Down Expand Up @@ -2370,6 +2372,7 @@ def build_tool_base_url(tool: str, workspace: str) -> str:
def build_opencode_base_urls(workspace: str) -> dict[str, str]:
return {
"anthropic": build_tool_base_url("claude", workspace) + "/v1",
"openai": build_tool_base_url("codex", workspace),
"gemini": build_tool_base_url("gemini", workspace) + "/v1beta",
"oss": f"{workspace}/ai-gateway/mlflow/v1",
}
Expand All @@ -2380,17 +2383,15 @@ def build_pi_base_urls(workspace: str) -> dict[str, str]:
# path (verified end-to-end). Each `api` type appends its own path suffix:
#
# - anthropic-messages appends `/v1/messages`
# - openai-responses appends `/responses`
# - openai-responses appends `/responses` (codex and OSS providers)
# - google-generative-ai appends `/v1beta/models/{id}:streamGenerateContent`
# - openai-completions appends `/chat/completions`
#
# So the baseUrls below stop just before the suffix Pi will tack on.
# Compat flags applied per-provider in agents/pi.py; required for `oss`
# only (MLflow rejects `store` and `tools[].function.strict`).
return {
"claude": build_tool_base_url("claude", workspace),
"openai": build_tool_base_url("codex", workspace),
"gemini": build_tool_base_url("gemini", workspace) + "/v1beta",
"oss": f"{workspace}/ai-gateway/mlflow/v1",
}


Expand Down
7 changes: 7 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

from ucode.databricks import (
build_shared_base_urls,
discover_model_services,
fetch_ai_gateway_claude_models,
fetch_codex_models,
fetch_gemini_models,
Expand Down Expand Up @@ -56,18 +57,24 @@ def e2e_state(e2e_workspace, e2e_token):
claude_models = fetch_ai_gateway_claude_models(e2e_workspace, e2e_token)
gemini_models = fetch_gemini_models(e2e_workspace, e2e_token)
codex_models = fetch_codex_models(e2e_workspace, e2e_token)
_, _, _, oss_models, _ = discover_model_services(e2e_workspace, e2e_token)

opencode_models: dict = {}
if claude_models:
opencode_models["anthropic"] = list(claude_models.values())
if codex_models:
opencode_models["openai"] = codex_models
if gemini_models:
opencode_models["gemini"] = gemini_models
if oss_models:
opencode_models["oss"] = oss_models

return {
"workspace": e2e_workspace,
"claude_models": claude_models,
"gemini_models": gemini_models,
"codex_models": codex_models,
"oss_models": oss_models,
"opencode_models": opencode_models,
"base_urls": build_shared_base_urls(e2e_workspace),
"managed_configs": {},
Expand Down
Loading