Skip to content
Closed
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
62 changes: 53 additions & 9 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-completions) → /ai-gateway/mlflow/v1

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

Expand All @@ -15,11 +16,11 @@
pi uses for every request. With this flag pi omits the per-tool field and
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: `supportsStore: false` + `supportsStrictMode: false` — the MLflow
chat-completions route rejects the OpenAI `store` field and
`tools[].function.strict`. OSS models also carry per-model `contextWindow`
and `maxTokens` (from `model_token_limits`) so pi clamps output to a value
the gateway accepts (e.g. GLM's 25k cap) instead of pi's 16k default.

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,17 @@
"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")
# `databricks-kimi` was written by an early build with `api: openai-responses`
# against the MLflow route (which speaks openai-completions), so its requests
# 404/401 and its baked token was never refreshed — strip it on every write.
LEGACY_PROVIDER_NAMES = ("databricks-anthropic", "databricks-codex", "databricks-kimi")


def is_update_available() -> tuple[str, str] | None:
Expand All @@ -86,6 +92,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 +104,34 @@ 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:
"""Per-model entry for an OSS model.

Pins `contextWindow`/`maxTokens` from the shared limits table when known so
pi clamps `max_tokens` to a value the MLflow route accepts (e.g. GLM's 25k
output cap) rather than pi's 16k default. Both fields are supplied together
because the limits table always provides both."""
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 +175,23 @@ 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-completions",
"apiKey": token,
"authHeader": True,
# The MLflow chat-completions route rejects the OpenAI `store` field
# and `tools[].function.strict`; opt out of both.
"compat": {"supportsStore": False, "supportsStrictMode": False},
"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 +218,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 +247,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; fall back to 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 +256,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
4 changes: 2 additions & 2 deletions src/ucode/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,7 @@
"claude": ("claude", "opencode", "copilot", "pi"),
"codex": ("codex", "copilot", "pi"),
"gemini": ("gemini", "opencode", "pi"),
"oss": ("opencode",),
"oss": ("opencode", "pi"),
}


Expand Down Expand Up @@ -337,7 +337,7 @@ def configure_shared_state(
)
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_oss = fetch_all or "opencode" in tools or "pi" in tools

claude_reason: str | None = None
gemini_reason: str | None = None
Expand Down
6 changes: 4 additions & 2 deletions src/ucode/databricks.py
Original file line number Diff line number Diff line change
Expand Up @@ -2135,12 +2135,14 @@ def build_pi_base_urls(workspace: str) -> dict[str, str]:
# - 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`).
# OSS families (kimi, glm) speak openai-completions to the MLflow route;
# compat flags applied per-provider in agents/pi.py (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
79 changes: 71 additions & 8 deletions tests/test_agent_pi.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,12 +20,17 @@ def _base_urls() -> dict[str, str]:
}


def _base_urls_with_oss() -> dict[str, str]:
return {**_base_urls(), "oss": f"{WS}/ai-gateway/mlflow/v1"}


def _empty() -> dict:
"""No-models input bundle for render_overlay."""
return {
"claude_models": {},
"codex_models": [],
"gemini_models": [],
"oss_models": [],
}


Expand All @@ -35,10 +40,11 @@ def _overlay(model: str, token: str = "tok", **kwargs):
return pi.render_overlay(
model,
token,
_base_urls(),
_base_urls_with_oss(),
bundle["claude_models"],
bundle["codex_models"],
bundle["gemini_models"],
bundle["oss_models"],
)


Expand Down Expand Up @@ -81,17 +87,25 @@ def test_gemini_provider_uses_google_generative_ai(self):
assert provider["api"] == "google-generative-ai"
assert provider["baseUrl"] == f"{WS}/ai-gateway/gemini/v1beta"

def test_all_three_providers_when_all_present(self):
def test_oss_provider_uses_openai_completions(self):
overlay, _ = _overlay("system.ai.kimi-k2-7-code", oss_models=["system.ai.kimi-k2-7-code"])
provider = overlay["providers"]["databricks-oss"]
assert provider["api"] == "openai-completions"
assert provider["baseUrl"] == f"{WS}/ai-gateway/mlflow/v1"

def test_all_providers_when_all_present(self):
overlay, _ = _overlay(
"claude-sonnet",
claude_models={"sonnet": "claude-sonnet"},
codex_models=["gpt-5"],
gemini_models=["gemini-2"],
oss_models=["system.ai.kimi-k2-7-code"],
)
assert set(overlay["providers"].keys()) == {
"databricks-claude",
"databricks-openai",
"databricks-gemini",
"databricks-oss",
}


Expand Down Expand Up @@ -129,6 +143,14 @@ def test_openai_and_gemini_have_no_compat_flags(self):
assert "compat" not in overlay["providers"]["databricks-openai"]
assert "compat" not in overlay["providers"]["databricks-gemini"]

def test_oss_disables_store_and_strict_mode(self):
# The MLflow chat-completions route rejects `store` and
# `tools[].function.strict`.
overlay, _ = _overlay("system.ai.kimi-k2-7-code", oss_models=["system.ai.kimi-k2-7-code"])
compat = overlay["providers"]["databricks-oss"]["compat"]
assert compat["supportsStore"] is False
assert compat["supportsStrictMode"] is False


class TestRenderOverlayAuthAndModels:
def test_token_in_api_key(self):
Expand Down Expand Up @@ -163,6 +185,29 @@ def test_gemini_models_listed(self):
ids = {m["id"] for m in overlay["providers"]["databricks-gemini"]["models"]}
assert ids == {"gemini-2", "gemini-2-pro"}

def test_oss_models_listed(self):
oss = ["system.ai.kimi-k2-7-code", "system.ai.glm-5-2"]
overlay, _ = _overlay("system.ai.kimi-k2-7-code", oss_models=oss)
ids = {m["id"] for m in overlay["providers"]["databricks-oss"]["models"]}
assert ids == {"system.ai.kimi-k2-7-code", "system.ai.glm-5-2"}

def test_oss_glm_carries_token_limits(self):
# GLM's output is capped well below pi's 16k default on the MLflow route;
# pin contextWindow/maxTokens from the shared limits table.
overlay, _ = _overlay("system.ai.glm-5-2", oss_models=["system.ai.glm-5-2"])
entry = next(
m for m in overlay["providers"]["databricks-oss"]["models"] if "glm" in m["id"]
)
assert entry["contextWindow"] == 200_000
assert entry["maxTokens"] == 25_000

def test_oss_kimi_omits_limits_when_unknown(self):
# Kimi has no entry in the limits table, so pi uses its own defaults.
overlay, _ = _overlay("system.ai.kimi-k2-7-code", oss_models=["system.ai.kimi-k2-7-code"])
entry = overlay["providers"]["databricks-oss"]["models"][0]
assert "contextWindow" not in entry
assert "maxTokens" not in entry


class TestRenderOverlayManagedKeys:
def test_managed_keys_include_model(self):
Expand Down Expand Up @@ -193,6 +238,10 @@ def test_prefixes_gemini_model(self):
overlay, _ = _overlay("gemini-2", gemini_models=["gemini-2"])
assert overlay["model"] == "databricks-gemini/gemini-2"

def test_prefixes_oss_model(self):
overlay, _ = _overlay("system.ai.kimi-k2-7-code", oss_models=["system.ai.kimi-k2-7-code"])
assert overlay["model"] == "databricks-oss/system.ai.kimi-k2-7-code"

def test_preserves_already_prefixed_model(self):
overlay, _ = _overlay(
"databricks-claude/claude-sonnet",
Expand Down Expand Up @@ -228,10 +277,22 @@ def test_falls_back_to_gemini(self):
state = {"claude_models": {}, "codex_models": [], "gemini_models": ["gemini-2"]}
assert pi.default_model(state) == "gemini-2"

def test_falls_back_to_oss(self):
state = {
"claude_models": {},
"codex_models": [],
"gemini_models": [],
"oss_models": ["system.ai.kimi-k2-7-code"],
}
assert pi.default_model(state) == "system.ai.kimi-k2-7-code"

def test_returns_none_when_empty(self):
assert pi.default_model({}) is None
assert (
pi.default_model({"claude_models": {}, "codex_models": [], "gemini_models": []}) is None
pi.default_model(
{"claude_models": {}, "codex_models": [], "gemini_models": [], "oss_models": []}
)
is None
)


Expand Down Expand Up @@ -279,10 +340,11 @@ def _setup(self, tmp_path, monkeypatch):
def _state(self, **overrides) -> dict:
state = {
"workspace": WS,
"base_urls": {"pi": _base_urls()},
"base_urls": {"pi": _base_urls_with_oss()},
"claude_models": {"sonnet": "claude-sonnet"},
"codex_models": [],
"gemini_models": [],
"oss_models": [],
"managed_configs": {},
}
state.update(overrides)
Expand Down Expand Up @@ -315,8 +377,9 @@ def test_stale_managed_providers_removed_before_merge(self, tmp_path, monkeypatc

def test_legacy_providers_removed_on_upgrade(self, tmp_path, monkeypatch):
"""Earlier ucode versions wrote `databricks-anthropic`, `databricks-codex`,
and `databricks-oss` providers. They must be stripped on the next write
so users don't end up with stale entries pointing at routes that 400."""
and `databricks-kimi` providers. They must be stripped on the next write
so users don't end up with stale entries pointing at routes that 400 (or
carrying a frozen, never-refreshed token)."""
pi_mod, config_file, _, _ = self._setup(tmp_path, monkeypatch)

config_file.write_text(
Expand All @@ -325,7 +388,7 @@ def test_legacy_providers_removed_on_upgrade(self, tmp_path, monkeypatch):
"providers": {
"databricks-anthropic": {"api": "anthropic-messages"},
"databricks-codex": {"api": "openai-responses"},
"databricks-oss": {"api": "openai-completions"},
"databricks-kimi": {"api": "openai-responses"},
}
}
),
Expand All @@ -339,7 +402,7 @@ def test_legacy_providers_removed_on_upgrade(self, tmp_path, monkeypatch):
pi_mod.write_tool_config(self._state(), "claude-sonnet", token="tok")

written_providers = json.loads(config_file.read_text()).get("providers", {})
for legacy in ("databricks-anthropic", "databricks-codex", "databricks-oss"):
for legacy in ("databricks-anthropic", "databricks-codex", "databricks-kimi"):
assert legacy not in written_providers
assert "databricks-claude" in written_providers

Expand Down
Loading