feat(grok): first-class xAI/Grok routing mode (UI + backend)

The Routing-mode toggle had Anthropic / Ollama / Self-Hosted / Mix but no way
to route the whole org to Grok. Add it end to end:

- backend: apply_mode("grok") + _apply_grok (GLOBAL default -> grok-build-0.1) +
  derive_mode "grok" detection; ApplyModeRequest/ModeResponse accept "grok".
- panel: a "Grok" routing-mode card (between Anthropic and Ollama, gated on the
  xAI key) + flipToGrok; a Grok group in the per-agent mix dropdown +
  catalogGrokOnly + a grok ProviderBadge variant; the mix-save key check and
  the AI-routing description now cover Grok.
- tests: integration derive_mode/apply_mode "grok" cases (+ grok provider row
  in the fixture).

Gated: ruff + mypy clean; panel typecheck + lint clean.
This commit is contained in:
Renn F
2026-06-18 10:13:14 +02:00
parent b0857915a1
commit 0e9bbc15db
5 changed files with 130 additions and 12 deletions
@@ -39,6 +39,7 @@ import {
Server, Server,
ShieldCheck, ShieldCheck,
Sparkles, Sparkles,
Zap,
} from "lucide-react"; } from "lucide-react";
import { toast } from "sonner"; import { toast } from "sonner";
import { AssignmentScope, ModelProvider } from "@/types"; import { AssignmentScope, ModelProvider } from "@/types";
@@ -175,6 +176,9 @@ export function AIRoutingCard() {
const catalogOllamaOnly = catalog.filter( const catalogOllamaOnly = catalog.filter(
(c: { provider_type: ModelProvider }) => c.provider_type === ModelProvider.OLLAMA_CLOUD, (c: { provider_type: ModelProvider }) => c.provider_type === ModelProvider.OLLAMA_CLOUD,
); );
const catalogGrokOnly = catalog.filter(
(c: { provider_type: ModelProvider }) => c.provider_type === ModelProvider.GROK,
);
const catalogAnthropicOnly = catalog.filter( const catalogAnthropicOnly = catalog.filter(
(c: { provider_type: ModelProvider }) => c.provider_type === ModelProvider.ANTHROPIC, (c: { provider_type: ModelProvider }) => c.provider_type === ModelProvider.ANTHROPIC,
); );
@@ -190,6 +194,20 @@ export function AIRoutingCard() {
} }
}; };
const flipToGrok = async () => {
if (!hasGrokKey) {
toast.error("Save the Grok (xAI) API key first");
return;
}
if (!confirm("Switch every agent to Grok? Clears any overrides.")) return;
try {
await applyMode.mutateAsync({ mode: "grok" });
toast.success("All agents now on Grok");
} catch (e) {
toast.error("Switch failed: " + errMsg(e));
}
};
const flipToOllama = async () => { const flipToOllama = async () => {
if (!hasOllamaKey) { if (!hasOllamaKey) {
toast.error("Save an Ollama API key first"); toast.error("Save an Ollama API key first");
@@ -236,6 +254,18 @@ export function AIRoutingCard() {
toast.error("Pick a model for at least one agent"); toast.error("Pick a model for at least one agent");
return; return;
} }
const needsGrok = Object.values(per_agent).some((m) =>
catalog.find(
(c: { model_name: string; provider_type: ModelProvider }) =>
c.model_name === m && c.provider_type === ModelProvider.GROK,
),
);
if (needsGrok && !hasGrokKey) {
toast.error(
"At least one agent is routed to a Grok model but no key is saved",
);
return;
}
const needsKey = Object.values(per_agent).some((m) => const needsKey = Object.values(per_agent).some((m) =>
catalog.find( catalog.find(
(c: { model_name: string; provider_type: ModelProvider }) => (c: { model_name: string; provider_type: ModelProvider }) =>
@@ -274,9 +304,9 @@ export function AIRoutingCard() {
</CardTitle> </CardTitle>
<CardDescription> <CardDescription>
Decide which model backs each agent. Anthropic uses the mounted Decide which model backs each agent. Anthropic uses the mounted
<code className="px-1"> ~/.claude </code> auth; Ollama Cloud uses <code className="px-1"> ~/.claude </code> auth; Grok (xAI) and Ollama
the API key you save below; Self-Hosted connects to any OpenAI-compatible Cloud use the API keys you save below; Self-Hosted connects to any
endpoint you run locally. OpenAI-compatible endpoint you run locally.
</CardDescription> </CardDescription>
</CardHeader> </CardHeader>
<CardContent className="space-y-6"> <CardContent className="space-y-6">
@@ -391,7 +421,7 @@ export function AIRoutingCard() {
{/* -------- Mode toggle -------- */} {/* -------- Mode toggle -------- */}
<section className="space-y-3"> <section className="space-y-3">
<Label className="text-sm font-medium">Routing mode</Label> <Label className="text-sm font-medium">Routing mode</Label>
<div className="grid grid-cols-2 md:grid-cols-4 gap-2"> <div className="grid grid-cols-2 md:grid-cols-3 lg:grid-cols-5 gap-2">
<ModeButton <ModeButton
icon={<ShieldCheck className="h-4 w-4" />} icon={<ShieldCheck className="h-4 w-4" />}
label="Anthropic" label="Anthropic"
@@ -400,6 +430,18 @@ export function AIRoutingCard() {
onClick={flipToAnthropic} onClick={flipToAnthropic}
disabled={applyMode.isPending} disabled={applyMode.isPending}
/> />
<ModeButton
icon={<Zap className="h-4 w-4" />}
label="Grok"
description={
hasGrokKey
? "Every agent uses Grok (grok-build-0.1)."
: "Save the Grok (xAI) key first."
}
active={currentMode === "grok"}
onClick={flipToGrok}
disabled={applyMode.isPending || !hasGrokKey}
/>
<ModeButton <ModeButton
icon={<Sparkles className="h-4 w-4" />} icon={<Sparkles className="h-4 w-4" />}
label="Ollama" label="Ollama"
@@ -545,6 +587,23 @@ export function AIRoutingCard() {
</SelectGroup> </SelectGroup>
)} )}
{/* Grok (xAI) models */}
{catalogGrokOnly.length > 0 && (
<SelectGroup>
<SelectLabel>
<ProviderBadge variant="grok" />
Grok (xAI)
</SelectLabel>
{catalogGrokOnly.map(
(c: { model_name: string; display_name: string }) => (
<SelectItem key={c.model_name} value={c.model_name}>
{c.display_name}
</SelectItem>
),
)}
</SelectGroup>
)}
{/* Ollama Cloud models */} {/* Ollama Cloud models */}
{catalogOllamaOnly.length > 0 && ( {catalogOllamaOnly.length > 0 && (
<SelectGroup> <SelectGroup>
@@ -659,17 +718,19 @@ function errMsg(e: unknown): string {
function ProviderBadge({ function ProviderBadge({
variant, variant,
}: { }: {
variant: "anthropic" | "ollama" | "self-hosted"; variant: "anthropic" | "grok" | "ollama" | "self-hosted";
}) { }) {
const styles: Record<string, string> = { const styles: Record<string, string> = {
anthropic: "bg-blue-500/20 text-blue-700 dark:text-blue-400", anthropic: "bg-blue-500/20 text-blue-700 dark:text-blue-400",
ollama: "bg-violet-500/20 text-violet-700 dark:text-violet-400", ollama: "bg-violet-500/20 text-violet-700 dark:text-violet-400",
"self-hosted": "bg-purple-500/20 text-purple-700 dark:text-purple-400", "self-hosted": "bg-purple-500/20 text-purple-700 dark:text-purple-400",
grok: "bg-teal-500/20 text-teal-700 dark:text-teal-400",
}; };
const labels: Record<string, string> = { const labels: Record<string, string> = {
anthropic: "A", anthropic: "A",
ollama: "O", ollama: "O",
"self-hosted": "S", "self-hosted": "S",
grok: "G",
}; };
return ( return (
<span <span
+6 -1
View File
@@ -29,7 +29,12 @@ export interface ModelAssignment {
model_name: string; model_name: string;
} }
export type RoutingMode = "anthropic" | "ollama" | "self_hosted" | "mix"; export type RoutingMode =
| "anthropic"
| "grok"
| "ollama"
| "self_hosted"
| "mix";
export interface ModeSnapshot { export interface ModeSnapshot {
mode: RoutingMode; mode: RoutingMode;
+3 -3
View File
@@ -6,7 +6,7 @@ Minimal surface that backs the Settings UI:
- set / clear / check the single Ollama Cloud API key - set / clear / check the single Ollama Cloud API key
- configure / test / discover the self-hosted (LOCAL) Ollama server - configure / test / discover the self-hosted (LOCAL) Ollama server
- read current routing assignments (so the UI renders Mix mode) - read current routing assignments (so the UI renders Mix mode)
- apply a routing mode (anthropic | ollama | mix | self_hosted) - apply a routing mode (anthropic | grok | ollama | mix | self_hosted)
""" """
from __future__ import annotations from __future__ import annotations
@@ -189,7 +189,7 @@ class ApplyModeRequest(BaseModel):
set GLOBAL default to `default_model` (a self-hosted model name). set GLOBAL default to `default_model` (a self-hosted model name).
""" """
mode: Literal["anthropic", "ollama", "mix", "self_hosted"] mode: Literal["anthropic", "grok", "ollama", "mix", "self_hosted"]
default_model: str | None = None default_model: str | None = None
per_agent: dict[str, str] | None = None per_agent: dict[str, str] | None = None
@@ -197,5 +197,5 @@ class ApplyModeRequest(BaseModel):
class ModeResponse(BaseModel): class ModeResponse(BaseModel):
"""Server-side view of the current mode + a snapshot of active rules.""" """Server-side view of the current mode + a snapshot of active rules."""
mode: Literal["anthropic", "ollama", "mix", "self_hosted"] mode: Literal["anthropic", "grok", "ollama", "mix", "self_hosted"]
assignments: list[AssignmentResponse] assignments: list[AssignmentResponse]
+28 -2
View File
@@ -311,7 +311,9 @@ class ModelRoutingService(BaseService):
) )
return row return row
async def derive_mode(self) -> Literal["anthropic", "ollama", "mix", "self_hosted"]: async def derive_mode(
self,
) -> Literal["anthropic", "grok", "ollama", "mix", "self_hosted"]:
"""Return the current "mode" label for the Settings UI. """Return the current "mode" label for the Settings UI.
Decision tree matches what `apply_mode` writes: Decision tree matches what `apply_mode` writes:
@@ -327,6 +329,8 @@ class ModelRoutingService(BaseService):
len(assignments) == 1 and assignments[0].scope == AssignmentScope.GLOBAL len(assignments) == 1 and assignments[0].scope == AssignmentScope.GLOBAL
) )
if only_global: if only_global:
if assignments[0].provider.type == ModelProvider.GROK:
return "grok"
if assignments[0].provider.type == ModelProvider.OLLAMA_CLOUD: if assignments[0].provider.type == ModelProvider.OLLAMA_CLOUD:
return "ollama" return "ollama"
if assignments[0].provider.type == ModelProvider.LOCAL: if assignments[0].provider.type == ModelProvider.LOCAL:
@@ -429,6 +433,8 @@ class ModelRoutingService(BaseService):
- "self_hosted": wipe all assignments, enable the LOCAL provider, - "self_hosted": wipe all assignments, enable the LOCAL provider,
and set the GLOBAL default to `default_model` (a self-hosted and set the GLOBAL default to `default_model` (a self-hosted
model name not validated against the static catalog). model name not validated against the static catalog).
- "grok": wipe all assignments, set the GLOBAL default to a
Grok (xAI) model (default grok-build-0.1). Requires the xAI key.
- "mix": apply per-agent map verbatim. Any agent not in the - "mix": apply per-agent map verbatim. Any agent not in the
map falls through to the GLOBAL default which is whatever it map falls through to the GLOBAL default which is whatever it
was (preserves prior state). Self-hosted model names (not in the was (preserves prior state). Self-hosted model names (not in the
@@ -436,6 +442,8 @@ class ModelRoutingService(BaseService):
""" """
if mode == "anthropic": if mode == "anthropic":
await self._apply_anthropic() await self._apply_anthropic()
elif mode == "grok":
await self._apply_grok(default_model)
elif mode == "ollama": elif mode == "ollama":
await self._apply_ollama(default_model) await self._apply_ollama(default_model)
elif mode == "self_hosted": elif mode == "self_hosted":
@@ -445,7 +453,7 @@ class ModelRoutingService(BaseService):
else: else:
raise ValueError( raise ValueError(
f"Unknown mode '{mode}'." f"Unknown mode '{mode}'."
" Use 'anthropic', 'ollama', 'self_hosted', or 'mix'." " Use 'anthropic', 'grok', 'ollama', 'self_hosted', or 'mix'."
) )
async def _apply_anthropic(self) -> None: async def _apply_anthropic(self) -> None:
@@ -454,6 +462,24 @@ class ModelRoutingService(BaseService):
await self.session.flush() await self.session.flush()
self.log.info("Mode applied: anthropic (all assignments cleared)") self.log.info("Mode applied: anthropic (all assignments cleared)")
async def _apply_grok(self, default_model: str | None) -> None:
"""Wipe assignments, set the GLOBAL default to a Grok (xAI) model.
``grok-build-0.1`` is in the catalog under the GROK provider, so the
upsert resolves to the seeded Grok provider row. Routing to Grok needs
the xAI key set (which enables the provider); without it, agents fall
back to the Anthropic path at spawn same contract as Ollama.
"""
await self.session.execute(sa_delete(ModelAssignmentTable))
await self.session.flush()
model_name = default_model or "grok-build-0.1"
await self.upsert_assignment(
scope=AssignmentScope.GLOBAL,
scope_value=None,
model_name=model_name,
)
self.log.info("Mode applied: grok", default_model=model_name)
async def _apply_ollama(self, default_model: str | None) -> None: async def _apply_ollama(self, default_model: str | None) -> None:
"""Wipe assignments, set the GLOBAL default to an Ollama Cloud model.""" """Wipe assignments, set the GLOBAL default to an Ollama Cloud model."""
await self.session.execute(sa_delete(ModelAssignmentTable)) await self.session.execute(sa_delete(ModelAssignmentTable))
+27 -1
View File
@@ -37,13 +37,19 @@ async def llm_setup(
type=ModelProvider.ANTHROPIC, type=ModelProvider.ANTHROPIC,
enabled=True, enabled=True,
) )
grok = ProviderConfigTable(
name="grok-test",
type=ModelProvider.GROK,
enabled=True,
base_url="https://api.x.ai/v1",
)
ollama = ProviderConfigTable( ollama = ProviderConfigTable(
name="ollama-test", name="ollama-test",
type=ModelProvider.OLLAMA_CLOUD, type=ModelProvider.OLLAMA_CLOUD,
enabled=True, enabled=True,
base_url="https://ollama.example.com", base_url="https://ollama.example.com",
) )
db_session.add_all([anthropic, ollama]) db_session.add_all([anthropic, grok, ollama])
await db_session.flush() await db_session.flush()
yield {"svc": ModelRoutingService(db_session)} yield {"svc": ModelRoutingService(db_session)}
@@ -177,6 +183,16 @@ async def test_derive_mode_ollama_when_only_ollama_global(llm_setup: dict) -> No
assert await svc.derive_mode() == "ollama" assert await svc.derive_mode() == "ollama"
@pytest.mark.asyncio
async def test_derive_mode_grok_when_only_grok_global(llm_setup: dict) -> None:
svc = llm_setup["svc"]
grok_model = _first_model_for_type(ModelProvider.GROK)
await svc.upsert_assignment(
scope=AssignmentScope.GLOBAL, scope_value=None, model_name=grok_model
)
assert await svc.derive_mode() == "grok"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_derive_mode_mix_with_per_agent(llm_setup: dict) -> None: async def test_derive_mode_mix_with_per_agent(llm_setup: dict) -> None:
svc = llm_setup["svc"] svc = llm_setup["svc"]
@@ -215,6 +231,16 @@ async def test_apply_mode_ollama_sets_global(llm_setup: dict) -> None:
assert assignments[0].scope == AssignmentScope.GLOBAL assert assignments[0].scope == AssignmentScope.GLOBAL
@pytest.mark.asyncio
async def test_apply_mode_grok_sets_global(llm_setup: dict) -> None:
svc = llm_setup["svc"]
await svc.apply_mode(mode="grok")
assignments = await svc.list_assignments()
assert len(assignments) == 1
assert assignments[0].scope == AssignmentScope.GLOBAL
assert assignments[0].provider.type == ModelProvider.GROK
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_apply_mode_mix_requires_per_agent(llm_setup: dict) -> None: async def test_apply_mode_mix_requires_per_agent(llm_setup: dict) -> None:
svc = llm_setup["svc"] svc = llm_setup["svc"]