AI Providers

This commit is contained in:
Renn F
2026-05-01 03:09:11 +02:00
parent d15b7ae561
commit 7970020021
11 changed files with 2149 additions and 0 deletions
+232
View File
@@ -0,0 +1,232 @@
"""Provider routing: provider_configs + model_assignments.
Revision ID: 004_provider_routing
Revises: 003_blocker_resolver_type
Create Date: 2026-04-21
Adds two tables so agents can be routed per-role / per-agent to different
model providers (Anthropic via mounted ~/.claude, Ollama Cloud via
ANTHROPIC_BASE_URL+AUTH_TOKEN env injection).
- `provider_configs`: one row per logical provider. The Anthropic default
is seeded with no base_url / no token — it's a pointer-only row; auth
stays in the agent container's mounted ~/.claude.
- `model_assignments`: scope (`global` | `role` | `agent_slug`) →
(provider, model_name). Unique on (scope, scope_value) with NULLS NOT
DISTINCT so the single `global` row can't be duplicated.
Zero rows in `model_assignments` leaves every spawn on the legacy
`ROLE_MODEL_MAP` path — fully backward-compatible.
"""
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision = "004_provider_routing"
down_revision = "003_blocker_resolver_type"
branch_labels = None
depends_on = None
# SQLAlchemy's default Enum binding serialises members by Python NAME
# (uppercase), matching the 003_blocker_resolver_type convention.
_PROVIDER_TYPES = ("ANTHROPIC", "OLLAMA_CLOUD", "OPENAI", "LOCAL")
_ASSIGNMENT_SCOPES = ("GLOBAL", "ROLE", "AGENT_SLUG")
def upgrade() -> None:
# `checkfirst=True` on SA's ENUM creation has been unreliable across
# PG versions; use raw `DO $$` + `EXCEPTION duplicate_object` so the
# migration is safely re-runnable on a DB that already has the types
# from a prior half-applied attempt.
op.execute(
sa.text(
"""
DO $$ BEGIN
CREATE TYPE modelprovider AS ENUM (
'ANTHROPIC', 'OLLAMA_CLOUD', 'OPENAI', 'LOCAL'
);
EXCEPTION WHEN duplicate_object THEN NULL;
END $$;
"""
)
)
op.execute(
sa.text(
"""
DO $$ BEGIN
CREATE TYPE assignmentscope AS ENUM (
'GLOBAL', 'ROLE', 'AGENT_SLUG'
);
EXCEPTION WHEN duplicate_object THEN NULL;
END $$;
"""
)
)
# Reference the pre-existing enums from the table columns without
# re-emitting their DDL. `create_type=False` is the critical flag.
provider_enum = postgresql.ENUM(
*_PROVIDER_TYPES, name="modelprovider", create_type=False
)
scope_enum = postgresql.ENUM(
*_ASSIGNMENT_SCOPES, name="assignmentscope", create_type=False
)
# Defensive: if a prior attempt half-created these tables, drop them
# clean before re-creating. Enum types survive this (they're owned by
# the database, not the tables).
op.execute(sa.text("DROP TABLE IF EXISTS model_assignments CASCADE"))
op.execute(sa.text("DROP TABLE IF EXISTS provider_configs CASCADE"))
op.create_table(
"provider_configs",
sa.Column(
"id",
postgresql.UUID(as_uuid=True),
primary_key=True,
nullable=False,
),
sa.Column("name", sa.String(length=100), nullable=False),
sa.Column("type", provider_enum, nullable=False),
sa.Column("base_url", sa.Text(), nullable=True),
sa.Column("auth_token_encrypted", sa.Text(), nullable=True),
sa.Column(
"enabled",
sa.Boolean(),
nullable=False,
server_default=sa.text("true"),
),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=True,
),
sa.UniqueConstraint("name", name="uq_provider_configs_name"),
)
op.create_index(
"ix_provider_configs_name",
"provider_configs",
["name"],
unique=False,
)
op.create_index(
"ix_provider_configs_enabled",
"provider_configs",
["enabled"],
unique=False,
)
op.create_table(
"model_assignments",
sa.Column(
"id",
postgresql.UUID(as_uuid=True),
primary_key=True,
nullable=False,
),
sa.Column("scope", scope_enum, nullable=False),
sa.Column("scope_value", sa.String(length=100), nullable=True),
sa.Column(
"provider_config_id",
postgresql.UUID(as_uuid=True),
nullable=False,
),
sa.Column("model_name", sa.String(length=100), nullable=False),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=True,
),
sa.ForeignKeyConstraint(
["provider_config_id"],
["provider_configs.id"],
ondelete="RESTRICT",
),
)
# NULLS NOT DISTINCT is PG 15+. roboco runs pgvector on PG 16, so fine.
# This stops the `global` row (scope_value=NULL) from being duplicated.
op.create_index(
"ux_model_assignments_scope_key",
"model_assignments",
["scope", "scope_value"],
unique=True,
postgresql_nulls_not_distinct=True,
)
op.create_index(
"ix_model_assignments_provider",
"model_assignments",
["provider_config_id"],
unique=False,
)
# Seed BOTH providers so the Settings UI is zero-setup: the user never
# "creates a provider" — they just pick a mode and (if using Ollama)
# paste the API key.
#
# * Anthropic — pointer-only. base_url + token stay NULL; spawn-time
# env injection is skipped and the container uses its mounted
# ~/.claude auth just like today.
#
# * Ollama Cloud — pre-seeded disabled. The key-input endpoint
# (`PUT /api/v1/providers/ollama-key`) flips enabled=true and
# stores the Fernet-encrypted token when the user saves their key.
op.execute(
sa.text(
"""
INSERT INTO provider_configs
(id, name, type, base_url, auth_token_encrypted, enabled, created_at)
VALUES
(
gen_random_uuid(),
'Anthropic (default)',
'ANTHROPIC',
NULL,
NULL,
true,
now()
),
(
gen_random_uuid(),
'Ollama Cloud',
'OLLAMA_CLOUD',
'https://ollama.com',
NULL,
false,
now()
)
"""
)
)
def downgrade() -> None:
# `if_exists=True` makes the downgrade safe on a partially-applied
# schema (e.g., an earlier upgrade that half-succeeded), since
# alembic `drop_table` doesn't take a checkfirst flag directly.
op.execute("DROP INDEX IF EXISTS ix_model_assignments_provider")
op.execute("DROP INDEX IF EXISTS ux_model_assignments_scope_key")
op.execute("DROP TABLE IF EXISTS model_assignments")
op.execute("DROP INDEX IF EXISTS ix_provider_configs_enabled")
op.execute("DROP INDEX IF EXISTS ix_provider_configs_name")
op.execute("DROP TABLE IF EXISTS provider_configs")
postgresql.ENUM(name="assignmentscope").drop(
op.get_bind(), checkfirst=True
)
postgresql.ENUM(name="modelprovider").drop(op.get_bind(), checkfirst=True)
+41
View File
@@ -0,0 +1,41 @@
"""Add blocker_raised_by to tasks.
Revision ID: 005_blocker_raised_by
Revises: 004_provider_routing
Create Date: 2026-04-22
Adds `tasks.blocker_raised_by` so `roboco_task_unblock` can restore the
task to the agent who actually raised the block/escalation. Without this,
escalations (which reassign the task to the escalation target for
resolution) leave the dev's identity lost — unblocking just flips status
back to `in_progress` with the PM still on the hook, so the orchestrator
never respawns the original dev and the task stalls.
NULL = never blocked, or legacy rows pre-migration. `unblock` treats NULL
as "no-op on assignee" to preserve back-compat.
"""
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "005_blocker_raised_by"
down_revision = "004_provider_routing"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"tasks",
sa.Column(
"blocker_raised_by",
sa.dialects.postgresql.UUID(as_uuid=True),
sa.ForeignKey("agents.id", ondelete="SET NULL"),
nullable=True,
),
)
def downgrade() -> None:
op.drop_column("tasks", "blocker_raised_by")
@@ -0,0 +1,16 @@
import { AIRoutingCard } from "@/components/settings/ai-routing-card";
export default function AIProvidersPage() {
return (
<div className="space-y-6 max-w-5xl">
<div>
<h1 className="text-3xl font-bold tracking-tight">AI Providers</h1>
<p className="text-muted-foreground">
Pick how roboco agents authenticate and which model each one runs on.
</p>
</div>
<AIRoutingCard />
</div>
);
}
@@ -0,0 +1,404 @@
"use client";
import { useEffect, useMemo, useState } from "react";
import {
useApplyMode,
useCatalog,
useOllamaKey,
useRoutingMode,
useSetOllamaKey,
} from "@/hooks/use-providers";
import {
Card,
CardContent,
CardDescription,
CardHeader,
CardTitle,
} from "@/components/ui/card";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import {
Select,
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Separator } from "@/components/ui/separator";
import {
AlertTriangle,
Cpu,
Key,
KeyRound,
ShieldCheck,
Sparkles,
} from "lucide-react";
import { toast } from "sonner";
import { AssignmentScope, ModelProvider } from "@/types";
import type { RoutingMode } from "@/lib/api/providers";
// Matches the roboco agents_config AGENT_ROLE_MAP / AGENT_TEAM_MAP.
// Hard-coded so Mix mode shows a stable 18-row picker without an extra
// server round-trip. Order mirrors the org chart in CLAUDE.md.
//
// NOTE: CEO is explicitly excluded — it's the human-in-the-loop seat
// (Renzo), not an LLM-backed agent. Routing it anywhere would be a
// no-op in spawn_agent but confusing in the UI.
const AGENTS: { slug: string; label: string }[] = [
{ slug: "product-owner", label: "Product Owner" },
{ slug: "head-marketing", label: "Head of Marketing" },
{ slug: "auditor", label: "Auditor" },
{ slug: "main-pm", label: "Main PM" },
{ slug: "be-pm", label: "Backend PM" },
{ slug: "be-dev-1", label: "Backend Dev 1" },
{ slug: "be-dev-2", label: "Backend Dev 2" },
{ slug: "be-qa", label: "Backend QA" },
{ slug: "be-doc", label: "Backend Documenter" },
{ slug: "fe-pm", label: "Frontend PM" },
{ slug: "fe-dev-1", label: "Frontend Dev 1" },
{ slug: "fe-dev-2", label: "Frontend Dev 2" },
{ slug: "fe-qa", label: "Frontend QA" },
{ slug: "fe-doc", label: "Frontend Documenter" },
{ slug: "ux-pm", label: "UX/UI PM" },
{ slug: "ux-dev-1", label: "UX/UI Dev" },
{ slug: "ux-qa", label: "UX/UI QA" },
{ slug: "ux-doc", label: "UX/UI Documenter" },
];
export function AIRoutingCard() {
const { data: catalog = [] } = useCatalog();
const { data: keyStatus } = useOllamaKey();
const { data: snapshot } = useRoutingMode();
const setKey = useSetOllamaKey();
const applyMode = useApplyMode();
const hasOllamaKey = !!keyStatus?.has_key;
const currentMode: RoutingMode = snapshot?.mode ?? "anthropic";
// --- API key input ---
const [apiKey, setApiKey] = useState("");
const [clearKey, setClearKey] = useState(false);
const saveKey = async () => {
try {
if (clearKey) {
await setKey.mutateAsync("");
toast.success("Ollama key cleared");
} else {
if (!apiKey.trim()) {
toast.error("Enter a key first");
return;
}
await setKey.mutateAsync(apiKey);
toast.success("Ollama key saved");
}
setApiKey("");
setClearKey(false);
} catch (e) {
toast.error("Save failed: " + errMsg(e));
}
};
// --- Mix mode state: agent_slug → model_name ---
const initialMix = useMemo(() => {
const map: Record<string, string> = {};
for (const a of snapshot?.assignments ?? []) {
if (a.scope === AssignmentScope.AGENT_SLUG && a.scope_value) {
map[a.scope_value] = a.model_name;
}
}
return map;
}, [snapshot]);
const [mixMap, setMixMap] = useState<Record<string, string>>(initialMix);
useEffect(() => {
// Reset local state when server returns a fresh snapshot (after save,
// mode switch, or initial load).
setMixMap(initialMix);
}, [initialMix]);
const catalogForMix = catalog;
const catalogOllamaOnly = catalog.filter(
(c: { provider_type: ModelProvider }) => c.provider_type === ModelProvider.OLLAMA_CLOUD,
);
// --- Mode toggle handlers ---
const flipToAnthropic = async () => {
if (!confirm("Switch every agent to Anthropic? Clears any overrides.")) return;
try {
await applyMode.mutateAsync({ mode: "anthropic" });
toast.success("All agents now on Anthropic");
} catch (e) {
toast.error("Switch failed: " + errMsg(e));
}
};
const flipToOllama = async () => {
if (!hasOllamaKey) {
toast.error("Save an Ollama API key first");
return;
}
if (!confirm("Switch every agent to Ollama? Clears any overrides.")) return;
try {
await applyMode.mutateAsync({ mode: "ollama" });
toast.success("All agents now on Ollama");
} catch (e) {
toast.error("Switch failed: " + errMsg(e));
}
};
const saveMix = async () => {
// Filter out empty picks (nothing selected = inherit global).
const per_agent: Record<string, string> = {};
for (const [slug, model] of Object.entries(mixMap)) {
if (model) per_agent[slug] = model;
}
if (Object.keys(per_agent).length === 0) {
toast.error("Pick a model for at least one agent");
return;
}
const needsKey = Object.values(per_agent).some((m) =>
catalog.find(
(c: { model_name: string; provider_type: ModelProvider }) =>
c.model_name === m &&
c.provider_type === ModelProvider.OLLAMA_CLOUD,
),
);
if (needsKey && !hasOllamaKey) {
toast.error(
"At least one agent is routed to an Ollama model but no key is saved",
);
return;
}
try {
await applyMode.mutateAsync({ mode: "mix", per_agent });
toast.success("Per-agent routing saved");
} catch (e) {
toast.error("Save failed: " + errMsg(e));
}
};
return (
<Card>
<CardHeader>
<CardTitle className="flex items-center gap-2">
<Cpu className="h-5 w-5" /> AI Routing
</CardTitle>
<CardDescription>
Decide which model backs each agent. Anthropic uses the mounted
<code className="px-1"> ~/.claude </code> auth; Ollama Cloud uses
the API key you save below.
</CardDescription>
</CardHeader>
<CardContent className="space-y-6">
{/* -------- Ollama key -------- */}
<section className="space-y-2">
<div className="flex items-center justify-between">
<Label className="text-sm font-medium">Ollama Cloud API key</Label>
{hasOllamaKey ? (
<span className="inline-flex items-center gap-1 rounded-full bg-emerald-500/10 px-2 py-0.5 text-xs font-medium text-emerald-600">
<KeyRound className="h-3 w-3" /> key set
</span>
) : (
<span className="inline-flex items-center gap-1 rounded-full bg-amber-500/10 px-2 py-0.5 text-xs font-medium text-amber-600">
<Key className="h-3 w-3" /> not set
</span>
)}
</div>
<div className="flex gap-2">
<Input
type="password"
value={apiKey}
onChange={(e) => setApiKey(e.target.value)}
placeholder={
hasOllamaKey ? "•••••••••••• (leave blank to keep)" : "ollama_xxx…"
}
disabled={clearKey}
/>
<Button onClick={saveKey} disabled={setKey.isPending}>
{setKey.isPending ? "Saving…" : "Save"}
</Button>
</div>
{hasOllamaKey ? (
<label className="flex items-center gap-2 text-xs text-muted-foreground">
<input
type="checkbox"
checked={clearKey}
onChange={(e) => {
setClearKey(e.target.checked);
if (e.target.checked) setApiKey("");
}}
/>
Clear the stored key
</label>
) : (
<p className="text-xs text-muted-foreground">
Stored Fernet-encrypted server-side; never returned by the API.
</p>
)}
</section>
<Separator />
{/* -------- Mode toggle -------- */}
<section className="space-y-3">
<Label className="text-sm font-medium">Routing mode</Label>
<div className="grid grid-cols-1 md:grid-cols-3 gap-2">
<ModeButton
icon={<ShieldCheck className="h-4 w-4" />}
label="Anthropic"
description="Every agent uses Anthropic (via mounted ~/.claude)."
active={currentMode === "anthropic"}
onClick={flipToAnthropic}
disabled={applyMode.isPending}
/>
<ModeButton
icon={<Sparkles className="h-4 w-4" />}
label="Ollama"
description={
hasOllamaKey
? "Every agent uses Ollama Cloud (Minimax M2.7 default)."
: "Save the Ollama key first."
}
active={currentMode === "ollama"}
onClick={flipToOllama}
disabled={applyMode.isPending || !hasOllamaKey}
/>
<ModeButton
icon={<Cpu className="h-4 w-4" />}
label="Mix"
description="Pick a model per agent (table appears below)."
active={currentMode === "mix"}
// "Mix" is engaged by picking models + Save — not a direct
// toggle. Clicking just scrolls awareness.
onClick={() => undefined}
disabled={false}
highlight={currentMode === "mix"}
/>
</div>
{currentMode === "mix" && !hasOllamaKey ? (
<p className="text-xs text-amber-600 flex items-center gap-1">
<AlertTriangle className="h-3 w-3" />
Some agents may already be routed to Ollama but no key is
saved those agents will fall back to Anthropic at spawn.
</p>
) : null}
</section>
{/* -------- Mix-mode per-agent picker -------- */}
<Separator />
<section className="space-y-3">
<div className="flex items-center justify-between">
<Label className="text-sm font-medium">
Per-agent override (mix mode)
</Label>
<Button
size="sm"
onClick={saveMix}
disabled={applyMode.isPending}
>
{applyMode.isPending ? "Saving…" : "Save mix"}
</Button>
</div>
<p className="text-xs text-muted-foreground">
Leave a row blank to inherit from the global mode. Saving
overwrites all per-agent overrides with what&apos;s picked here.
</p>
<div className="divide-y rounded-md border">
{AGENTS.map((a) => (
<div
key={a.slug}
className="grid grid-cols-[1fr_280px] items-center gap-2 px-3 py-2"
>
<div>
<div className="font-mono text-sm">{a.slug}</div>
<div className="text-xs text-muted-foreground">
{a.label}
</div>
</div>
<Select
value={mixMap[a.slug] ?? ""}
onValueChange={(v: string) =>
setMixMap((prev) => ({
...prev,
[a.slug]: v === "__clear__" ? "" : v,
}))
}
>
<SelectTrigger>
<SelectValue placeholder="(inherit)" />
</SelectTrigger>
<SelectContent>
<SelectItem value="__clear__">(inherit global)</SelectItem>
{catalogForMix.map((c: { model_name: string; display_name: string }) => (
<SelectItem key={c.model_name} value={c.model_name}>
{c.display_name} {c.model_name}
</SelectItem>
))}
</SelectContent>
</Select>
</div>
))}
</div>
{catalogOllamaOnly.length === 0 ? (
<p className="text-xs text-muted-foreground">
Ollama catalog empty check /api/v1/providers/catalog.
</p>
) : null}
</section>
</CardContent>
</Card>
);
}
// ---------------------------------------------------------------------------
function ModeButton({
icon,
label,
description,
active,
disabled,
onClick,
highlight,
}: {
icon: React.ReactNode;
label: string;
description: string;
active: boolean;
disabled: boolean;
onClick: () => void | Promise<void>;
highlight?: boolean;
}) {
return (
<button
type="button"
onClick={onClick}
disabled={disabled}
className={
"rounded-md border p-3 text-left transition-colors " +
(active || highlight
? "border-primary bg-primary/5"
: "hover:bg-muted") +
(disabled ? " cursor-not-allowed opacity-50" : "")
}
>
<div className="flex items-center gap-2 text-sm font-medium">
{icon}
<span>{label}</span>
{active ? (
<span className="ml-auto rounded-full bg-primary/15 px-2 py-0.5 text-xs text-primary">
active
</span>
) : null}
</div>
<p className="mt-1 text-xs text-muted-foreground">{description}</p>
</button>
);
}
function errMsg(e: unknown): string {
return e instanceof Error ? e.message : "Unknown error";
}
+58
View File
@@ -0,0 +1,58 @@
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query";
import {
providersApi,
type ApplyModePayload,
} from "@/lib/api/providers";
export const providerKeys = {
all: ["providers"] as const,
catalog: () => [...providerKeys.all, "catalog"] as const,
ollamaKey: () => [...providerKeys.all, "ollama-key"] as const,
mode: () => [...providerKeys.all, "mode"] as const,
};
export function useCatalog() {
return useQuery({
queryKey: providerKeys.catalog(),
queryFn: () => providersApi.catalog(),
staleTime: 5 * 60_000, // 5 minutes — static list
});
}
export function useOllamaKey() {
return useQuery({
queryKey: providerKeys.ollamaKey(),
queryFn: () => providersApi.getOllamaKey(),
staleTime: 60_000,
});
}
export function useSetOllamaKey() {
const qc = useQueryClient();
return useMutation({
mutationFn: (apiKey: string) => providersApi.setOllamaKey(apiKey),
onSuccess: () => {
qc.invalidateQueries({ queryKey: providerKeys.ollamaKey() });
// Applying a mode also reads this so refresh it too.
qc.invalidateQueries({ queryKey: providerKeys.mode() });
},
});
}
export function useRoutingMode() {
return useQuery({
queryKey: providerKeys.mode(),
queryFn: () => providersApi.getMode(),
staleTime: 30_000,
});
}
export function useApplyMode() {
const qc = useQueryClient();
return useMutation({
mutationFn: (payload: ApplyModePayload) => providersApi.applyMode(payload),
onSuccess: () => {
qc.invalidateQueries({ queryKey: providerKeys.mode() });
},
});
}
+162
View File
@@ -0,0 +1,162 @@
"""
Provider Routes
Thin HTTP plumbing for the Settings UI's AI-routing panel. Four endpoints
cover the whole UX: fetch the catalog, get / set the Ollama key, fetch the
current mode + assignments, apply a mode change. No provider CRUD — the
two providers (Anthropic + Ollama Cloud) are pre-seeded by migration 004.
"""
from fastapi import APIRouter, HTTPException, status
from roboco.api.deps import CurrentAgentContext, DbSession, require_pm_or_above
from roboco.api.schemas.provider import (
ApplyModeRequest,
CatalogEntryResponse,
ModeResponse,
OllamaKeyStatus,
SetOllamaKeyRequest,
assignment_to_response,
)
from roboco.models.base import ModelProvider
from roboco.models.llm_catalog import MODEL_CATALOG
from roboco.services.base import NotFoundError
from roboco.services.llm import get_model_routing_service
from roboco.services.provider import get_provider_service
router = APIRouter()
# =============================================================================
# CATALOG
# =============================================================================
@router.get("/catalog", response_model=list[CatalogEntryResponse])
async def get_catalog(
agent: CurrentAgentContext,
) -> list[CatalogEntryResponse]:
"""Return the preset list of selectable models.
Order matches display order in the UI. Static — served from constants
so the UI never needs to hit the DB for the model dropdown.
"""
require_pm_or_above(agent.role, "view the model catalog")
return [
CatalogEntryResponse(
model_name=entry.model_name,
provider_type=entry.provider_type,
display_name=entry.display_name,
)
for entry in MODEL_CATALOG
]
# =============================================================================
# OLLAMA API KEY (the one and only secret the user types)
# =============================================================================
@router.get("/ollama-key", response_model=OllamaKeyStatus)
async def get_ollama_key_status(
db: DbSession,
agent: CurrentAgentContext,
) -> OllamaKeyStatus:
"""Return whether the Ollama Cloud key is set + enabled."""
require_pm_or_above(agent.role, "view the Ollama key status")
provider_svc = get_provider_service(db)
providers = await provider_svc.list_providers(include_disabled=True)
ollama = next(
(p for p in providers if p.type == ModelProvider.OLLAMA_CLOUD),
None,
)
if ollama is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=("Ollama Cloud provider not seeded. Run alembic upgrade head."),
)
return OllamaKeyStatus(
has_key=bool(ollama.auth_token_encrypted),
enabled=ollama.enabled,
)
@router.put("/ollama-key", response_model=OllamaKeyStatus)
async def set_ollama_key(
data: SetOllamaKeyRequest,
db: DbSession,
agent: CurrentAgentContext,
) -> OllamaKeyStatus:
"""Set or clear the Ollama Cloud API key.
Empty string → clears and disables the provider. Any other value →
Fernet-encrypts + marks enabled. This is the only secret the user
types anywhere.
"""
require_pm_or_above(agent.role, "set the Ollama key")
routing = get_model_routing_service(db)
try:
provider = await routing.set_ollama_api_key(data.api_key)
except NotFoundError as e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(e)) from e
await db.commit()
return OllamaKeyStatus(
has_key=bool(provider.auth_token_encrypted),
enabled=provider.enabled,
)
# =============================================================================
# MODE (the three-way toggle)
# =============================================================================
@router.get("", response_model=ModeResponse)
async def get_current_mode(
db: DbSession,
agent: CurrentAgentContext,
) -> ModeResponse:
"""Return the current mode + all live assignments for UI rendering."""
require_pm_or_above(agent.role, "view routing state")
routing = get_model_routing_service(db)
mode = await routing.derive_mode()
assignments = await routing.list_assignments()
return ModeResponse(
mode=mode, # type: ignore[arg-type]
assignments=[assignment_to_response(a) for a in assignments],
)
@router.post("", response_model=ModeResponse)
async def apply_mode(
data: ApplyModeRequest,
db: DbSession,
agent: CurrentAgentContext,
) -> ModeResponse:
"""Apply a mode change atomically.
Returns the new assignment snapshot so the UI can re-render without a
second round-trip.
"""
require_pm_or_above(agent.role, "change routing mode")
routing = get_model_routing_service(db)
try:
await routing.apply_mode(
mode=data.mode,
default_model=data.default_model,
per_agent=data.per_agent,
)
except ValueError as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)
) from e
except NotFoundError as e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(e)) from e
await db.commit()
mode = await routing.derive_mode()
assignments = await routing.list_assignments()
return ModeResponse(
mode=mode, # type: ignore[arg-type]
assignments=[assignment_to_response(a) for a in assignments],
)
+114
View File
@@ -0,0 +1,114 @@
"""
Providers API Schemas
Minimal surface that backs the Settings UI:
- fetch the preset catalog of selectable models
- set / clear / check the single Ollama Cloud API key
- read current routing assignments (so the UI renders Mix mode)
- apply a routing mode (anthropic | ollama | mix)
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Literal
from uuid import UUID # noqa: TC003 (pydantic needs the type at runtime)
from pydantic import BaseModel, Field
from roboco.models.base import AssignmentScope, ModelProvider # noqa: TC001
from roboco.utils.converters import require_uuid
if TYPE_CHECKING:
from roboco.db.tables import ModelAssignmentTable
# =============================================================================
# CATALOG
# =============================================================================
class CatalogEntryResponse(BaseModel):
"""One selectable model in the Settings dropdown."""
model_name: str
provider_type: ModelProvider
display_name: str
# =============================================================================
# OLLAMA API KEY
# =============================================================================
class OllamaKeyStatus(BaseModel):
"""Whether the Ollama Cloud provider has a stored token."""
has_key: bool
enabled: bool
class SetOllamaKeyRequest(BaseModel):
"""Set or clear the Ollama Cloud API key.
Pass an empty string to clear. Pass a non-empty string to save
(encrypted with Fernet) and mark the Ollama provider enabled.
"""
api_key: str = Field(default="")
# =============================================================================
# MODEL ASSIGNMENTS (read-only for the UI)
# =============================================================================
class AssignmentResponse(BaseModel):
"""Routing rule with the provider summary flattened for UI rendering."""
id: UUID
scope: AssignmentScope
scope_value: str | None
provider_type: ModelProvider
model_name: str
def assignment_to_response(
row: ModelAssignmentTable,
) -> AssignmentResponse:
"""Convert a ModelAssignmentTable row + joined provider to a response."""
return AssignmentResponse(
id=require_uuid(row.id),
scope=row.scope,
scope_value=row.scope_value,
provider_type=row.provider.type,
model_name=row.model_name,
)
# =============================================================================
# MODE APPLY
# =============================================================================
class ApplyModeRequest(BaseModel):
"""Apply a routing mode in one atomic call.
- mode="anthropic": clear every assignment; spawns fall through to
ROLE_MODEL_MAP + mounted ~/.claude.
- mode="ollama": clear every assignment; set GLOBAL default to
`default_model` (if omitted, the service picks a sensible default).
- mode="mix": clear existing per-agent pins; upsert the `per_agent`
map verbatim. Role + GLOBAL rows are left untouched so the user can
layer with an existing partial setup.
"""
mode: Literal["anthropic", "ollama", "mix"]
default_model: str | None = None
per_agent: dict[str, str] | None = None
class ModeResponse(BaseModel):
"""Server-side view of the current mode + a snapshot of active rules."""
mode: Literal["anthropic", "ollama", "mix"]
assignments: list[AssignmentResponse]
+125
View File
@@ -0,0 +1,125 @@
"""
Model Catalog
Preset list of (model_name, provider_type, display_name) that the Settings
UI renders. Users never type a model name — they pick from this list, and
the router maps back to the correct pre-seeded provider row.
**Anthropic entries derive from `runtime.MODEL_MAP`** — that's the single
source of truth for which Claude versions are supported. Bumping a model
version there (e.g. `claude-opus-4-6` → `claude-opus-4-7`) updates the
catalog automatically without a second edit here.
Ollama Cloud entries are hand-maintained because Ollama's cloud tags
don't live in the rest of the codebase.
"""
from __future__ import annotations
from dataclasses import dataclass
from roboco.models.base import ModelProvider
from roboco.models.runtime import MODEL_MAP
@dataclass(frozen=True)
class CatalogEntry:
"""One selectable model in the Settings dropdown."""
model_name: str
provider_type: ModelProvider
display_name: str
# Display labels for the Anthropic short names. Order of this tuple is the
# render order in the UI dropdown; entries missing from MODEL_MAP are
# silently skipped so we don't expose a model we can't route to.
_ANTHROPIC_DISPLAY: tuple[tuple[str, str], ...] = (
("opus", "Claude Opus"),
("sonnet", "Claude Sonnet"),
("haiku", "Claude Haiku"),
)
def _build_anthropic_entries() -> tuple[CatalogEntry, ...]:
"""Expand MODEL_MAP into catalog entries with full version in the label.
Keeps the UI honest: a user picking "Claude Opus" sees exactly which
underlying Claude Code model id will be used at spawn
(e.g. "Claude Opus · claude-opus-4-6"), so version bumps in
`runtime.MODEL_MAP` are immediately visible.
"""
entries: list[CatalogEntry] = []
for short_name, label in _ANTHROPIC_DISPLAY:
full_id = MODEL_MAP.get(short_name)
if not full_id:
continue
entries.append(
CatalogEntry(
model_name=short_name,
provider_type=ModelProvider.ANTHROPIC,
display_name=f"{label} · {full_id}",
)
)
return tuple(entries)
MODEL_CATALOG: tuple[CatalogEntry, ...] = (
*_build_anthropic_entries(),
# --- Ollama Cloud (verbatim tags) ---
# Pro plan active as of 2026-04-22. Drop any entry that stops working —
# the catalog is the single source of truth the Settings dropdown renders from.
CatalogEntry("glm-5.1:cloud", ModelProvider.OLLAMA_CLOUD, "GLM 5.1"),
CatalogEntry("kimi-k2.6:cloud", ModelProvider.OLLAMA_CLOUD, "Kimi K2.6"),
CatalogEntry("minimax-m2.7:cloud", ModelProvider.OLLAMA_CLOUD, "Minimax M2.7"),
)
# Fast lookup by model_name. Enforces uniqueness at import time.
MODEL_CATALOG_BY_NAME: dict[str, CatalogEntry] = {
e.model_name: e for e in MODEL_CATALOG
}
assert len(MODEL_CATALOG_BY_NAME) == len(MODEL_CATALOG), "duplicate model_name"
def provider_type_for_model(model_name: str) -> ModelProvider | None:
"""Return the provider type for a catalog entry, or None if unknown."""
entry = MODEL_CATALOG_BY_NAME.get(model_name)
return entry.provider_type if entry else None
# Defaults per role when the user flips to "pure Ollama" mode.
# Assignments reflect the 2026-04 public benchmarks for each cloud tag:
# Kimi K2.6 — HLE 44.9%, AIME 95.6%, Agent Swarm (100 parallel sub-agents),
# 200-300 sequential tool calls. Best at reasoning, orchestration, tool use.
# MiniMax M2.7 — SWE-Bench 73.8%, SWE-Pro 56.2%, 10B active params (fastest,
# cheapest). Explicitly "built for Max coding & agentic workflows".
# GLM 5.1 — SWE-Bench 77.8% (highest of the three, 94.6% of Claude Opus 4.6),
# self-correcting across hundreds of iterations, strong creative writing.
OLLAMA_ROLE_DEFAULTS: dict[str, str] = {
# High-volume agentic coding — M2.7 is purpose-built for this.
"developer": "minimax-m2.7:cloud",
# Deep code review — GLM 5.1 has the highest SWE-Bench and iterates thoroughly.
"qa": "glm-5.1:cloud",
# Orchestration + tool coordination — Kimi K2.6's Agent Swarm is the exact fit.
"cell_pm": "kimi-k2.6:cloud",
"main_pm": "kimi-k2.6:cloud",
# Quality reasoning — Kimi K2.6 leads HLE by a wide margin.
"auditor": "kimi-k2.6:cloud",
# Product reasoning — same profile as PM work.
"product_owner": "kimi-k2.6:cloud",
# Writing with code-context — GLM 5.1's creative writing + SWE-Bench combo.
"documenter": "glm-5.1:cloud",
# Stylistic writing — GLM 5.1's creative-writing strength.
"head_marketing": "glm-5.1:cloud",
# CEO is human-in-the-loop; keep an entry in case someone forces
# a route to it, but the Settings UI intentionally excludes it.
"ceo": "kimi-k2.6:cloud",
}
# The Ollama model picked for "pure Ollama" mode's GLOBAL row when the
# caller doesn't override. Kimi K2.6 wins as the generalist because it has
# the strongest reasoning/tool-use profile and can fall back to coding/writing
# adequately if a role ends up mapped to the global default.
OLLAMA_DEFAULT_MODEL: str = "kimi-k2.6:cloud"
+369
View File
@@ -0,0 +1,369 @@
"""
Model Routing Service
Resolves (provider, model) for a given agent at spawn time using the
scoped rows in `model_assignments`:
AGENT_SLUG override > ROLE override > GLOBAL default
If none apply, falls back to the legacy `ROLE_MODEL_MAP` + implicit
Anthropic provider so deployments with zero rows behave exactly as
before. Decryption failures are contained: the service logs the error
and downgrades to the legacy path rather than failing the spawn.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, ClassVar, cast
from sqlalchemy import delete as sa_delete
from sqlalchemy import select
from roboco.agents_config import get_agent_role
from roboco.db.tables import ModelAssignmentTable, ProviderConfigTable
from roboco.models.base import AssignmentScope, ModelProvider
from roboco.models.llm_catalog import (
MODEL_CATALOG_BY_NAME,
OLLAMA_DEFAULT_MODEL,
)
from roboco.models.runtime import MODEL_MAP, ROLE_MODEL_MAP
from roboco.services.base import BaseService, NotFoundError
from roboco.services.provider import ProviderService, ProviderUpdate
from roboco.utils.converters import require_uuid
from roboco.utils.crypto import EncryptionError
if TYPE_CHECKING:
from uuid import UUID
from sqlalchemy.ext.asyncio import AsyncSession
@dataclass(frozen=True)
class AgentRoute:
"""Resolved routing for a single agent spawn.
`base_url` / `auth_token` being `None` means "Anthropic default":
orchestrator injects no `ANTHROPIC_*` env vars and the container
uses its mounted `~/.claude` auth (legacy behaviour).
"""
provider_id: UUID | None
provider_type: ModelProvider
base_url: str | None
auth_token: str | None
model_name: str
@dataclass(frozen=True)
class _ResolvedAssignment:
"""Internal — one resolved `model_assignments` row joined to provider."""
provider: ProviderConfigTable
model_name: str
class ModelRoutingService(BaseService):
"""Resolves per-agent routes from `model_assignments` + legacy fallback."""
service_name: ClassVar[str] = "model_routing"
async def resolve_for_agent(self, agent_slug: str) -> AgentRoute:
"""Resolve routing for `agent_slug` using the precedence ladder.
Never raises for a normal agent — decrypt failures and missing
agents both downgrade to the legacy Anthropic path, because a
stalled spawn is worse than a routing miss.
"""
role = get_agent_role(agent_slug) or ""
# 1) agent override
resolved = await self._find_assignment(
scope=AssignmentScope.AGENT_SLUG, scope_value=agent_slug
)
# 2) role override
if resolved is None and role:
resolved = await self._find_assignment(
scope=AssignmentScope.ROLE, scope_value=role
)
# 3) global default
if resolved is None:
resolved = await self._find_assignment(
scope=AssignmentScope.GLOBAL, scope_value=None
)
if resolved is not None and resolved.provider.enabled:
try:
return await self._route_from_assignment(resolved)
except EncryptionError:
self.log.error(
"Provider token decrypt failed; falling back to legacy path",
provider_id=str(resolved.provider.id),
agent_slug=agent_slug,
)
# 4) legacy fallback: role-default short name through MODEL_MAP.
short = ROLE_MODEL_MAP.get(role, "sonnet")
return AgentRoute(
provider_id=None,
provider_type=ModelProvider.ANTHROPIC,
base_url=None,
auth_token=None,
model_name=MODEL_MAP.get(short, short),
)
# =========================================================================
# ASSIGNMENT CRUD (consumed by api/routes/provider.py)
# =========================================================================
async def list_assignments(self) -> list[ModelAssignmentTable]:
result = await self.session.execute(
select(ModelAssignmentTable).order_by(
ModelAssignmentTable.scope, ModelAssignmentTable.scope_value
)
)
return list(result.scalars().all())
async def get_assignment(
self, *, scope: AssignmentScope, scope_value: str | None
) -> ModelAssignmentTable | None:
query = select(ModelAssignmentTable).where(ModelAssignmentTable.scope == scope)
if scope_value is None:
query = query.where(ModelAssignmentTable.scope_value.is_(None))
else:
query = query.where(ModelAssignmentTable.scope_value == scope_value)
result = await self.session.execute(query)
return result.scalar_one_or_none()
async def upsert_assignment(
self,
*,
scope: AssignmentScope,
scope_value: str | None,
model_name: str,
) -> ModelAssignmentTable:
"""Insert-or-update (by unique (scope, scope_value)).
Provider is derived from `MODEL_CATALOG` — the UI never picks a
provider separately, so the service looks up the pre-seeded
provider row for the catalog entry's type.
"""
self._validate_scope(scope, scope_value)
entry = MODEL_CATALOG_BY_NAME.get(model_name)
if entry is None:
raise ValueError(
f"Unknown model '{model_name}'. Use one from "
"GET /api/v1/providers/catalog."
)
provider = await self._get_seeded_provider(entry.provider_type)
row = await self.get_assignment(scope=scope, scope_value=scope_value)
if row is None:
row = ModelAssignmentTable(
scope=scope,
scope_value=scope_value,
provider_config_id=provider.id,
model_name=model_name,
)
self.session.add(row)
else:
row.provider_config_id = cast("Any", provider.id)
row.model_name = model_name
await self.session.flush()
self.log.info(
"Assignment upserted",
scope=scope.value,
scope_value=scope_value,
provider_type=entry.provider_type.value,
model_name=model_name,
)
return row
async def derive_mode(self) -> str:
"""Return the current "mode" label for the Settings UI.
Decision tree matches what `apply_mode` writes:
- no assignments at all → "anthropic"
- only a global row, Ollama → "ollama"
- anything else → "mix"
"""
assignments = await self.list_assignments()
if not assignments:
return "anthropic"
only_global = (
len(assignments) == 1 and assignments[0].scope == AssignmentScope.GLOBAL
)
is_ollama = (
only_global and assignments[0].provider.type == ModelProvider.OLLAMA_CLOUD
)
if is_ollama:
return "ollama"
return "mix"
async def set_ollama_api_key(self, api_key: str) -> ProviderConfigTable:
"""Set / clear the Ollama Cloud provider's API key.
Empty string clears + disables; a real key encrypts + enables.
Operates on the single pre-seeded Ollama row — no provider
creation happens here.
"""
provider = await self._get_seeded_provider(ModelProvider.OLLAMA_CLOUD)
provider_svc = ProviderService(self.session)
await provider_svc.update_provider(
require_uuid(provider.id),
ProviderUpdate(
auth_token=api_key if api_key else None,
clear_auth_token=not api_key,
enabled=bool(api_key),
),
)
# Re-fetch for the caller.
return await self._get_seeded_provider(ModelProvider.OLLAMA_CLOUD)
async def _get_seeded_provider(
self, provider_type: ModelProvider
) -> ProviderConfigTable:
"""Find the single seeded provider row for `provider_type`.
Migration `004_provider_routing` seeds exactly one row per type —
we just look it up. Raises NotFoundError if the seed is missing
(e.g., migration hasn't been applied).
"""
result = await self.session.execute(
select(ProviderConfigTable).where(ProviderConfigTable.type == provider_type)
)
row = result.scalar_one_or_none()
if row is None:
raise NotFoundError(
resource_type="Provider",
resource_id=f"type={provider_type.value}",
)
return row
async def delete_assignment(
self, *, scope: AssignmentScope, scope_value: str | None
) -> None:
row = await self.get_assignment(scope=scope, scope_value=scope_value)
if row is None:
raise NotFoundError(
resource_type="ModelAssignment",
resource_id=f"{scope.value}:{scope_value or '-'}",
)
await self.session.delete(row)
await self.session.flush()
self.log.info(
"Assignment deleted",
scope=scope.value,
scope_value=scope_value,
)
async def apply_mode(
self,
*,
mode: str,
default_model: str | None = None,
per_agent: dict[str, str] | None = None,
) -> None:
"""Apply a routing "mode" in a single transactional call.
Modes:
- "anthropic": wipe all assignments so every spawn falls through
to the legacy ROLE_MODEL_MAP + mounted ~/.claude path.
- "ollama": wipe role/agent overrides, set GLOBAL to the given
Ollama model (default: Kimi K2.6). CEO-type pins can be layered
back manually if the user wants them.
- "mix": apply per-agent map verbatim. Any agent not in the
map falls through to the GLOBAL default — which is whatever it
was (preserves prior state).
"""
if mode == "anthropic":
await self.session.execute(sa_delete(ModelAssignmentTable))
await self.session.flush()
self.log.info("Mode applied: anthropic (all assignments cleared)")
return
if mode == "ollama":
await self.session.execute(sa_delete(ModelAssignmentTable))
await self.session.flush()
await self.upsert_assignment(
scope=AssignmentScope.GLOBAL,
scope_value=None,
model_name=default_model or OLLAMA_DEFAULT_MODEL,
)
self.log.info(
"Mode applied: ollama",
default_model=default_model or OLLAMA_DEFAULT_MODEL,
)
return
if mode == "mix":
if not per_agent:
raise ValueError("mix mode requires a per_agent map")
# Clear existing agent-slug overrides so the new map is
# authoritative; leave role + global alone.
await self.session.execute(
sa_delete(ModelAssignmentTable).where(
ModelAssignmentTable.scope == AssignmentScope.AGENT_SLUG
)
)
await self.session.flush()
for agent_slug, model_name in per_agent.items():
if not model_name:
continue
await self.upsert_assignment(
scope=AssignmentScope.AGENT_SLUG,
scope_value=agent_slug,
model_name=model_name,
)
self.log.info("Mode applied: mix", agents=len(per_agent))
return
raise ValueError(f"Unknown mode '{mode}'. Use 'anthropic', 'ollama', or 'mix'.")
# =========================================================================
# INTERNAL
# =========================================================================
async def _find_assignment(
self, *, scope: AssignmentScope, scope_value: str | None
) -> _ResolvedAssignment | None:
row = await self.get_assignment(scope=scope, scope_value=scope_value)
if row is None:
return None
# Relationship is lazy="joined" in the ORM so `.provider` is loaded.
return _ResolvedAssignment(provider=row.provider, model_name=row.model_name)
async def _route_from_assignment(self, resolved: _ResolvedAssignment) -> AgentRoute:
provider = resolved.provider
# Decrypt only when the provider has a stored token (ollama_cloud).
# Anthropic providers have `auth_token_encrypted=NULL` and use the
# container's mounted credentials — no env injection needed.
provider_uuid = require_uuid(provider.id)
auth_token: str | None = None
if provider.auth_token_encrypted:
provider_svc = ProviderService(self.session)
auth_token = await provider_svc.get_decrypted_token(provider_uuid)
return AgentRoute(
provider_id=provider_uuid,
provider_type=provider.type,
base_url=provider.base_url,
auth_token=auth_token,
model_name=resolved.model_name,
)
@staticmethod
def _validate_scope(scope: AssignmentScope, scope_value: str | None) -> None:
if scope == AssignmentScope.GLOBAL and scope_value is not None:
raise ValueError("global scope must have scope_value=None")
if (
scope in (AssignmentScope.ROLE, AssignmentScope.AGENT_SLUG)
and not scope_value
):
raise ValueError(f"{scope.value} scope requires a non-empty scope_value")
def get_model_routing_service(session: AsyncSession) -> ModelRoutingService:
"""Get a ModelRoutingService instance."""
return ModelRoutingService(session)
+217
View File
@@ -0,0 +1,217 @@
"""
Provider Service
CRUD for `provider_configs` rows (logical model-provider connections).
Mirrors the Fernet encryption patterns in `ProjectService` — empty-string
clears the token, `None` leaves unchanged, non-empty re-encrypts.
API route modules translate their pydantic request models into the
dataclasses defined here at the boundary, so services never depend on
api.schemas.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, ClassVar
from sqlalchemy import select
from roboco.db.tables import ModelAssignmentTable, ProviderConfigTable
from roboco.services.base import BaseService, ConflictError, NotFoundError
from roboco.utils.crypto import EncryptionError, decrypt_token, encrypt_token
if TYPE_CHECKING:
from uuid import UUID
from sqlalchemy.ext.asyncio import AsyncSession
from roboco.models.base import ModelProvider
@dataclass(frozen=True)
class ProviderCreate:
"""Service-side shape for creating a provider config."""
name: str
type: ModelProvider
base_url: str | None = None
auth_token: str | None = None # plaintext; encrypted before persist
enabled: bool = True
@dataclass(frozen=True)
class ProviderUpdate:
"""Service-side shape for updating a provider config.
`auth_token` is tri-state: `None` leaves unchanged, `""` clears, any
other value re-encrypts. Matches `ProjectService.update` semantics.
"""
name: str | None = None
base_url: str | None = None
# `_SENTINEL` is the marker for "no change"; routes translate their
# pydantic model — with Python `None` as both "unset" and "clear to
# NULL" depending on field — into explicit values.
auth_token: str | None = None
clear_auth_token: bool = False # if True, force token → NULL
enabled: bool | None = None
class ProviderService(BaseService):
"""Manages `provider_configs` rows + their Fernet-encrypted tokens."""
service_name: ClassVar[str] = "provider"
# =========================================================================
# QUERIES
# =========================================================================
async def list_providers(
self, *, include_disabled: bool = False
) -> list[ProviderConfigTable]:
query = select(ProviderConfigTable)
if not include_disabled:
query = query.where(ProviderConfigTable.enabled.is_(True))
query = query.order_by(ProviderConfigTable.name)
result = await self.session.execute(query)
return list(result.scalars().all())
async def get_provider(self, provider_id: UUID) -> ProviderConfigTable | None:
result = await self.session.execute(
select(ProviderConfigTable).where(ProviderConfigTable.id == provider_id)
)
return result.scalar_one_or_none()
async def get_provider_or_raise(self, provider_id: UUID) -> ProviderConfigTable:
provider = await self.get_provider(provider_id)
if not provider:
raise NotFoundError(resource_type="Provider", resource_id=str(provider_id))
return provider
async def get_by_name(self, name: str) -> ProviderConfigTable | None:
result = await self.session.execute(
select(ProviderConfigTable).where(ProviderConfigTable.name == name)
)
return result.scalar_one_or_none()
# =========================================================================
# MUTATIONS
# =========================================================================
async def create_provider(self, data: ProviderCreate) -> ProviderConfigTable:
"""Create a provider. Raises ConflictError on duplicate name."""
existing = await self.get_by_name(data.name)
if existing:
raise ConflictError(
f"Provider with name '{data.name}' already exists",
resource_type="provider",
)
encrypted: str | None = None
if data.auth_token:
try:
encrypted = encrypt_token(data.auth_token)
except EncryptionError as e:
self.log.error("Failed to encrypt auth token", error=str(e))
raise
row = ProviderConfigTable(
name=data.name,
type=data.type,
base_url=data.base_url,
auth_token_encrypted=encrypted,
enabled=data.enabled,
)
self.session.add(row)
await self.session.flush()
self.log.info(
"Provider created",
provider_id=str(row.id),
name=data.name,
type=data.type.value,
has_auth_token=bool(encrypted),
)
return row
async def update_provider(
self, provider_id: UUID, data: ProviderUpdate
) -> ProviderConfigTable | None:
"""Apply a patch to a provider. Tri-state semantics on `auth_token`."""
row = await self.get_provider(provider_id)
if not row:
return None
if data.name is not None:
# Duplicate-name check only when actually changing the name.
if data.name != row.name:
dup = await self.get_by_name(data.name)
if dup and dup.id != row.id:
raise ConflictError(
f"Provider with name '{data.name}' already exists",
resource_type="provider",
)
row.name = data.name
if data.base_url is not None:
# Empty string → clear to NULL (matches git-token convention).
row.base_url = data.base_url or None
if data.enabled is not None:
row.enabled = data.enabled
if data.clear_auth_token:
row.auth_token_encrypted = None
self.log.info("Provider auth token cleared", provider_id=str(row.id))
elif data.auth_token:
try:
row.auth_token_encrypted = encrypt_token(data.auth_token)
except EncryptionError as e:
self.log.error("Failed to encrypt auth token", error=str(e))
raise
self.log.info("Provider auth token updated", provider_id=str(row.id))
await self.session.flush()
return row
async def delete_provider(self, provider_id: UUID) -> None:
"""Delete a provider. 409 if any assignment references it."""
row = await self.get_provider_or_raise(provider_id)
ref_count_q = select(ModelAssignmentTable).where(
ModelAssignmentTable.provider_config_id == provider_id
)
ref_result = await self.session.execute(ref_count_q)
if ref_result.first():
raise ConflictError(
"Provider is referenced by one or more model assignments; "
"remove those first.",
resource_type="provider",
)
await self.session.delete(row)
await self.session.flush()
self.log.info("Provider deleted", provider_id=str(provider_id))
# =========================================================================
# DECRYPTION
# =========================================================================
async def get_decrypted_token(self, provider_id: UUID) -> str | None:
"""Return the decrypted token for a provider, or None if unset."""
row = await self.get_provider(provider_id)
if not row or not row.auth_token_encrypted:
return None
try:
return decrypt_token(row.auth_token_encrypted)
except EncryptionError as e:
self.log.error(
"Failed to decrypt provider auth token",
provider_id=str(provider_id),
error=str(e),
)
raise
def get_provider_service(session: AsyncSession) -> ProviderService:
"""Get a ProviderService instance."""
return ProviderService(session)
+411
View File
@@ -0,0 +1,411 @@
"""
Test Runner Service
Orchestrates CI/CD-style commands (tests, lint, format, typecheck, build)
in agent workspaces. The API routes are thin adapters over this service —
all workspace resolution, subprocess dispatch, and output parsing happens
here so routes only handle HTTP translation.
"""
from __future__ import annotations
import asyncio
import re
import shlex
import subprocess
import time
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any, ClassVar
from roboco.api.schemas.test import (
BuildRequest,
BuildResponse,
FormatRequest,
FormatResponse,
LintIssue,
LintRequest,
LintResponse,
TestRunRequest,
TestRunResponse,
TestStatusResponse,
TypecheckError,
TypecheckRequest,
TypecheckResponse,
)
from roboco.config import settings
from roboco.services.base import (
BaseService,
NotFoundError,
ServiceError,
ServiceUnavailableError,
ValidationError,
)
from roboco.services.project import get_project_service
from roboco.services.workspace import WorkspaceError, get_workspace_service
if TYPE_CHECKING:
from uuid import UUID
from sqlalchemy.ext.asyncio import AsyncSession
# Command timeout in seconds — long enough for full test runs.
_CMD_TIMEOUT = 300
# Minimum parts for lint output parsing (file:line:col:message).
_LINT_PARTS_MIN = 4
# Minimum parts for type error parsing (file:line:message).
_TYPE_ERROR_PARTS_MIN = 3
@dataclass(frozen=True)
class _ProjectContext:
"""Resolved (project, workspace) pair for a command run."""
project: Any
workspace: Path
class TestRunnerService(BaseService):
"""Runs project CI commands in an agent's workspace."""
service_name: ClassVar[str] = "test_runner"
# =========================================================================
# WORKSPACE RESOLUTION
# =========================================================================
async def _resolve_legacy_workspace(self, project: Any, project_slug: str) -> Path:
"""Legacy single-path project config (no agent_id given)."""
workspace_path = getattr(project, "workspace_path", None)
if not workspace_path:
raise ValidationError(
f"Project '{project_slug}' has no workspace configured and "
"no agent_id provided for dynamic workspace resolution"
)
workspace = Path(workspace_path)
if not workspace.exists():
raise ValidationError(f"Workspace path does not exist: {workspace}")
return workspace
async def _resolve_agent_workspace(
self, project: Any, project_slug: str, agent_id: UUID
) -> Path:
"""Resolve workspace via WorkspaceService (ensure/resolve per setting)."""
workspace_service = get_workspace_service(self.session)
try:
if settings.workspace_auto_clone:
return await workspace_service.ensure_workspace(
project_slug=project_slug,
agent_id=agent_id,
git_url=project.git_url,
default_branch=project.default_branch or "main",
)
workspace = await workspace_service.resolve_workspace(
project_slug=project_slug,
agent_id=agent_id,
)
if not workspace.exists():
raise ValidationError(
f"Workspace does not exist: {workspace}. "
"Clone the repository first or enable auto_clone."
)
return workspace
except WorkspaceError as e:
raise ValidationError(str(e)) from e
async def _load_project_and_workspace(
self,
project_slug: str,
agent_id: UUID | None,
) -> _ProjectContext:
"""Fetch project + resolve its workspace. Raises typed errors."""
service = get_project_service(self.session)
project = await service.get_by_slug(project_slug)
if not project:
raise NotFoundError(resource_type="Project", resource_id=project_slug)
if agent_id is None:
workspace = await self._resolve_legacy_workspace(project, project_slug)
else:
workspace = await self._resolve_agent_workspace(
project, project_slug, agent_id
)
return _ProjectContext(project=project, workspace=workspace)
# =========================================================================
# COMMAND RUNNER
# =========================================================================
async def _run_command(
self,
workspace: Path,
command: str,
timeout: int = _CMD_TIMEOUT,
) -> subprocess.CompletedProcess[str]:
"""Run a shell command in the workspace (non-blocking).
Commands are tokenized with shlex — no shell features (pipes,
redirects, env expansion) supported. Project-configured commands
are expected to be simple exe + args.
"""
argv = shlex.split(command)
def _run() -> subprocess.CompletedProcess[str]:
return subprocess.run(
argv,
check=False,
cwd=workspace,
capture_output=True,
text=True,
timeout=timeout,
)
try:
return await asyncio.to_thread(_run)
except subprocess.TimeoutExpired as e:
raise ServiceUnavailableError(
service_name="test_runner",
reason=f"Command timed out after {timeout}s: {command}",
) from e
def _project_cmd(self, project: Any, attr: str) -> str | None:
"""Fetch a configured command for the project, or None if unset.
Prior behavior raised ValidationError on a missing command, but
QA agents ended up hitting 400s on every probe when the project
simply hadn't opted into that CI tool (e.g., lint_command=null
because the task is a README edit). Now callers get `None` and
return a `skipped=True` success response — no gate violation,
clear signal to the agent that there's nothing to run.
"""
cmd = getattr(project, attr, None)
return str(cmd) if cmd else None
@staticmethod
def _skip_reason(project_slug: str, label: str) -> str:
return (
f"Project '{project_slug}' has no {label} configured — check "
"skipped. Review the change manually if it's relevant."
)
# =========================================================================
# OUTPUT PARSERS
# =========================================================================
@staticmethod
def _parse_pytest_counts(output: str) -> tuple[int, int, int]:
"""Parse (passed, failed, skipped) from pytest-style output."""
if "passed" not in output:
return 0, 0, 0
def _first_int(pattern: str) -> int:
match = re.search(pattern, output)
return int(match.group(1)) if match else 0
return (
_first_int(r"(\d+) passed"),
_first_int(r"(\d+) failed"),
_first_int(r"(\d+) skipped"),
)
@staticmethod
def _parse_lint_line(line: str) -> LintIssue | None:
"""Parse a single ruff-style lint output line."""
if "::" in line or not line.strip():
return None
parts = line.split(":", 3)
if len(parts) < _LINT_PARTS_MIN:
return None
try:
return LintIssue(
file=parts[0],
line=int(parts[1]),
column=int(parts[2]),
code=parts[3].split()[0] if parts[3].strip() else "E",
message=parts[3].strip(),
)
except (ValueError, IndexError):
return None
def _parse_lint_output(self, output: str) -> list[LintIssue]:
return [
issue
for line in output.split("\n")
if (issue := self._parse_lint_line(line)) is not None
]
@staticmethod
def _parse_typecheck_line(line: str) -> TypecheckError | None:
"""Parse a single mypy-style type error line."""
if ": error:" not in line:
return None
parts = line.split(":", 2)
if len(parts) < _TYPE_ERROR_PARTS_MIN:
return None
try:
return TypecheckError(
file=parts[0],
line=int(parts[1]),
message=parts[2].replace(" error:", "").strip(),
)
except (ValueError, IndexError):
return None
def _parse_typecheck_output(self, output: str) -> list[TypecheckError]:
return [
err
for line in output.split("\n")
if (err := self._parse_typecheck_line(line)) is not None
]
@staticmethod
def _build_format_cmd(base_cmd: str, data: FormatRequest) -> str:
cmd = base_cmd
if data.check_only:
cmd = f"{cmd} --check"
if data.path:
cmd = f"{cmd} {data.path}"
return cmd
# =========================================================================
# PUBLIC ORCHESTRATION
# =========================================================================
async def get_status(self, project_slug: str, agent_id: UUID) -> TestStatusResponse:
"""Placeholder status endpoint — validates workspace resolution."""
await self._load_project_and_workspace(project_slug, agent_id)
return TestStatusResponse(
project_slug=project_slug,
passed=True,
summary=(
"No test results stored yet. Run roboco_test_run() to execute tests."
),
last_run=None,
)
async def run_tests(self, agent_id: UUID, data: TestRunRequest) -> TestRunResponse:
ctx = await self._load_project_and_workspace(data.project_slug, agent_id)
base = self._project_cmd(ctx.project, "test_command")
if not base:
return TestRunResponse(
project_slug=data.project_slug,
passed=True,
skipped=True,
skip_reason=self._skip_reason(data.project_slug, "test_command"),
)
cmd = base
if data.test_path:
cmd = f"{cmd} {data.test_path}"
if data.verbose:
cmd = f"{cmd} -v"
result = await self._run_command(ctx.workspace, cmd)
output = result.stdout + result.stderr
passed, failed, skipped = self._parse_pytest_counts(output)
return TestRunResponse(
project_slug=data.project_slug,
passed=result.returncode == 0,
passed_count=passed,
failed_count=failed,
skipped_count=skipped,
output=output[:10000],
failures=[],
)
async def run_lint(self, agent_id: UUID, data: LintRequest) -> LintResponse:
ctx = await self._load_project_and_workspace(data.project_slug, agent_id)
base = self._project_cmd(ctx.project, "lint_command")
if not base:
return LintResponse(
project_slug=data.project_slug,
passed=True,
skipped=True,
skip_reason=self._skip_reason(data.project_slug, "lint_command"),
)
cmd = base
if data.fix:
cmd = f"{cmd} --fix"
if data.path:
cmd = f"{cmd} {data.path}"
result = await self._run_command(ctx.workspace, cmd)
output = result.stdout + result.stderr
return LintResponse(
project_slug=data.project_slug,
passed=result.returncode == 0,
issues=self._parse_lint_output(output),
fixed_count=0,
)
async def run_format(self, agent_id: UUID, data: FormatRequest) -> FormatResponse:
ctx = await self._load_project_and_workspace(data.project_slug, agent_id)
base = self._project_cmd(ctx.project, "format_command")
if not base:
return FormatResponse(
project_slug=data.project_slug,
skipped=True,
skip_reason=self._skip_reason(data.project_slug, "format_command"),
)
result = await self._run_command(
ctx.workspace, self._build_format_cmd(base, data)
)
output = result.stdout + result.stderr
files_modified = output.count("reformatted") if not data.check_only else 0
files_unchanged = output.count("unchanged") or output.count("already formatted")
return FormatResponse(
project_slug=data.project_slug,
files_modified=files_modified,
files_unchanged=files_unchanged,
)
async def run_typecheck(
self, agent_id: UUID, data: TypecheckRequest
) -> TypecheckResponse:
ctx = await self._load_project_and_workspace(data.project_slug, agent_id)
base = self._project_cmd(ctx.project, "typecheck_command")
if not base:
return TypecheckResponse(
project_slug=data.project_slug,
passed=True,
skipped=True,
skip_reason=self._skip_reason(data.project_slug, "typecheck_command"),
)
cmd = f"{base} {data.path}" if data.path else base
result = await self._run_command(ctx.workspace, cmd)
output = result.stdout + result.stderr
return TypecheckResponse(
project_slug=data.project_slug,
passed=result.returncode == 0,
errors=self._parse_typecheck_output(output),
)
async def run_build(self, agent_id: UUID, data: BuildRequest) -> BuildResponse:
ctx = await self._load_project_and_workspace(data.project_slug, agent_id)
base = self._project_cmd(ctx.project, "build_command")
if not base:
return BuildResponse(
project_slug=data.project_slug,
success=True,
skipped=True,
skip_reason=self._skip_reason(data.project_slug, "build_command"),
)
start = time.time()
result = await self._run_command(ctx.workspace, base)
duration = time.time() - start
return BuildResponse(
project_slug=data.project_slug,
success=result.returncode == 0,
duration_seconds=round(duration, 2),
output=(result.stdout + result.stderr)[:10000],
)
def get_test_runner_service(session: AsyncSession) -> TestRunnerService:
"""Factory for TestRunnerService."""
return TestRunnerService(session)
__all__ = ["ServiceError", "TestRunnerService", "get_test_runner_service"]