From 6f28027687a8ffd71876f16999c879af1e5b047c Mon Sep 17 00:00:00 2001 From: Mohamed Boudra Date: Mon, 13 Apr 2026 19:42:01 +0700 Subject: [PATCH] =?UTF-8?q?feat:=20provider=20profiles=20=E2=80=94=20custo?= =?UTF-8?q?m=20provider=20definitions=20(#290)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: add provider profiles for custom provider definitions Users can define custom providers in config.json that appear as first-class entries alongside built-ins. A provider can override a built-in (custom binary, env, models) or create a new one by extending a base via `extends`. Generic ACP transport supported via `extends: "acp"`. Providers can be hidden with `enabled: false`. Hardcoded models merge with runtime-fetched ones. - Config schema with Zod validation, auto-migration from old format - Dynamic provider registry replaces static provider lists - GenericACPAgentClient for user-defined ACP providers - Snapshot entries carry label/description/defaultModeId over the wire - MCP tools accept dynamic provider IDs - App derives provider definitions from snapshot with static fallback - CLI `provider ls` calls daemon with label column - Schedule/session rehydration validates providers against registry * fix: accept any provider status in CLI provider ls test The test now connects to a real daemon where providers may be loading or unavailable, not just the static "available" fallback. * ci: re-trigger CI checks * style: fix checkout-git.ts formatting to match CI Biome version * fix: relax provider ls test assertions for daemon-backed responses The daemon snapshot may not include all 5 built-in providers in CI (some require external binaries). Assert at least the core 3 (claude, codex, opencode) instead of all 5. * fix app combobox dropdown positioning flash * refactor: stop merging models in provider registry, use override models directly Override models now replace instead of merge with base provider models. Also add icon/colorTier fallback from definition modes in fetchModes. * refactor: make provider definitions fully dynamic from server snapshots Remove static AGENT_PROVIDER_DEFINITIONS fallbacks from the client — providers, modes, icons, and color tiers now flow entirely from runtime snapshots. Add icon and colorTier to AgentMode schema so the server can advertise mode visuals directly. Fix setAgentMode to persist modeId in agent config so the selected mode survives session reload. Simplify model merging so profile models replace runtime models instead of prepending. * docs: add ad-hoc daemon testing guide --- docs/AD-HOC-DAEMON-TESTING.md | 159 +++++ .../agent-form/agent-form-dropdowns.tsx | 6 +- .../app/src/components/agent-status-bar.tsx | 29 +- .../components/provider-diagnostic-sheet.tsx | 7 +- packages/app/src/components/ui/combobox.tsx | 77 ++- .../src/hooks/use-agent-form-state.test.ts | 129 +++- .../app/src/hooks/use-agent-form-state.ts | 34 +- .../src/screens/agent/draft-agent-screen.tsx | 11 +- packages/app/src/screens/settings-screen.tsx | 5 +- .../app/src/utils/provider-definitions.ts | 47 ++ packages/cli/src/commands/provider/ls.ts | 48 +- packages/cli/tests/15-provider.test.ts | 21 +- .../src/server/agent/agent-management-mcp.ts | 4 +- .../src/server/agent/agent-manager.test.ts | 150 +++++ .../server/src/server/agent/agent-manager.ts | 10 +- .../src/server/agent/agent-sdk-types.ts | 5 + .../src/server/agent/mcp-server.test.ts | 77 +++ .../server/src/server/agent/mcp-server.ts | 4 +- .../server/src/server/agent/mcp-shared.ts | 12 +- .../agent/provider-launch-config.test.ts | 176 +++++- .../server/agent/provider-launch-config.ts | 76 +++ .../src/server/agent/provider-manifest.ts | 27 +- .../server/agent/provider-registry.test.ts | 584 ++++++++++++++++++ .../src/server/agent/provider-registry.ts | 415 ++++++++++++- .../agent/provider-snapshot-manager.test.ts | 148 ++++- .../server/agent/provider-snapshot-manager.ts | 19 +- .../agent/providers/generic-acp-agent.ts | 46 ++ packages/server/src/server/bootstrap.ts | 10 +- packages/server/src/server/config.ts | 60 +- packages/server/src/server/exports.ts | 7 +- .../src/server/persisted-config.test.ts | 292 ++++++++- .../server/src/server/persisted-config.ts | 124 +++- .../src/server/persistence-hooks.test.ts | 19 + .../server/src/server/persistence-hooks.ts | 24 +- .../server/src/server/schedule/service.ts | 9 +- packages/server/src/server/session.ts | 72 ++- .../websocket-server.notifications.test.ts | 1 + .../websocket-server.relay-reconnect.test.ts | 1 + .../server/src/server/websocket-server.ts | 10 +- packages/server/src/shared/messages.ts | 5 + 40 files changed, 2762 insertions(+), 198 deletions(-) create mode 100644 docs/AD-HOC-DAEMON-TESTING.md create mode 100644 packages/app/src/utils/provider-definitions.ts create mode 100644 packages/server/src/server/agent/provider-registry.test.ts create mode 100644 packages/server/src/server/agent/providers/generic-acp-agent.ts diff --git a/docs/AD-HOC-DAEMON-TESTING.md b/docs/AD-HOC-DAEMON-TESTING.md new file mode 100644 index 000000000..30c0d2e65 --- /dev/null +++ b/docs/AD-HOC-DAEMON-TESTING.md @@ -0,0 +1,159 @@ +# Ad-hoc daemon testing + +Spin up an isolated daemon programmatically without touching the main daemon on port 6767. + +## Quick start + +```typescript +import os from "node:os"; +import path from "node:path"; +import { mkdir, mkdtemp, rm } from "node:fs/promises"; +import pino from "pino"; +import { createPaseoDaemon } from "./bootstrap.js"; +import { DaemonClient } from "./test-utils/daemon-client.js"; + +const logger = pino({ level: "warn" }); +const paseoHomeRoot = await mkdtemp(path.join(os.tmpdir(), "paseo-test-")); +const paseoHome = path.join(paseoHomeRoot, ".paseo"); +await mkdir(paseoHome, { recursive: true }); +const staticDir = await mkdtemp(path.join(os.tmpdir(), "paseo-static-")); + +const daemon = await createPaseoDaemon( + { + listen: "127.0.0.1:0", // OS picks a free port + paseoHome, + corsAllowedOrigins: [], + allowedHosts: true, + mcpEnabled: false, + staticDir, + mcpDebug: false, + agentClients: {}, + agentStoragePath: path.join(paseoHome, "agents"), + relayEnabled: false, + relayEndpoint: "relay.paseo.sh:443", + appBaseUrl: "https://app.paseo.sh", + // Add custom config here, e.g.: + // providerOverrides: { ... }, + }, + logger, +); + +await daemon.start(); +const target = daemon.getListenTarget(); +const port = target!.type === "tcp" ? target!.port : null; + +const client = new DaemonClient({ + url: `ws://127.0.0.1:${port}/ws`, + appVersion: "0.1.54", // see gotcha #1 +}); +await client.connect(); +await client.fetchAgents({ subscribe: { subscriptionId: "test" } }); + +// ... do your testing ... + +await client.close(); +await daemon.stop(); +await rm(paseoHomeRoot, { recursive: true, force: true }); +await rm(staticDir, { recursive: true, force: true }); +``` + +Run with: +```bash +npx tsx packages/server/src/server/your-script.ts +``` + +## Using the test helper + +For simpler cases, `createTestPaseoDaemon` + `DaemonClient` handles temp dirs and port selection: + +```typescript +import { createTestPaseoDaemon } from "./test-utils/paseo-daemon.js"; +import { DaemonClient } from "./test-utils/daemon-client.js"; + +const daemon = await createTestPaseoDaemon(); +const client = new DaemonClient({ + url: `ws://127.0.0.1:${daemon.port}/ws`, + appVersion: "0.1.54", +}); +await client.connect(); +await client.fetchAgents({ subscribe: { subscriptionId: "test" } }); + +// ... test ... + +await client.close(); +await daemon.close(); // stops daemon + cleans up temp dirs +``` + +The test helper does **not** expose `providerOverrides`. Use `createPaseoDaemon` directly when you need it (see quick start above). + +## Common client methods + +```typescript +// Provider discovery +const snapshot = await client.getProvidersSnapshot({ cwd: "/tmp" }); +const models = await client.listProviderModels("claude"); +const modes = await client.listProviderModes("claude"); + +// Agent lifecycle +const agent = await client.createAgent({ provider: "claude", cwd: "/tmp" }); +await client.sendMessage(agent.id, "Hello"); +const updated = await client.waitForAgentUpsert(agent.id, (s) => s.status === "idle"); +``` + +## Gotchas + +### 1. appVersion gates provider visibility + +The daemon hides non-legacy providers (anything other than claude, codex, opencode) from clients that don't send an `appVersion >= 0.1.45`. The `DaemonClient` sends no version by default, so custom providers like ACP-based ones will be invisible in snapshot responses. + +Always pass `appVersion`: +```typescript +const client = new DaemonClient({ + url: `ws://127.0.0.1:${port}/ws`, + appVersion: "0.1.54", +}); +``` + +### 2. Provider snapshots are async + +After the daemon starts, providers are probed in the background. The first `getProvidersSnapshot()` call will likely return `status: "loading"` for most providers. Poll until the provider you care about is no longer loading: + +```typescript +let snapshot = await client.getProvidersSnapshot({ cwd: "/tmp" }); +for (let i = 0; i < 20; i++) { + const entry = snapshot.entries.find((e) => e.provider === "gemini"); + if (entry && entry.status !== "loading") break; + await new Promise((r) => setTimeout(r, 2_000)); + snapshot = await client.getProvidersSnapshot({ cwd: "/tmp" }); +} +``` + +### 3. fetchAgents is required before most operations + +Call `client.fetchAgents()` after connecting. The daemon session expects this handshake before it processes other requests — without it, messages like `get_providers_snapshot_request` will silently hang. + +### 4. listen: "127.0.0.1:0" for port allocation + +Always use port `0` so the OS picks a free port. Never hardcode a port — it will collide with the main daemon or other test runs. + +### 5. Script must live inside packages/server + +The test utilities use relative imports through the TypeScript project. Place your script somewhere under `packages/server/src/` and import from there. Scripts outside the repo will fail with module resolution errors. + +### 6. Cleanup on failure + +Wrap your test logic in try/finally to ensure the daemon stops and temp dirs are cleaned up, even if an assertion fails: + +```typescript +try { + // ... test logic ... +} finally { + await client.close(); + await daemon.stop().catch(() => undefined); + await rm(paseoHomeRoot, { recursive: true, force: true }); +} +``` + +### 7. ACP providers spawn real processes + +When testing ACP providers (e.g., Gemini with `extends: "acp"`), the daemon will spawn real processes to probe for models and modes. The binary must be installed and on PATH. Probing can take 5-15 seconds depending on the provider. diff --git a/packages/app/src/components/agent-form/agent-form-dropdowns.tsx b/packages/app/src/components/agent-form/agent-form-dropdowns.tsx index 5c5b596c8..d8a12afd3 100644 --- a/packages/app/src/components/agent-form/agent-form-dropdowns.tsx +++ b/packages/app/src/components/agent-form/agent-form-dropdowns.tsx @@ -557,7 +557,11 @@ export function AgentConfigRow({ const effectiveSelectedThinkingOption = selectedThinkingOptionId || thinkingSelectOptions[0]?.id || ""; - const selectedModeVisuals = getModeVisuals(selectedProvider, effectiveSelectedMode); + const selectedModeVisuals = getModeVisuals( + selectedProvider, + effectiveSelectedMode, + providerDefinitions, + ); const ModeIcon = MODE_ICON_MAP[selectedModeVisuals?.icon ?? "ShieldCheck"]; const modeIconColor = MODE_COLOR_MAP[selectedModeVisuals?.colorTier ?? "safe"]; diff --git a/packages/app/src/components/agent-status-bar.tsx b/packages/app/src/components/agent-status-bar.tsx index f18afc048..f79392b9d 100644 --- a/packages/app/src/components/agent-status-bar.tsx +++ b/packages/app/src/components/agent-status-bar.tsx @@ -17,6 +17,7 @@ import { getProviderIcon } from "@/components/provider-icons"; import { CombinedModelSelector } from "@/components/combined-model-selector"; import { useSessionStore } from "@/stores/session-store"; import { useProvidersSnapshot } from "@/hooks/use-providers-snapshot"; +import { resolveProviderDefinition } from "@/utils/provider-definitions"; import { buildFavoriteModelKey, mergeProviderPreferences, @@ -40,7 +41,6 @@ import type { } from "@server/server/agent/agent-sdk-types"; import type { AgentProviderDefinition } from "@server/server/agent/provider-manifest"; import { - AGENT_PROVIDER_DEFINITIONS, getModeVisuals, type AgentModeColorTier, type AgentModeIcon, @@ -60,10 +60,6 @@ type StatusOption = { type StatusSelector = "provider" | "mode" | "model" | "thinking" | `feature-${string}`; -const PROVIDER_DEFINITION_MAP = new Map( - AGENT_PROVIDER_DEFINITIONS.map((definition) => [definition.id, definition]), -); - type ControlledAgentStatusBarProps = { provider: string; providerOptions?: StatusOption[]; @@ -81,7 +77,7 @@ type ControlledAgentStatusBarProps = { onSelectThinkingOption?: (thinkingOptionId: string) => void; disabled?: boolean; isModelLoading?: boolean; - providerDefinitions?: AgentProviderDefinition[]; + providerDefinitions: AgentProviderDefinition[]; allProviderModels?: Map; canSelectModelProvider?: (providerId: string) => boolean; favoriteKeys?: Set; @@ -252,7 +248,9 @@ function ControlledStatusBar({ thinkingOptions?.[0]?.label ?? "Unknown", ); - const modeVisuals = selectedModeId ? getModeVisuals(provider, selectedModeId) : undefined; + const modeVisuals = selectedModeId + ? getModeVisuals(provider, selectedModeId, providerDefinitions) + : undefined; const ModeIconComponent = modeVisuals?.icon ? MODE_ICONS[modeVisuals.icon] : null; const modeIconColor = getModeIconColor(modeVisuals?.colorTier, theme.colors.palette); const ProviderIcon = getProviderIcon(provider); @@ -300,9 +298,7 @@ function ControlledStatusBar({ ); return map; }, [modelOptions, provider]); - const effectiveProviderDefinitions = - providerDefinitions ?? - (PROVIDER_DEFINITION_MAP.has(provider) ? [PROVIDER_DEFINITION_MAP.get(provider)!] : []); + const effectiveProviderDefinitions = providerDefinitions; const effectiveAllProviderModels = allProviderModels ?? fallbackAllProviderModels; const canSelectProviderInModelMenu = canSelectModelProvider ?? (() => true); const comboboxThinkingOptions = useMemo( @@ -322,7 +318,7 @@ function ControlledStatusBar({ active: boolean; onPress: () => void; }) => { - const visuals = getModeVisuals(provider, option.id); + const visuals = getModeVisuals(provider, option.id, providerDefinitions); const IconComponent = visuals?.icon ? MODE_ICONS[visuals.icon] : ShieldCheck; return ( ); }, - [provider, theme.colors.foreground], + [provider, providerDefinitions, theme.colors.foreground], ); const handleOpenChange = useCallback( @@ -745,7 +741,7 @@ function ControlledStatusBar({ {modeOptions.map((mode) => { - const visuals = getModeVisuals(provider, mode.id); + const visuals = getModeVisuals(provider, mode.id, providerDefinitions); const Icon = visuals?.icon ? MODE_ICONS[visuals.icon] : ShieldCheck; return ( { - const definition = AGENT_PROVIDER_DEFINITIONS.find((d) => d.id === agent?.provider); + const definition = agent?.provider + ? resolveProviderDefinition(agent.provider, snapshotEntries) + : undefined; return definition ? [definition] : []; - }, [agent?.provider]); + }, [agent?.provider, snapshotEntries]); const agentProviderModels = useMemo(() => { const map = new Map(); @@ -1128,6 +1126,7 @@ export function DraftAgentStatusBar({ /> (null); const [loading, setLoading] = useState(false); - const providerLabel = - AGENT_PROVIDER_DEFINITIONS.find((d) => d.id === provider)?.label ?? provider; + const providerLabel = resolveProviderLabel(provider, snapshotEntries); const fetchDiagnostic = useCallback(async () => { if (!client || !provider) return; diff --git a/packages/app/src/components/ui/combobox.tsx b/packages/app/src/components/ui/combobox.tsx index 929dbb5e9..a646bac20 100644 --- a/packages/app/src/components/ui/combobox.tsx +++ b/packages/app/src/components/ui/combobox.tsx @@ -268,7 +268,7 @@ export function Combobox({ const isMobile = useIsCompactFormFactor(); const effectiveOptionsPosition = isMobile ? "below-search" : optionsPosition; const isDesktopAboveSearch = !isMobile && isWeb && effectiveOptionsPosition === "above-search"; - const { height: windowHeight } = useWindowDimensions(); + const { height: windowHeight, width: windowWidth } = useWindowDimensions(); const bottomSheetRef = useRef(null); const hasPresentedBottomSheetRef = useRef(false); const snapPoints = useMemo(() => ["60%", "90%"], []); @@ -276,11 +276,13 @@ export function Combobox({ null, ); const [referenceWidth, setReferenceWidth] = useState(null); + const [referenceLeft, setReferenceLeft] = useState(null); const [referenceTop, setReferenceTop] = useState(null); const [referenceAtOrigin, setReferenceAtOrigin] = useState(false); const [searchQuery, setSearchQuery] = useState(""); const [activeIndex, setActiveIndex] = useState(-1); const desktopOptionsScrollRef = useRef(null); + const [desktopContentWidth, setDesktopContentWidth] = useState(null); const isControlled = typeof open === "boolean"; const [internalOpen, setInternalOpen] = useState(false); @@ -324,7 +326,7 @@ export function Combobox({ const middleware = useMemo( () => [ - floatingOffset(isWeb ? 0 : 4), + floatingOffset(isWeb ? 5 : 4), ...(isWeb ? [] : [flip({ padding: collisionPadding })]), ...(isDesktopAboveSearch ? [] : [shift({ padding: collisionPadding })]), floatingSize({ @@ -338,6 +340,9 @@ export function Combobox({ }); setReferenceWidth((prev) => { const next = rects.reference.width; + if (!(next > 0)) { + return prev; + } if (prev === next) return prev; return next; }); @@ -359,15 +364,18 @@ export function Combobox({ useEffect(() => { if (!isOpen || isMobile) { setAvailableSize(null); + setDesktopContentWidth(null); + setReferenceLeft(null); setReferenceWidth(null); return; } const raf = requestAnimationFrame(() => void update()); return () => cancelAnimationFrame(raf); - }, [desktopPlacement, isMobile, update, isOpen]); + }, [desktopPlacement, isMobile, isOpen, update]); useEffect(() => { if (!isOpen || isMobile) { + setReferenceLeft(null); setReferenceAtOrigin(false); setReferenceTop(null); return; @@ -381,9 +389,16 @@ export function Combobox({ } const measure = () => { - referenceEl.measureInWindow((x, y) => { + referenceEl.measureInWindow((x, y, width, height) => { + setReferenceLeft((prev) => (prev === x ? prev : x)); setReferenceAtOrigin(Math.abs(x) <= 1 && Math.abs(y) <= 1); setReferenceTop((prev) => (prev === y ? prev : y)); + setReferenceWidth((prev) => { + if (!(width > 0)) { + return prev; + } + return prev === width ? prev : width; + }); }); }; @@ -398,32 +413,46 @@ export function Combobox({ isDesktopAboveSearch && referenceTop !== null ? Math.max(windowHeight - referenceTop, collisionPadding) : null; - const hasResolvedDesktopPosition = - referenceWidth !== null && - floatingLeft !== null && - (isDesktopAboveSearch ? desktopAboveSearchBottom !== null : floatingTop !== null) && - ((floatingTop ?? 0) !== 0 || floatingLeft !== 0 || referenceAtOrigin); - const shouldHideDesktopContent = desktopPreventInitialFlash && !hasResolvedDesktopPosition; - const shouldUseDesktopFade = !desktopPreventInitialFlash; - // For top-placed popups: once position resolves, use bottom-based CSS positioning - // so height changes grow upward naturally without floating-ui needing to reposition. - const useStableBottom = + const hasNonZeroFloatingPosition = (floatingTop ?? 0) !== 0 || floatingLeft !== 0; + const useMeasuredTopStartPosition = !isDesktopAboveSearch && IS_WEB && !isMobile && - hasResolvedDesktopPosition && - desktopPlacement.startsWith("top") && - referenceTop !== null; + desktopPlacement === "top-start" && + referenceTop !== null && + referenceLeft !== null && + desktopContentWidth !== null; + const clampedMeasuredTopStartLeft = useMeasuredTopStartPosition + ? Math.max( + collisionPadding, + Math.min(windowWidth - desktopContentWidth - collisionPadding, referenceLeft), + ) + : null; + const measuredTopStartBottom = useMeasuredTopStartPosition + ? Math.max(windowHeight - referenceTop + 5, collisionPadding) + : null; + const hasResolvedDesktopPosition = + referenceWidth !== null && + referenceWidth > 0 && + (isDesktopAboveSearch + ? floatingLeft !== null && desktopAboveSearchBottom !== null + : useMeasuredTopStartPosition + ? clampedMeasuredTopStartLeft !== null && measuredTopStartBottom !== null + : floatingLeft !== null && + floatingTop !== null && + (hasNonZeroFloatingPosition || !referenceAtOrigin)); + const shouldHideDesktopContent = desktopPreventInitialFlash && !hasResolvedDesktopPosition; + const shouldUseDesktopFade = !desktopPreventInitialFlash; const desktopPositionStyle = isDesktopAboveSearch ? { left: floatingLeft ?? 0, bottom: desktopAboveSearchBottom ?? 0, } - : useStableBottom + : useMeasuredTopStartPosition ? { - left: floatingLeft ?? 0, - bottom: Math.max(windowHeight - referenceTop!, collisionPadding), + left: clampedMeasuredTopStartLeft ?? 0, + bottom: measuredTopStartBottom ?? 0, } : floatingStyles; @@ -728,7 +757,13 @@ export function Combobox({ ]} ref={refs.setFloating} collapsable={false} - onLayout={() => update()} + onLayout={(event) => { + const { width, height } = event.nativeEvent.layout; + setDesktopContentWidth((prev) => (prev === width ? prev : width)); + if (!useMeasuredTopStartPosition || !hasResolvedDesktopPosition) { + void update(); + } + }} > {children ? ( <> diff --git a/packages/app/src/hooks/use-agent-form-state.test.ts b/packages/app/src/hooks/use-agent-form-state.test.ts index d909d817c..8b9457797 100644 --- a/packages/app/src/hooks/use-agent-form-state.test.ts +++ b/packages/app/src/hooks/use-agent-form-state.test.ts @@ -1,12 +1,117 @@ import { describe, expect, it } from "vitest"; import { __private__ } from "./use-agent-form-state"; -import { - AGENT_PROVIDER_DEFINITIONS, - type AgentProviderDefinition, -} from "@server/server/agent/provider-manifest"; -import type { AgentModelDefinition, AgentProvider } from "@server/server/agent/agent-sdk-types"; +import { buildProviderDefinitions } from "@/utils/provider-definitions"; +import type { AgentProviderDefinition } from "@server/server/agent/provider-manifest"; +import type { + AgentModelDefinition, + AgentProvider, + ProviderSnapshotEntry, +} from "@server/server/agent/agent-sdk-types"; + +const TEST_CODEX_DEFINITION: AgentProviderDefinition = { + id: "codex", + label: "Codex", + description: "Codex test provider", + defaultModeId: "auto", + modes: [ + { id: "auto", label: "Auto", icon: "ShieldAlert", colorTier: "moderate" }, + { id: "full-access", label: "Full Access", icon: "ShieldAlert", colorTier: "dangerous" }, + ], +}; + +const TEST_CLAUDE_DEFINITION: AgentProviderDefinition = { + id: "claude", + label: "Claude", + description: "Claude test provider", + defaultModeId: "default", + modes: [ + { id: "default", label: "Always Ask", icon: "ShieldCheck", colorTier: "safe" }, + { id: "acceptEdits", label: "Accept File Edits", icon: "ShieldAlert", colorTier: "moderate" }, + { id: "plan", label: "Plan Mode", icon: "ShieldCheck", colorTier: "planning" }, + { id: "bypassPermissions", label: "Bypass", icon: "ShieldAlert", colorTier: "dangerous" }, + ], +}; + +function makeProviderMap( + ...definitions: AgentProviderDefinition[] +): Map { + return new Map(definitions.map((d) => [d.id as AgentProvider, d])); +} + +const codexProviderMap = makeProviderMap(TEST_CODEX_DEFINITION); +const claudeProviderMap = makeProviderMap(TEST_CLAUDE_DEFINITION); describe("useAgentFormState", () => { + describe("buildProviderDefinitions", () => { + it("returns empty array when snapshot data is unavailable", () => { + expect(buildProviderDefinitions(undefined)).toEqual([]); + expect(buildProviderDefinitions([])).toEqual([]); + }); + + it("builds custom provider definitions from snapshot metadata", () => { + const entries: ProviderSnapshotEntry[] = [ + { + provider: "zai", + status: "ready", + label: "ZAI", + description: "Claude with ZAI config", + defaultModeId: "default", + modes: [ + { + id: "default", + label: "Default", + description: "Safe mode", + icon: "ShieldCheck", + colorTier: "safe", + }, + ], + }, + { + provider: "claude", + status: "ready", + label: "Claude", + description: "Anthropic Claude", + defaultModeId: "default", + modes: [{ id: "default", label: "Always Ask", icon: "ShieldCheck", colorTier: "safe" }], + }, + ]; + + const definitions = buildProviderDefinitions(entries); + + expect(definitions).toEqual([ + { + id: "zai", + label: "ZAI", + description: "Claude with ZAI config", + defaultModeId: "default", + modes: [ + { + id: "default", + label: "Default", + description: "Safe mode", + icon: "ShieldCheck", + colorTier: "safe", + }, + ], + }, + { + id: "claude", + label: "Claude", + description: "Anthropic Claude", + defaultModeId: "default", + modes: [ + { + id: "default", + label: "Always Ask", + icon: "ShieldCheck", + colorTier: "safe", + }, + ], + }, + ]); + }); + }); + describe("__private__.combineInitialValues", () => { it("returns undefined when no initial values and no initial server id", () => { expect(__private__.combineInitialValues(undefined, null)).toBeUndefined(); @@ -78,6 +183,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + codexProviderMap, ); expect(resolved.model).toBe("gpt-5.3-codex"); @@ -106,6 +212,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + codexProviderMap, ); expect(resolved.model).toBe("gpt-5.3-codex"); @@ -134,6 +241,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + codexProviderMap, ); expect(resolved.thinkingOptionId).toBe("xhigh"); @@ -161,6 +269,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + codexProviderMap, ); expect(resolved.model).toBe("gpt-5.3-codex"); @@ -188,6 +297,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + codexProviderMap, ); expect(resolved.model).toBe("gpt-5.3-codex"); @@ -215,6 +325,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + codexProviderMap, ); expect(resolved.model).toBe("gpt-5.3-codex"); @@ -256,6 +367,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), + claudeProviderMap, ); expect(resolved.model).toBe("default"); @@ -263,11 +375,6 @@ describe("useAgentFormState", () => { }); it("resolves provider only from allowed provider map", () => { - const allowedProviderMap = new Map( - AGENT_PROVIDER_DEFINITIONS.filter((definition) => definition.id === "claude").map( - (definition) => [definition.id as AgentProvider, definition], - ), - ); const resolved = __private__.resolveFormState( undefined, { provider: "codex" }, @@ -289,7 +396,7 @@ describe("useAgentFormState", () => { workingDir: "", }, new Set(), - allowedProviderMap, + claudeProviderMap, ); expect(resolved.provider).toBe("claude"); diff --git a/packages/app/src/hooks/use-agent-form-state.ts b/packages/app/src/hooks/use-agent-form-state.ts index a208c9a8b..b1df8ae9d 100644 --- a/packages/app/src/hooks/use-agent-form-state.ts +++ b/packages/app/src/hooks/use-agent-form-state.ts @@ -1,8 +1,5 @@ import { useCallback, useEffect, useMemo, useRef, useState } from "react"; -import { - AGENT_PROVIDER_DEFINITIONS, - type AgentProviderDefinition, -} from "@server/server/agent/provider-manifest"; +import type { AgentProviderDefinition } from "@server/server/agent/provider-manifest"; import type { AgentMode, AgentModelDefinition, @@ -10,6 +7,7 @@ import type { ProviderSnapshotEntry, } from "@server/server/agent/agent-sdk-types"; import { useHosts } from "@/runtime/host-runtime"; +import { buildProviderDefinitions } from "@/utils/provider-definitions"; import { useProvidersSnapshot } from "./use-providers-snapshot"; import { useFormPreferences, @@ -99,13 +97,8 @@ type UseAgentFormStateResult = { persistFormPreferences: () => Promise; }; -const allProviderDefinitions = AGENT_PROVIDER_DEFINITIONS; -const allProviderDefinitionMap = new Map( - allProviderDefinitions.map((definition) => [definition.id, definition]), -); -const fallbackDefinition = allProviderDefinitions[0]; -const DEFAULT_PROVIDER: AgentProvider = fallbackDefinition?.id ?? "claude"; -const DEFAULT_MODE_FOR_DEFAULT_PROVIDER = fallbackDefinition?.defaultModeId ?? ""; +const DEFAULT_PROVIDER: AgentProvider = "claude"; +const DEFAULT_MODE_FOR_DEFAULT_PROVIDER = "default"; function normalizeSelectedModelId(modelId: string | null | undefined): string { const normalized = typeof modelId === "string" ? modelId.trim() : ""; @@ -180,7 +173,7 @@ function resolveFormState( userModified: UserModifiedFields, currentState: FormState, validServerIds: Set, - allowedProviderMap: Map = allProviderDefinitionMap, + allowedProviderMap: Map, ): FormState { // Start with current state - we only update non-user-modified fields const result = { ...currentState }; @@ -376,10 +369,10 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg } = useProvidersSnapshot(formState.serverId); const allProviderEntries = useMemo(() => snapshotEntries ?? [], [snapshotEntries]); - const snapshotProviderDefinitions = useMemo(() => { - const snapshotProviders = new Set((snapshotEntries ?? []).map((entry) => entry.provider)); - return allProviderDefinitions.filter((definition) => snapshotProviders.has(definition.id)); - }, [snapshotEntries]); + const snapshotProviderDefinitions = useMemo( + () => buildProviderDefinitions(snapshotEntries), + [snapshotEntries], + ); const snapshotProviderDefinitionMap = useMemo( () => new Map( @@ -388,17 +381,18 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg [snapshotProviderDefinitions], ); const snapshotSelectableProviderDefinitionMap = useMemo(() => { + if (!snapshotEntries?.length) { + return snapshotProviderDefinitionMap; + } const readyProviders = new Set( - (snapshotEntries ?? []) - .filter((entry) => entry.status === "ready") - .map((entry) => entry.provider), + snapshotEntries.filter((entry) => entry.status === "ready").map((entry) => entry.provider), ); return new Map( snapshotProviderDefinitions .filter((definition) => readyProviders.has(definition.id)) .map((definition) => [definition.id, definition]), ); - }, [snapshotEntries, snapshotProviderDefinitions]); + }, [snapshotEntries, snapshotProviderDefinitionMap, snapshotProviderDefinitions]); const snapshotAllProviderModels = useMemo(() => { const map = new Map(); for (const entry of snapshotEntries ?? []) { diff --git a/packages/app/src/screens/agent/draft-agent-screen.tsx b/packages/app/src/screens/agent/draft-agent-screen.tsx index 6e19d7db9..1cfcede0d 100644 --- a/packages/app/src/screens/agent/draft-agent-screen.tsx +++ b/packages/app/src/screens/agent/draft-agent-screen.tsx @@ -49,7 +49,6 @@ import type { AgentCapabilityFlags, AgentSessionConfig, } from "@server/server/agent/agent-sdk-types"; -import { AGENT_PROVIDER_DEFINITIONS } from "@server/server/agent/provider-manifest"; import { prepareWorkspaceTab } from "@/utils/workspace-navigation"; import { TitlebarDragRegion } from "@/components/desktop/titlebar-drag-region"; import { useKeyboardShiftStyle } from "@/hooks/use-keyboard-shift-style"; @@ -68,10 +67,6 @@ const DRAFT_CAPABILITIES: AgentCapabilityFlags = { supportsReasoningStream: false, supportsToolInvocations: false, }; -const PROVIDER_DEFINITION_MAP = new Map( - AGENT_PROVIDER_DEFINITIONS.map((definition) => [definition.id, definition]), -); - function getParamValue(value: string | string[] | undefined) { if (typeof value === "string") { const trimmed = value.trim(); @@ -92,16 +87,14 @@ function getValidProvider(value: string | undefined) { if (!value) { return undefined; } - return PROVIDER_DEFINITION_MAP.has(value as AgentProvider) ? (value as AgentProvider) : undefined; + return value as AgentProvider; } function getValidMode(provider: AgentProvider | undefined, value: string | undefined) { if (!provider || !value) { return undefined; } - const definition = PROVIDER_DEFINITION_MAP.get(provider); - const modes = definition?.modes ?? []; - return modes.some((mode) => mode.id === value) ? value : undefined; + return value; } type DraftAgentParams = { diff --git a/packages/app/src/screens/settings-screen.tsx b/packages/app/src/screens/settings-screen.tsx index a70f80cd7..638d7a9a1 100644 --- a/packages/app/src/screens/settings-screen.tsx +++ b/packages/app/src/screens/settings-screen.tsx @@ -68,10 +68,10 @@ import { useIsLocalDaemon } from "@/hooks/use-is-local-daemon"; import { useDaemonConfig } from "@/hooks/use-daemon-config"; import { useProvidersSnapshot } from "@/hooks/use-providers-snapshot"; import { useIsCompactFormFactor } from "@/constants/layout"; -import { AGENT_PROVIDER_DEFINITIONS } from "@server/server/agent/provider-manifest"; import { getProviderIcon } from "@/components/provider-icons"; import { ProviderDiagnosticSheet } from "@/components/provider-diagnostic-sheet"; import { StatusBadge } from "@/components/ui/status-badge"; +import { buildProviderDefinitions } from "@/utils/provider-definitions"; import { isWeb } from "@/constants/platform"; // --------------------------------------------------------------------------- @@ -523,6 +523,7 @@ function ProvidersSection({ routeServerId }: ProvidersSectionProps) { const isConnected = useHostRuntimeIsConnected(routeServerId); const { entries, isLoading, isFetching, refresh } = useProvidersSnapshot(routeServerId); const [diagnosticProvider, setDiagnosticProvider] = useState(null); + const providerDefinitions = buildProviderDefinitions(entries); const hasServer = routeServerId.length > 0; @@ -558,7 +559,7 @@ function ProvidersSection({ routeServerId }: ProvidersSectionProps) { ) : ( - {AGENT_PROVIDER_DEFINITIONS.map((def) => { + {providerDefinitions.map((def) => { const entry = entries?.find((e) => e.provider === def.id); const status = entry?.status ?? "unavailable"; const ProviderIcon = getProviderIcon(def.id); diff --git a/packages/app/src/utils/provider-definitions.ts b/packages/app/src/utils/provider-definitions.ts new file mode 100644 index 000000000..c7cdb5e41 --- /dev/null +++ b/packages/app/src/utils/provider-definitions.ts @@ -0,0 +1,47 @@ +import type { ProviderSnapshotEntry, AgentProvider } from "@server/server/agent/agent-sdk-types"; +import { + type AgentModeColorTier, + type AgentModeIcon, + type AgentProviderDefinition, + type AgentProviderModeDefinition, +} from "@server/server/agent/provider-manifest"; + +function buildProviderModes(entry: ProviderSnapshotEntry): AgentProviderModeDefinition[] { + const entryModes = entry.modes ?? []; + + return entryModes.map((mode) => ({ + ...mode, + icon: (mode.icon ?? "ShieldCheck") as AgentModeIcon, + colorTier: (mode.colorTier ?? "moderate") as AgentModeColorTier, + })); +} + +export function buildProviderDefinitions( + snapshotEntries: ProviderSnapshotEntry[] | undefined, +): AgentProviderDefinition[] { + if (!snapshotEntries?.length) { + return []; + } + + return snapshotEntries.map((entry) => ({ + id: entry.provider, + label: entry.label ?? entry.provider, + description: entry.description ?? "", + defaultModeId: entry.defaultModeId ?? null, + modes: buildProviderModes(entry), + })); +} + +export function resolveProviderLabel( + provider: string, + snapshotEntries: ProviderSnapshotEntry[] | undefined, +): string { + return snapshotEntries?.find((entry) => entry.provider === provider)?.label ?? provider; +} + +export function resolveProviderDefinition( + provider: string, + snapshotEntries: ProviderSnapshotEntry[] | undefined, +): AgentProviderDefinition | undefined { + return buildProviderDefinitions(snapshotEntries).find((definition) => definition.id === provider); +} diff --git a/packages/cli/src/commands/provider/ls.ts b/packages/cli/src/commands/provider/ls.ts index 7e8bcfbb5..a4d8b7279 100644 --- a/packages/cli/src/commands/provider/ls.ts +++ b/packages/cli/src/commands/provider/ls.ts @@ -1,10 +1,12 @@ import type { Command } from "commander"; import type { CommandOptions, ListResult, OutputSchema } from "../../output/index.js"; import { AGENT_PROVIDER_DEFINITIONS } from "@getpaseo/server"; +import { tryConnectToDaemon } from "../../utils/client.js"; /** Provider list item for display */ export interface ProviderListItem { provider: string; + label: string; status: string; defaultMode: string; modes: string; @@ -13,16 +15,22 @@ export interface ProviderListItem { /** Derive provider list from the manifest — single source of truth */ const PROVIDERS: ProviderListItem[] = AGENT_PROVIDER_DEFINITIONS.map((def) => ({ provider: def.id, + label: def.label, status: "available", defaultMode: def.defaultModeId ?? "default", modes: def.modes.map((m) => m.label).join(", "), })); +function getStaticProviders(): ProviderListItem[] { + return PROVIDERS; +} + /** Schema for provider ls output */ export const providerLsSchema: OutputSchema = { idField: "provider", columns: [ { header: "PROVIDER", field: "provider", width: 12 }, + { header: "LABEL", field: "label", width: 16 }, { header: "STATUS", field: "status", @@ -45,13 +53,39 @@ export interface ProviderLsOptions extends CommandOptions { } export async function runLsCommand( - _options: ProviderLsOptions, + options: ProviderLsOptions, _command: Command, ): Promise { - // Provider data is static - no daemon connection needed - return { - type: "list", - data: PROVIDERS, - schema: providerLsSchema, - }; + const client = await tryConnectToDaemon({ host: options.host }); + + if (!client) { + return { + type: "list", + data: getStaticProviders(), + schema: providerLsSchema, + }; + } + + try { + const snapshot = await client.getProvidersSnapshot(); + return { + type: "list", + data: snapshot.entries.map((entry) => ({ + provider: entry.provider, + label: entry.label ?? entry.provider, + status: entry.status === "ready" ? "available" : entry.status, + defaultMode: entry.defaultModeId ?? "default", + modes: (entry.modes ?? []).map((mode) => mode.label).join(", "), + })), + schema: providerLsSchema, + }; + } catch { + return { + type: "list", + data: getStaticProviders(), + schema: providerLsSchema, + }; + } finally { + await client.close().catch(() => {}); + } } diff --git a/packages/cli/tests/15-provider.test.ts b/packages/cli/tests/15-provider.test.ts index c51f1d6a1..f08f9b420 100644 --- a/packages/cli/tests/15-provider.test.ts +++ b/packages/cli/tests/15-provider.test.ts @@ -129,7 +129,12 @@ try { assert(result.stdout.includes("claude"), "output should include claude"); assert(result.stdout.includes("codex"), "output should include codex"); assert(result.stdout.includes("opencode"), "output should include opencode"); - assert(result.stdout.includes("available"), "output should show available status"); + assert( + result.stdout.includes("available") || + result.stdout.includes("loading") || + result.stdout.includes("unavailable"), + "output should show a provider status", + ); console.log("✓ provider ls lists all providers\n"); } @@ -140,7 +145,7 @@ try { assert.strictEqual(result.exitCode, 0, "should exit 0"); const data = JSON.parse(result.stdout.trim()); assert(Array.isArray(data), "output should be an array"); - assert.strictEqual(data.length, 5, "should have 5 providers"); + assert(data.length >= 3, `should have at least 3 providers, got ${data.length}`); assert( data.some((p: { provider: string }) => p.provider === "claude"), "should include claude", @@ -153,14 +158,6 @@ try { data.some((p: { provider: string }) => p.provider === "opencode"), "should include opencode", ); - assert( - data.some((p: { provider: string }) => p.provider === "copilot"), - "should include copilot", - ); - assert( - data.some((p: { provider: string }) => p.provider === "pi"), - "should include pi", - ); console.log("✓ provider ls --json outputs valid JSON\n"); } @@ -170,12 +167,10 @@ try { const result = await ctx.paseo(["provider", "ls", "--quiet"]); assert.strictEqual(result.exitCode, 0, "should exit 0"); const lines = result.stdout.trim().split("\n"); - assert.strictEqual(lines.length, 5, "should have 5 lines"); + assert(lines.length >= 3, `should have at least 3 lines, got ${lines.length}`); assert(lines.includes("claude"), "should include claude"); assert(lines.includes("codex"), "should include codex"); assert(lines.includes("opencode"), "should include opencode"); - assert(lines.includes("copilot"), "should include copilot"); - assert(lines.includes("pi"), "should include pi"); console.log("✓ provider ls --quiet outputs provider names only\n"); } diff --git a/packages/server/src/server/agent/agent-management-mcp.ts b/packages/server/src/server/agent/agent-management-mcp.ts index 9ada46dec..7a009f5e7 100644 --- a/packages/server/src/server/agent/agent-management-mcp.ts +++ b/packages/server/src/server/agent/agent-management-mcp.ts @@ -48,7 +48,7 @@ import type { TerminalManager } from "../../terminal/terminal-manager.js"; import { createAgentWorktree, runAsyncWorktreeBootstrap } from "../worktree-bootstrap.js"; import type { ScheduleService } from "../schedule/service.js"; import { ScheduleSummarySchema, StoredScheduleSchema } from "../schedule/types.js"; -import { AGENT_PROVIDER_DEFINITIONS, type ProviderDefinition } from "./provider-registry.js"; +import type { ProviderDefinition } from "./provider-registry.js"; import { AgentModelSchema, AgentProviderEnum, @@ -832,7 +832,7 @@ export async function createAgentManagementMcpServer( async () => ({ content: [], structuredContent: ensureValidJson({ - providers: AGENT_PROVIDER_DEFINITIONS.map((provider) => ({ + providers: Object.values(providerRegistry ?? {}).map((provider) => ({ id: provider.id, label: provider.label, modes: provider.modes.map((mode) => ({ diff --git a/packages/server/src/server/agent/agent-manager.test.ts b/packages/server/src/server/agent/agent-manager.test.ts index 2a104e09d..a10735909 100644 --- a/packages/server/src/server/agent/agent-manager.test.ts +++ b/packages/server/src/server/agent/agent-manager.test.ts @@ -273,6 +273,156 @@ describe("AgentManager", () => { expect(snapshot.config.modeId).toBe("auto"); }); + test("setAgentMode persists the selected mode across session reload", async () => { + const workdir = mkdtempSync(join(tmpdir(), "agent-manager-test-")); + const storagePath = join(workdir, "agents"); + const storage = new AgentStorage(storagePath, logger); + + class ModeAwareSession implements AgentSession { + readonly provider = "codex" as const; + readonly capabilities = TEST_CAPABILITIES; + readonly id = randomUUID(); + private currentMode: string | null; + + constructor(private readonly config: AgentSessionConfig) { + this.currentMode = config.modeId ?? null; + } + + async run(): Promise { + return { sessionId: this.id, finalText: "", timeline: [] }; + } + + async startTurn(): Promise<{ turnId: string }> { + return { turnId: "turn-1" }; + } + + subscribe(): () => void { + return () => {}; + } + + async *streamHistory(): AsyncGenerator {} + + async getRuntimeInfo() { + return { + provider: this.provider, + sessionId: this.id, + model: this.config.model ?? null, + modeId: this.currentMode, + }; + } + + async getAvailableModes() { + return []; + } + + async getCurrentMode() { + return this.currentMode; + } + + async setMode(modeId: string): Promise { + this.currentMode = modeId; + } + + getPendingPermissions() { + return []; + } + + async respondToPermission(): Promise {} + + describePersistence() { + return { provider: this.provider, sessionId: this.id }; + } + + async interrupt(): Promise {} + async close(): Promise {} + } + + class ModeAwareClient implements AgentClient { + readonly provider = "codex" as const; + readonly capabilities = TEST_CAPABILITIES; + + async isAvailable(): Promise { + return true; + } + + async createSession(config: AgentSessionConfig): Promise { + return new ModeAwareSession(config); + } + + async resumeSession( + _handle: AgentPersistenceHandle, + config?: Partial, + ): Promise { + return new ModeAwareSession({ + provider: "codex", + cwd: config?.cwd ?? workdir, + modeId: config?.modeId, + model: config?.model, + }); + } + + async listModels() { + return [{ provider: "codex", id: "gpt-5.4", label: "GPT-5.4", isDefault: true }]; + } + } + + const manager = new AgentManager({ + clients: { + codex: new ModeAwareClient(), + }, + registry: storage, + logger, + idFactory: () => "00000000-0000-4000-8000-000000000301", + }); + + const snapshot = await manager.createAgent({ + provider: "codex", + cwd: workdir, + modeId: "auto", + }); + + await manager.setAgentMode(snapshot.id, "full-access"); + + const beforeReload = manager.getAgent(snapshot.id); + expect(beforeReload?.config.modeId).toBe("full-access"); + expect(beforeReload?.currentModeId).toBe("full-access"); + + const reloaded = await manager.reloadAgentSession(snapshot.id); + expect(reloaded.config.modeId).toBe("full-access"); + expect(reloaded.currentModeId).toBe("full-access"); + }); + + test("listProviderAvailability uses registered client keys, including custom providers", async () => { + const customClient: AgentClient = { + provider: "zai", + capabilities: TEST_CAPABILITIES, + async isAvailable() { + return true; + }, + async createSession() { + throw new Error("not implemented"); + }, + async resumeSession() { + throw new Error("not implemented"); + }, + }; + + const manager = new AgentManager({ + clients: { + zai: customClient, + }, + logger, + }); + + await expect(manager.listProviderAvailability()).resolves.toEqual([ + { + provider: "zai", + available: true, + error: null, + }, + ]); + }); + test("createAgent passes daemon launch env through the provider launch context", async () => { const workdir = mkdtempSync(join(tmpdir(), "agent-manager-test-")); const storagePath = join(workdir, "agents"); diff --git a/packages/server/src/server/agent/agent-manager.ts b/packages/server/src/server/agent/agent-manager.ts index 014459b98..ba201b3ce 100644 --- a/packages/server/src/server/agent/agent-manager.ts +++ b/packages/server/src/server/agent/agent-manager.ts @@ -33,7 +33,7 @@ import type { PersistedAgentDescriptor, } from "./agent-sdk-types.js"; import type { AgentStorage } from "./agent-storage.js"; -import { AGENT_PROVIDER_IDS, getAgentProviderDefinition } from "./provider-manifest.js"; +import { getAgentProviderDefinition } from "./provider-manifest.js"; export { AGENT_LIFECYCLE_STATUSES, type AgentLifecycleStatus }; @@ -362,6 +362,10 @@ export class AgentManager { this.clients.set(provider, client); } + getRegisteredProviderIds(): AgentProvider[] { + return Array.from(this.clients.keys()); + } + setAgentAttentionCallback(callback: AgentAttentionCallback): void { this.onAgentAttention = callback; } @@ -501,8 +505,7 @@ export class AgentManager { } async listProviderAvailability(): Promise { - const checks = AGENT_PROVIDER_IDS.map(async (providerId) => { - const provider = providerId as AgentProvider; + const checks = Array.from(this.clients.keys()).map(async (provider) => { const client = this.clients.get(provider); if (!client) { return { @@ -950,6 +953,7 @@ export class AgentManager { async setAgentMode(agentId: string, modeId: string): Promise { const agent = this.requireAgent(agentId); await agent.session.setMode(modeId); + agent.config.modeId = modeId; agent.currentModeId = modeId; // Update runtimeInfo to reflect the new mode if (agent.runtimeInfo) { diff --git a/packages/server/src/server/agent/agent-sdk-types.ts b/packages/server/src/server/agent/agent-sdk-types.ts index 5530b3059..19e1cddf5 100644 --- a/packages/server/src/server/agent/agent-sdk-types.ts +++ b/packages/server/src/server/agent/agent-sdk-types.ts @@ -43,6 +43,8 @@ export type AgentMode = { id: string; label: string; description?: string; + icon?: string; + colorTier?: string; }; export type ProviderStatus = "ready" | "loading" | "error" | "unavailable"; @@ -73,6 +75,9 @@ export interface ProviderSnapshotEntry { models?: AgentModelDefinition[]; modes?: AgentMode[]; fetchedAt?: string; + label?: string; + description?: string; + defaultModeId?: string | null; } export type AgentFeatureToggle = { diff --git a/packages/server/src/server/agent/mcp-server.test.ts b/packages/server/src/server/agent/mcp-server.test.ts index 2d3a649dc..f2394be1b 100644 --- a/packages/server/src/server/agent/mcp-server.test.ts +++ b/packages/server/src/server/agent/mcp-server.test.ts @@ -7,6 +7,7 @@ import { createTestLogger } from "../../test-utils/test-logger.js"; import { createAgentMcpServer } from "./mcp-server.js"; import type { AgentManager, ManagedAgent } from "./agent-manager.js"; import type { AgentStorage } from "./agent-storage.js"; +import type { ProviderDefinition } from "./provider-registry.js"; type TestDeps = { agentManager: AgentManager; @@ -53,6 +54,20 @@ function createTestDeps(): TestDeps { }; } +function createProviderDefinition(overrides: Partial): ProviderDefinition { + return { + id: "claude", + label: "Claude", + description: "Test provider", + defaultModeId: "default", + modes: [], + createClient: vi.fn() as ProviderDefinition["createClient"], + fetchModels: vi.fn().mockResolvedValue([]), + fetchModes: vi.fn().mockResolvedValue([]), + ...overrides, + }; +} + describe("create_agent MCP tool", () => { const logger = createTestLogger(); const existingCwd = process.cwd(); @@ -213,6 +228,22 @@ describe("create_agent MCP tool", () => { ); }); + it("accepts custom provider IDs in create_agent input validation", async () => { + const { agentManager, agentStorage } = createTestDeps(); + const server = await createAgentMcpServer({ agentManager, agentStorage, logger }); + const tool = (server as any)._registeredTools["create_agent"]; + + const parsed = await tool.inputSchema.safeParseAsync({ + cwd: existingCwd, + title: "Custom provider agent", + initialMode: "default", + agentType: "zai", + initialPrompt: "Do work", + }); + + expect(parsed.success).toBe(true); + }); + it("allows caller agents to override cwd and applies caller context labels", async () => { const { agentManager, agentStorage, spies } = createTestDeps(); const baseDir = await mkdtemp(join(tmpdir(), "paseo-mcp-test-")); @@ -302,6 +333,52 @@ describe("create_agent MCP tool", () => { }); }); +describe("provider listing MCP tool", () => { + const logger = createTestLogger(); + + it("returns providers from the registry, including custom providers", async () => { + const { agentManager, agentStorage } = createTestDeps(); + const providerRegistry = { + claude: createProviderDefinition({ + id: "claude", + label: "Claude", + modes: [{ id: "default", label: "Default", description: "Built-in mode" }], + }), + zai: createProviderDefinition({ + id: "zai", + label: "ZAI", + description: "Custom Claude profile", + defaultModeId: "default", + modes: [{ id: "default", label: "Default", description: "Custom mode" }], + }), + }; + + const server = await createAgentMcpServer({ + agentManager, + agentStorage, + providerRegistry, + logger, + }); + const tool = (server as any)._registeredTools["list_providers"]; + const response = await tool.callback({}); + + expect(response.structuredContent).toEqual({ + providers: [ + { + id: "claude", + label: "Claude", + modes: [{ id: "default", label: "Default", description: "Built-in mode" }], + }, + { + id: "zai", + label: "ZAI", + modes: [{ id: "default", label: "Default", description: "Custom mode" }], + }, + ], + }); + }); +}); + describe("speak MCP tool", () => { const logger = createTestLogger(); diff --git a/packages/server/src/server/agent/mcp-server.ts b/packages/server/src/server/agent/mcp-server.ts index 422674651..b634eaa6f 100644 --- a/packages/server/src/server/agent/mcp-server.ts +++ b/packages/server/src/server/agent/mcp-server.ts @@ -29,7 +29,7 @@ import { captureTerminalLines } from "../../terminal/terminal.js"; import { createAgentWorktree, runAsyncWorktreeBootstrap } from "../worktree-bootstrap.js"; import type { ScheduleService } from "../schedule/service.js"; import { ScheduleSummarySchema, StoredScheduleSchema } from "../schedule/types.js"; -import { AGENT_PROVIDER_DEFINITIONS, type ProviderDefinition } from "./provider-registry.js"; +import type { ProviderDefinition } from "./provider-registry.js"; import { deletePaseoWorktree, listPaseoWorktrees } from "../../utils/worktree.js"; import { AgentModelSchema, @@ -1362,7 +1362,7 @@ export async function createAgentMcpServer(options: AgentMcpServerOptions): Prom async () => ({ content: [], structuredContent: ensureValidJson({ - providers: AGENT_PROVIDER_DEFINITIONS.map((provider) => ({ + providers: Object.values(providerRegistry ?? {}).map((provider) => ({ id: provider.id, label: provider.label, modes: provider.modes.map((mode) => ({ diff --git a/packages/server/src/server/agent/mcp-shared.ts b/packages/server/src/server/agent/mcp-shared.ts index a018b80ae..79d7172aa 100644 --- a/packages/server/src/server/agent/mcp-shared.ts +++ b/packages/server/src/server/agent/mcp-shared.ts @@ -1,20 +1,14 @@ import { z } from "zod"; import type { Logger } from "pino"; -import type { AgentPromptInput, AgentProvider, AgentPermissionRequest } from "./agent-sdk-types.js"; +import type { AgentPromptInput, AgentPermissionRequest } from "./agent-sdk-types.js"; import type { AgentManager, ManagedAgent, WaitForAgentResult } from "./agent-manager.js"; import { curateAgentActivity } from "./activity-curator.js"; -import { AGENT_PROVIDER_DEFINITIONS } from "./provider-registry.js"; import type { AgentStorage } from "./agent-storage.js"; import { serializeAgentSnapshot } from "../messages.js"; import { StoredScheduleSchema } from "../schedule/types.js"; -export const AgentProviderEnum = z.enum( - AGENT_PROVIDER_DEFINITIONS.map((definition) => definition.id) as [ - AgentProvider, - ...AgentProvider[], - ], -); +export const AgentProviderEnum = z.string(); export const AgentStatusEnum = z.enum(["initializing", "idle", "running", "error", "closed"]); @@ -22,6 +16,8 @@ export const ProviderModeSchema = z.object({ id: z.string(), label: z.string(), description: z.string().optional(), + icon: z.string().optional(), + colorTier: z.string().optional(), }); export const ProviderSummarySchema = z.object({ diff --git a/packages/server/src/server/agent/provider-launch-config.test.ts b/packages/server/src/server/agent/provider-launch-config.test.ts index a172ee0a7..1f210a44d 100644 --- a/packages/server/src/server/agent/provider-launch-config.test.ts +++ b/packages/server/src/server/agent/provider-launch-config.test.ts @@ -1,8 +1,10 @@ import { describe, expect, test, vi } from "vitest"; import { - resolveProviderCommandPrefix, applyProviderEnv, + migrateProviderSettings, + ProviderOverrideSchema, + resolveProviderCommandPrefix, type ProviderRuntimeSettings, } from "./provider-launch-config.js"; @@ -103,3 +105,175 @@ describe("applyProviderEnv", () => { expect(env.CLAUDE_CODE_ENABLE_SDK_FILE_CHECKPOINTING).toBeUndefined(); }); }); + +describe("ProviderOverrideSchema", () => { + test("accepts built-in override fields", () => { + const parsed = ProviderOverrideSchema.parse({ + command: ["custom-claude", "--json"], + env: { + FOO: "bar", + }, + enabled: false, + order: 2, + }); + + expect(parsed.command).toEqual(["custom-claude", "--json"]); + expect(parsed.env?.FOO).toBe("bar"); + expect(parsed.enabled).toBe(false); + expect(parsed.order).toBe(2); + }); + + test("accepts models with thinking options", () => { + const parsed = ProviderOverrideSchema.parse({ + models: [ + { + id: "zai-fast", + label: "ZAI Fast", + isDefault: true, + thinkingOptions: [ + { + id: "deep", + label: "Deep", + description: "Higher effort", + }, + ], + }, + ], + }); + + expect(parsed.models).toEqual([ + { + id: "zai-fast", + label: "ZAI Fast", + isDefault: true, + thinkingOptions: [ + { + id: "deep", + label: "Deep", + description: "Higher effort", + }, + ], + }, + ]); + }); +}); + +describe("migrateProviderSettings", () => { + const builtinProviderIds = ["claude", "codex", "copilot", "opencode", "pi"]; + + test("passes through entries already in the new format", () => { + const migrated = migrateProviderSettings( + { + zai: { + extends: "claude", + label: "ZAI", + command: ["zai"], + env: { + ZAI_KEY: "secret", + }, + }, + }, + builtinProviderIds, + ); + + expect(migrated).toEqual({ + zai: { + extends: "claude", + label: "ZAI", + command: ["zai"], + env: { + ZAI_KEY: "secret", + }, + }, + }); + }); + + test("migrates mode replace to command argv", () => { + const migrated = migrateProviderSettings( + { + claude: { + command: { + mode: "replace", + argv: ["docker", "run", "--rm", "claude"], + }, + }, + }, + builtinProviderIds, + ); + + expect(migrated).toEqual({ + claude: { + command: ["docker", "run", "--rm", "claude"], + }, + }); + }); + + test("migrates mode default by dropping command", () => { + const migrated = migrateProviderSettings( + { + codex: { + command: { + mode: "default", + }, + env: { + FOO: "bar", + }, + }, + }, + builtinProviderIds, + ); + + expect(migrated).toEqual({ + codex: { + env: { + FOO: "bar", + }, + }, + }); + }); + + test("drops append mode entries because they cannot be auto-migrated", () => { + const migrated = migrateProviderSettings( + { + claude: { + command: { + mode: "append", + args: ["--debug"], + }, + env: { + FOO: "bar", + }, + }, + }, + builtinProviderIds, + ); + + expect(migrated).toEqual({}); + }); + + test("preserves legacy env while migrating old entries", () => { + const migrated = migrateProviderSettings( + { + opencode: { + command: { + mode: "replace", + argv: ["opencode"], + }, + env: { + PATH: "/custom/bin", + }, + }, + }, + builtinProviderIds, + ); + + expect(migrated).toEqual({ + opencode: { + command: ["opencode"], + env: { + PATH: "/custom/bin", + }, + }, + }); + }); +}); diff --git a/packages/server/src/server/agent/provider-launch-config.ts b/packages/server/src/server/agent/provider-launch-config.ts index 66bf83944..cf7de3a9b 100644 --- a/packages/server/src/server/agent/provider-launch-config.ts +++ b/packages/server/src/server/agent/provider-launch-config.ts @@ -37,6 +37,38 @@ export const ProviderRuntimeSettingsSchema = z }) .strict(); +const ProviderProfileThinkingOptionSchema = z + .object({ + id: z.string(), + label: z.string(), + description: z.string().optional(), + isDefault: z.boolean().optional(), + }) + .strict(); + +export const ProviderProfileModelSchema = z + .object({ + id: z.string().min(1), + label: z.string().min(1), + description: z.string().optional(), + isDefault: z.boolean().optional(), + thinkingOptions: z.array(ProviderProfileThinkingOptionSchema).optional(), + }) + .strict(); + +export const ProviderOverrideSchema = z + .object({ + extends: z.string().optional(), + label: z.string().optional(), + description: z.string().optional(), + command: z.array(z.string().min(1)).min(1).optional(), + env: z.record(z.string()).optional(), + models: z.array(ProviderProfileModelSchema).optional(), + enabled: z.boolean().optional(), + order: z.number().optional(), + }) + .strict(); + export const AgentProviderRuntimeSettingsMapSchema = z.record( AgentProviderSchema, ProviderRuntimeSettingsSchema, @@ -44,6 +76,8 @@ export const AgentProviderRuntimeSettingsMapSchema = z.record( export type ProviderCommand = z.infer; export type ProviderRuntimeSettings = z.infer; +export type ProviderProfileModel = z.infer; +export type ProviderOverride = z.infer; export type AgentProviderRuntimeSettingsMap = Partial< Record >; @@ -77,6 +111,48 @@ export async function resolveProviderCommandPrefix( }; } +export function migrateProviderSettings( + raw: Record, + builtinProviderIds: string[], +): Record { + const migrated: Record = {}; + const builtinProviderIdSet = new Set(builtinProviderIds); + + for (const [providerId, value] of Object.entries(raw)) { + const parsedNew = ProviderOverrideSchema.safeParse(value); + if (parsedNew.success) { + migrated[providerId] = parsedNew.data; + continue; + } + + const parsedOld = ProviderRuntimeSettingsSchema.safeParse(value); + if (!parsedOld.success) { + continue; + } + + const nextEntry: ProviderOverride = {}; + const command = parsedOld.data.command; + if (command?.mode === "append") { + console.warn( + `[Config] Skipping legacy agents.providers.${providerId}.command append mode during provider override migration because it cannot be auto-migrated.`, + ); + continue; + } + if (command?.mode === "replace") { + nextEntry.command = command.argv; + } + if (parsedOld.data.env) { + nextEntry.env = parsedOld.data.env; + } + if (!builtinProviderIdSet.has(providerId) && nextEntry.extends === undefined) { + delete nextEntry.extends; + } + migrated[providerId] = nextEntry; + } + + return migrated; +} + // Env vars that indicate a running Claude Code session. If the daemon itself is // launched from inside Claude Code (e.g. by a Paseo agent), these leak into // child processes and cause "cannot be launched inside another session" errors. diff --git a/packages/server/src/server/agent/provider-manifest.ts b/packages/server/src/server/agent/provider-manifest.ts index 8d27d2daf..5884fd8e9 100644 --- a/packages/server/src/server/agent/provider-manifest.ts +++ b/packages/server/src/server/agent/provider-manifest.ts @@ -9,7 +9,7 @@ export interface AgentModeVisuals { colorTier: AgentModeColorTier; } -export interface AgentProviderModeDefinition extends AgentMode, AgentModeVisuals {} +export type AgentProviderModeDefinition = Omit & AgentModeVisuals; // TODO: `modes` should not be static. Providers (especially ACP) report their // own modes at runtime via session/new. We should fetch modes from the provider @@ -168,24 +168,35 @@ export const AGENT_PROVIDER_DEFINITIONS: AgentProviderDefinition[] = [ }, ]; -export function getAgentProviderDefinition(provider: string): AgentProviderDefinition { - const definition = AGENT_PROVIDER_DEFINITIONS.find((entry) => entry.id === provider); +export function getAgentProviderDefinition( + provider: string, + definitions: AgentProviderDefinition[] = AGENT_PROVIDER_DEFINITIONS, +): AgentProviderDefinition { + const definition = definitions.find((entry) => entry.id === provider); if (!definition) { throw new Error(`Unknown agent provider: ${provider}`); } return definition; } -export const AGENT_PROVIDER_IDS = AGENT_PROVIDER_DEFINITIONS.map((d) => d.id); +export const BUILTIN_PROVIDER_IDS = AGENT_PROVIDER_DEFINITIONS.map((d) => d.id); +export const AGENT_PROVIDER_IDS = BUILTIN_PROVIDER_IDS; export const AgentProviderSchema = z.string(); -export function isValidAgentProvider(value: string): boolean { - return AGENT_PROVIDER_IDS.includes(value); +export function isValidAgentProvider( + value: string, + validIds: Iterable = BUILTIN_PROVIDER_IDS, +): boolean { + return Array.isArray(validIds) ? validIds.includes(value) : new Set(validIds).has(value); } -export function getModeVisuals(provider: string, modeId: string): AgentModeVisuals | undefined { - const definition = AGENT_PROVIDER_DEFINITIONS.find((entry) => entry.id === provider); +export function getModeVisuals( + provider: string, + modeId: string, + definitions: AgentProviderDefinition[], +): AgentModeVisuals | undefined { + const definition = definitions.find((entry) => entry.id === provider); const mode = definition?.modes.find((m) => m.id === modeId); if (!mode) return undefined; return { icon: mode.icon, colorTier: mode.colorTier }; diff --git a/packages/server/src/server/agent/provider-registry.test.ts b/packages/server/src/server/agent/provider-registry.test.ts new file mode 100644 index 000000000..90e6265aa --- /dev/null +++ b/packages/server/src/server/agent/provider-registry.test.ts @@ -0,0 +1,584 @@ +import { beforeEach, describe, expect, test, vi } from "vitest"; + +import { createTestLogger } from "../../test-utils/test-logger.js"; +import type { AgentModelDefinition } from "./agent-sdk-types.js"; + +const mockState = vi.hoisted(() => { + type ConstructorEntry = { + runtimeSettings?: unknown; + }; + + return { + constructorArgs: { + claude: [] as ConstructorEntry[], + codex: [] as ConstructorEntry[], + copilot: [] as ConstructorEntry[], + opencode: [] as ConstructorEntry[], + pi: [] as ConstructorEntry[], + genericAcp: [] as Array<{ + command: string[]; + env?: Record; + }>, + }, + runtimeModels: new Map(), + reset() { + for (const key of Object.keys(this.constructorArgs) as Array< + keyof typeof this.constructorArgs + >) { + this.constructorArgs[key] = []; + } + this.runtimeModels.clear(); + }, + }; +}); + +vi.mock("./providers/claude-agent.js", () => ({ + ClaudeAgentClient: class ClaudeAgentClient { + readonly capabilities = { + supportsStreaming: true, + supportsSessionPersistence: true, + supportsDynamicModes: true, + supportsMcpServers: true, + supportsReasoningStream: true, + supportsToolInvocations: true, + }; + readonly provider = "claude"; + readonly runtimeSettings?: unknown; + + constructor(options: { runtimeSettings?: unknown }) { + this.runtimeSettings = options.runtimeSettings; + mockState.constructorArgs.claude.push({ + runtimeSettings: options.runtimeSettings, + }); + } + + async createSession(): Promise { + throw new Error("not implemented"); + } + + async resumeSession(): Promise { + throw new Error("not implemented"); + } + + async listModels(): Promise { + return mockState.runtimeModels.get(this.provider) ?? []; + } + + async listModes(): Promise<[]> { + return []; + } + + async isAvailable(): Promise { + return true; + } + }, +})); + +vi.mock("./providers/codex-app-server-agent.js", () => ({ + CodexAppServerAgentClient: class CodexAppServerAgentClient { + readonly capabilities = { + supportsStreaming: true, + supportsSessionPersistence: true, + supportsDynamicModes: true, + supportsMcpServers: true, + supportsReasoningStream: true, + supportsToolInvocations: true, + }; + readonly provider = "codex"; + readonly runtimeSettings?: unknown; + + constructor(_logger: unknown, runtimeSettings?: unknown) { + this.runtimeSettings = runtimeSettings; + mockState.constructorArgs.codex.push({ runtimeSettings }); + } + + async createSession(): Promise { + throw new Error("not implemented"); + } + + async resumeSession(): Promise { + throw new Error("not implemented"); + } + + async listModels(): Promise { + return mockState.runtimeModels.get(this.provider) ?? []; + } + + async listModes(): Promise<[]> { + return []; + } + + async isAvailable(): Promise { + return true; + } + }, +})); + +vi.mock("./providers/copilot-acp-agent.js", () => ({ + CopilotACPAgentClient: class CopilotACPAgentClient { + readonly capabilities = { + supportsStreaming: true, + supportsSessionPersistence: true, + supportsDynamicModes: true, + supportsMcpServers: true, + supportsReasoningStream: true, + supportsToolInvocations: true, + }; + readonly provider = "copilot"; + readonly runtimeSettings?: unknown; + + constructor(options: { runtimeSettings?: unknown }) { + this.runtimeSettings = options.runtimeSettings; + mockState.constructorArgs.copilot.push({ + runtimeSettings: options.runtimeSettings, + }); + } + + async createSession(): Promise { + throw new Error("not implemented"); + } + + async resumeSession(): Promise { + throw new Error("not implemented"); + } + + async listModels(): Promise { + return mockState.runtimeModels.get(this.provider) ?? []; + } + + async listModes(): Promise<[]> { + return []; + } + + async isAvailable(): Promise { + return true; + } + }, +})); + +vi.mock("./providers/opencode-agent.js", () => ({ + OpenCodeAgentClient: class OpenCodeAgentClient { + readonly capabilities = { + supportsStreaming: true, + supportsSessionPersistence: true, + supportsDynamicModes: true, + supportsMcpServers: true, + supportsReasoningStream: true, + supportsToolInvocations: true, + }; + readonly provider = "opencode"; + readonly runtimeSettings?: unknown; + + constructor(_logger: unknown, runtimeSettings?: unknown) { + this.runtimeSettings = runtimeSettings; + mockState.constructorArgs.opencode.push({ runtimeSettings }); + } + + async createSession(): Promise { + throw new Error("not implemented"); + } + + async resumeSession(): Promise { + throw new Error("not implemented"); + } + + async listModels(): Promise { + return mockState.runtimeModels.get(this.provider) ?? []; + } + + async listModes(): Promise<[]> { + return []; + } + + async isAvailable(): Promise { + return true; + } + }, + OpenCodeServerManager: { + getInstance: vi.fn(() => ({ + shutdown: vi.fn(), + })), + }, +})); + +vi.mock("./providers/pi-acp-agent.js", () => ({ + PiACPAgentClient: class PiACPAgentClient { + readonly capabilities = { + supportsStreaming: true, + supportsSessionPersistence: true, + supportsDynamicModes: true, + supportsMcpServers: true, + supportsReasoningStream: true, + supportsToolInvocations: true, + }; + readonly provider = "pi"; + readonly runtimeSettings?: unknown; + + constructor(options: { runtimeSettings?: unknown }) { + this.runtimeSettings = options.runtimeSettings; + mockState.constructorArgs.pi.push({ + runtimeSettings: options.runtimeSettings, + }); + } + + async createSession(): Promise { + throw new Error("not implemented"); + } + + async resumeSession(): Promise { + throw new Error("not implemented"); + } + + async listModels(): Promise { + return mockState.runtimeModels.get(this.provider) ?? []; + } + + async listModes(): Promise<[]> { + return []; + } + + async isAvailable(): Promise { + return true; + } + }, +})); + +vi.mock("./providers/generic-acp-agent.js", () => ({ + GenericACPAgentClient: class GenericACPAgentClient { + readonly capabilities = { + supportsStreaming: true, + supportsSessionPersistence: true, + supportsDynamicModes: true, + supportsMcpServers: true, + supportsReasoningStream: true, + supportsToolInvocations: true, + }; + readonly provider = "acp"; + readonly runtimeSettings?: unknown; + + constructor(options: { command: string[]; env?: Record }) { + this.runtimeSettings = { + command: { + mode: "replace", + argv: options.command, + }, + env: options.env, + }; + mockState.constructorArgs.genericAcp.push({ + command: options.command, + env: options.env, + }); + } + + async createSession(): Promise { + throw new Error("not implemented"); + } + + async resumeSession(): Promise { + throw new Error("not implemented"); + } + + async listModels(): Promise { + return mockState.runtimeModels.get(this.provider) ?? []; + } + + async listModes(): Promise<[]> { + return []; + } + + async isAvailable(): Promise { + return true; + } + }, +})); + +import { AGENT_PROVIDER_DEFINITIONS, buildProviderRegistry } from "./provider-registry.js"; + +describe("buildProviderRegistry", () => { + const logger = createTestLogger(); + + beforeEach(() => { + mockState.reset(); + }); + + test("builds registry with no overrides — same as built-in count", () => { + const registry = buildProviderRegistry(logger); + + expect(Object.keys(registry)).toHaveLength(AGENT_PROVIDER_DEFINITIONS.length); + }); + + test("built-in override applies command", () => { + buildProviderRegistry(logger, { + providerOverrides: { + claude: { + command: ["/opt/custom-claude", "--verbose"], + }, + }, + }); + + expect(mockState.constructorArgs.claude[0]).toEqual({ + runtimeSettings: { + command: { + mode: "replace", + argv: ["/opt/custom-claude", "--verbose"], + }, + env: undefined, + }, + }); + }); + + test("built-in override applies env", () => { + buildProviderRegistry(logger, { + providerOverrides: { + claude: { + env: { + CLAUDE_CONFIG_DIR: "/tmp/claude", + }, + }, + }, + }); + + expect(mockState.constructorArgs.claude[0]).toEqual({ + runtimeSettings: { + command: undefined, + env: { + CLAUDE_CONFIG_DIR: "/tmp/claude", + }, + }, + }); + }); + + test("new provider extending claude appears in registry", () => { + const registry = buildProviderRegistry(logger, { + providerOverrides: { + zai: { + extends: "claude", + label: "ZAI", + description: "Claude with ZAI defaults", + }, + }, + }); + + expect(registry.zai).toBeDefined(); + expect(registry.zai.label).toBe("ZAI"); + expect(registry.zai.description).toBe("Claude with ZAI defaults"); + expect(registry.zai.createClient(logger).provider).toBe("zai"); + }); + + test("new provider extending acp uses GenericACPAgentClient", () => { + const registry = buildProviderRegistry(logger, { + providerOverrides: { + "my-agent": { + extends: "acp", + label: "My Agent", + command: ["my-agent", "--acp"], + env: { + ACP_TOKEN: "secret", + }, + }, + }, + }); + + expect(registry["my-agent"].createClient(logger).provider).toBe("my-agent"); + expect(mockState.constructorArgs.genericAcp).toEqual([ + { + command: ["my-agent", "--acp"], + env: { + ACP_TOKEN: "secret", + }, + }, + { + command: ["my-agent", "--acp"], + env: { + ACP_TOKEN: "secret", + }, + }, + ]); + }); + + test('extends: "acp" without command throws', () => { + expect(() => + buildProviderRegistry(logger, { + providerOverrides: { + "my-agent": { + extends: "acp", + label: "My Agent", + }, + }, + }), + ).toThrowError("ACP provider 'my-agent' requires a command"); + }); + + test("custom provider without label throws", () => { + expect(() => + buildProviderRegistry(logger, { + providerOverrides: { + zai: { + extends: "claude", + }, + }, + }), + ).toThrowError("Custom provider 'zai' requires a label"); + }); + + test("enabled: false excludes provider from registry", () => { + const registry = buildProviderRegistry(logger, { + providerOverrides: { + claude: { + enabled: false, + }, + }, + }); + + expect(registry.claude).toBeUndefined(); + }); + + test("extension inherits base override — override claude command, zai extends claude gets overridden command", () => { + buildProviderRegistry(logger, { + providerOverrides: { + claude: { + command: ["/opt/custom-claude"], + }, + zai: { + extends: "claude", + label: "ZAI", + }, + }, + }); + + expect(mockState.constructorArgs.claude).toHaveLength(2); + expect( + mockState.constructorArgs.claude.every( + (entry) => + (entry.runtimeSettings as { command?: { argv?: string[] } }).command?.argv?.[0] === + "/opt/custom-claude", + ), + ).toBe(true); + }); + + describe("model merging", () => { + test("profile models replace runtime models", async () => { + mockState.runtimeModels.set("claude", [ + { + provider: "claude", + id: "runtime-pro", + label: "Runtime Pro", + }, + ]); + + const registry = buildProviderRegistry(logger, { + providerOverrides: { + claude: { + models: [ + { + id: "profile-fast", + label: "Profile Fast", + }, + ], + }, + }, + }); + + const models = await registry.claude.fetchModels(); + + expect(models.map((model) => model.id)).toEqual(["profile-fast"]); + }); + + test("profile models exclude runtime models entirely", async () => { + mockState.runtimeModels.set("claude", [ + { + provider: "claude", + id: "shared-model", + label: "Runtime Label", + }, + { + provider: "claude", + id: "runtime-only", + label: "Runtime Only", + }, + ]); + + const registry = buildProviderRegistry(logger, { + providerOverrides: { + claude: { + models: [ + { + id: "shared-model", + label: "Profile Label", + }, + ], + }, + }, + }); + + const models = await registry.claude.fetchModels(); + + expect(models).toEqual([ + { + provider: "claude", + id: "shared-model", + label: "Profile Label", + }, + ]); + }); + + test("profile isDefault preserved without runtime models", async () => { + mockState.runtimeModels.set("claude", [ + { + provider: "claude", + id: "runtime-default", + label: "Runtime Default", + isDefault: true, + }, + ]); + + const registry = buildProviderRegistry(logger, { + providerOverrides: { + claude: { + models: [ + { + id: "profile-default", + label: "Profile Default", + isDefault: true, + }, + ], + }, + }, + }); + + const models = await registry.claude.fetchModels(); + + expect(models).toEqual([ + { + provider: "claude", + id: "profile-default", + label: "Profile Default", + isDefault: true, + }, + ]); + }); + + test("no profile models — runtime models returned as-is", async () => { + mockState.runtimeModels.set("claude", [ + { + provider: "claude", + id: "runtime-default", + label: "Runtime Default", + isDefault: true, + }, + ]); + + const registry = buildProviderRegistry(logger); + const models = await registry.claude.fetchModels(); + + expect(models).toEqual([ + { + provider: "claude", + id: "runtime-default", + label: "Runtime Default", + isDefault: true, + }, + ]); + }); + }); +}); diff --git a/packages/server/src/server/agent/provider-registry.ts b/packages/server/src/server/agent/provider-registry.ts index 94b183f71..0707a8534 100644 --- a/packages/server/src/server/agent/provider-registry.ts +++ b/packages/server/src/server/agent/provider-registry.ts @@ -1,22 +1,34 @@ +import type { Logger } from "pino"; + import type { AgentClient, AgentMode, AgentModelDefinition, + AgentPersistenceHandle, AgentProvider, + AgentRuntimeInfo, + AgentSession, + AgentStreamEvent, ListModelsOptions, ListModesOptions, + ListPersistedAgentsOptions, + PersistedAgentDescriptor, } from "./agent-sdk-types.js"; -import type { AgentProviderRuntimeSettingsMap } from "./provider-launch-config.js"; -import type { Logger } from "pino"; - +import type { + AgentProviderRuntimeSettingsMap, + ProviderOverride, + ProviderProfileModel, + ProviderRuntimeSettings, +} from "./provider-launch-config.js"; import { ClaudeAgentClient } from "./providers/claude-agent.js"; import { CodexAppServerAgentClient } from "./providers/codex-app-server-agent.js"; -import { OpenCodeAgentClient, OpenCodeServerManager } from "./providers/opencode-agent.js"; import { CopilotACPAgentClient } from "./providers/copilot-acp-agent.js"; +import { GenericACPAgentClient } from "./providers/generic-acp-agent.js"; +import { OpenCodeAgentClient, OpenCodeServerManager } from "./providers/opencode-agent.js"; import { PiACPAgentClient } from "./providers/pi-acp-agent.js"; - import { AGENT_PROVIDER_DEFINITIONS, + BUILTIN_PROVIDER_IDS, getAgentProviderDefinition, type AgentProviderDefinition, } from "./provider-manifest.js"; @@ -31,30 +43,42 @@ export interface ProviderDefinition extends AgentProviderDefinition { fetchModes: (options?: ListModesOptions) => Promise; } -type BuildProviderRegistryOptions = { +export type BuildProviderRegistryOptions = { runtimeSettings?: AgentProviderRuntimeSettingsMap; + providerOverrides?: Record; }; type ProviderClientFactory = ( logger: Logger, - runtimeSettings?: AgentProviderRuntimeSettingsMap, + runtimeSettings?: ProviderRuntimeSettings, ) => AgentClient; +type ResolvedProvider = { + definition: AgentProviderDefinition; + runtimeSettings?: ProviderRuntimeSettings; + profileModels: ProviderProfileModel[]; + enabled: boolean; + createBaseClient: (logger: Logger) => AgentClient; +}; + const PROVIDER_CLIENT_FACTORIES: Record = { claude: (logger, runtimeSettings) => new ClaudeAgentClient({ logger, - runtimeSettings: runtimeSettings?.claude, + runtimeSettings, }), - codex: (logger, runtimeSettings) => new CodexAppServerAgentClient(logger, runtimeSettings?.codex), + codex: (logger, runtimeSettings) => new CodexAppServerAgentClient(logger, runtimeSettings), copilot: (logger, runtimeSettings) => new CopilotACPAgentClient({ logger, - runtimeSettings: runtimeSettings?.copilot, + runtimeSettings, }), - opencode: (logger, runtimeSettings) => new OpenCodeAgentClient(logger, runtimeSettings?.opencode), + opencode: (logger, runtimeSettings) => new OpenCodeAgentClient(logger, runtimeSettings), pi: (logger, runtimeSettings) => - new PiACPAgentClient({ logger, runtimeSettings: runtimeSettings?.pi }), + new PiACPAgentClient({ + logger, + runtimeSettings, + }), }; function getProviderClientFactory(provider: string): ProviderClientFactory { @@ -65,31 +89,366 @@ function getProviderClientFactory(provider: string): ProviderClientFactory { return factory; } +function toRuntimeSettings(override?: ProviderOverride): ProviderRuntimeSettings | undefined { + if (!override?.command && !override?.env) { + return undefined; + } + + return { + command: override.command + ? { + mode: "replace", + argv: override.command, + } + : undefined, + env: override.env, + }; +} + +function mergeRuntimeSettings( + base: ProviderRuntimeSettings | undefined, + override: ProviderRuntimeSettings | undefined, +): ProviderRuntimeSettings | undefined { + if (!base && !override) { + return undefined; + } + + return { + command: override?.command ?? base?.command, + env: + base?.env || override?.env + ? { + ...(base?.env ?? {}), + ...(override?.env ?? {}), + } + : undefined, + }; +} + +function applyOverrideToDefinition( + definition: AgentProviderDefinition, + override?: ProviderOverride, +): AgentProviderDefinition { + if (!override) { + return definition; + } + + return { + ...definition, + label: override.label ?? definition.label, + description: override.description ?? definition.description, + }; +} + +function createDerivedDefinition( + providerId: string, + baseDefinition: AgentProviderDefinition, + override: ProviderOverride, +): AgentProviderDefinition { + if (!override.label) { + throw new Error(`Custom provider '${providerId}' requires a label`); + } + + return { + ...baseDefinition, + id: providerId, + label: override.label, + description: override.description ?? baseDefinition.description, + }; +} + +function mapPersistenceHandle( + provider: AgentProvider, + handle: AgentPersistenceHandle | null, +): AgentPersistenceHandle | null { + if (!handle) { + return null; + } + + return { + ...handle, + provider, + }; +} + +function mapRuntimeInfo(provider: AgentProvider, runtimeInfo: AgentRuntimeInfo): AgentRuntimeInfo { + return { + ...runtimeInfo, + provider, + }; +} + +function mapStreamEvent(provider: AgentProvider, event: AgentStreamEvent): AgentStreamEvent { + return { + ...event, + provider, + }; +} + +function mapPersistedAgentDescriptor( + provider: AgentProvider, + descriptor: PersistedAgentDescriptor, +): PersistedAgentDescriptor { + return { + ...descriptor, + provider, + persistence: { + ...descriptor.persistence, + provider, + }, + }; +} + +function mapModel(provider: AgentProvider, model: AgentModelDefinition): AgentModelDefinition { + return { + ...model, + provider, + }; +} + +function mergeModels( + provider: AgentProvider, + profileModels: ProviderProfileModel[], + runtimeModels: AgentModelDefinition[], +): AgentModelDefinition[] { + if (profileModels.length === 0) { + return runtimeModels.map((model) => mapModel(provider, model)); + } + + return profileModels.map((model) => ({ + ...model, + provider, + })); +} + +function wrapSessionProvider(provider: AgentProvider, inner: AgentSession): AgentSession { + return { + provider, + id: inner.id, + capabilities: inner.capabilities, + get features() { + return inner.features; + }, + run: (prompt, options) => inner.run(prompt, options), + startTurn: (prompt, options) => inner.startTurn(prompt, options), + subscribe: (callback) => inner.subscribe((event) => callback(mapStreamEvent(provider, event))), + async *streamHistory() { + for await (const event of inner.streamHistory()) { + yield mapStreamEvent(provider, event); + } + }, + getRuntimeInfo: async () => mapRuntimeInfo(provider, await inner.getRuntimeInfo()), + getAvailableModes: () => inner.getAvailableModes(), + getCurrentMode: () => inner.getCurrentMode(), + setMode: (modeId) => inner.setMode(modeId), + getPendingPermissions: () => inner.getPendingPermissions(), + respondToPermission: (requestId, response) => inner.respondToPermission(requestId, response), + describePersistence: () => mapPersistenceHandle(provider, inner.describePersistence()), + interrupt: () => inner.interrupt(), + close: () => inner.close(), + listCommands: inner.listCommands?.bind(inner), + setModel: inner.setModel?.bind(inner), + setThinkingOption: inner.setThinkingOption?.bind(inner), + setFeature: inner.setFeature?.bind(inner), + }; +} + +function wrapClientProvider(provider: AgentProvider, inner: AgentClient): AgentClient { + const listPersistedAgents = inner.listPersistedAgents?.bind(inner); + + return { + provider, + capabilities: inner.capabilities, + createSession: async (config, launchContext) => + wrapSessionProvider( + provider, + await inner.createSession( + { + ...config, + provider: inner.provider, + }, + launchContext, + ), + ), + resumeSession: async (handle, overrides, launchContext) => + wrapSessionProvider( + provider, + await inner.resumeSession( + { + ...handle, + provider: inner.provider, + }, + overrides + ? { + ...overrides, + provider: inner.provider, + } + : undefined, + launchContext, + ), + ), + listModels: async (options) => + (await inner.listModels(options)).map((model) => mapModel(provider, model)), + listModes: inner.listModes?.bind(inner), + listPersistedAgents: listPersistedAgents + ? async (options?: ListPersistedAgentsOptions) => + (await listPersistedAgents(options)).map((descriptor) => + mapPersistedAgentDescriptor(provider, descriptor), + ) + : undefined, + isAvailable: () => inner.isAvailable(), + getDiagnostic: inner.getDiagnostic?.bind(inner), + }; +} + +function createRegistryEntry( + logger: Logger, + provider: AgentProvider, + resolved: ResolvedProvider, +): ProviderDefinition { + const modelClient = resolved.createBaseClient(logger); + + return { + ...resolved.definition, + createClient: (providerLogger: Logger) => { + const inner = resolved.createBaseClient(providerLogger); + return inner.provider === provider ? inner : wrapClientProvider(provider, inner); + }, + fetchModels: async (options?: ListModelsOptions) => + mergeModels(provider, resolved.profileModels, await modelClient.listModels(options)), + fetchModes: async (options?: ListModesOptions) => { + const modes = modelClient.listModes + ? await modelClient.listModes(options) + : resolved.definition.modes; + return modes.map((mode) => { + if (mode.icon && mode.colorTier) return mode; + const definitionMode = resolved.definition.modes.find((d) => d.id === mode.id); + if (!definitionMode) return mode; + return { + ...mode, + icon: mode.icon ?? definitionMode.icon, + colorTier: mode.colorTier ?? definitionMode.colorTier, + }; + }); + }, + }; +} + +function buildResolvedBuiltinProviders( + providerOverrides: Record, + runtimeSettings: AgentProviderRuntimeSettingsMap | undefined, +): Map { + const resolvedProviders = new Map(); + + for (const definition of AGENT_PROVIDER_DEFINITIONS) { + const override = providerOverrides[definition.id]; + const factory = getProviderClientFactory(definition.id); + const mergedRuntimeSettings = mergeRuntimeSettings( + runtimeSettings?.[definition.id], + toRuntimeSettings(override), + ); + + resolvedProviders.set(definition.id, { + definition: applyOverrideToDefinition(definition, override), + runtimeSettings: mergedRuntimeSettings, + profileModels: override?.models ?? [], + enabled: override?.enabled !== false, + createBaseClient: (logger) => factory(logger, mergedRuntimeSettings), + }); + } + + return resolvedProviders; +} + +function addDerivedProviders( + resolvedProviders: Map, + providerOverrides: Record, +): void { + for (const [providerId, override] of Object.entries(providerOverrides)) { + if (BUILTIN_PROVIDER_IDS.includes(providerId)) { + continue; + } + + if (!override.extends) { + throw new Error(`Custom provider '${providerId}' requires an extends value`); + } + + if (override.extends === "acp") { + if (!override.command) { + throw new Error(`ACP provider '${providerId}' requires a command`); + } + + resolvedProviders.set(providerId, { + definition: createDerivedDefinition( + providerId, + { + id: providerId, + label: override.label ?? providerId, + description: override.description ?? "Custom ACP provider", + defaultModeId: null, + modes: [], + }, + override, + ), + runtimeSettings: toRuntimeSettings(override), + profileModels: override.models ?? [], + enabled: override.enabled !== false, + createBaseClient: (logger) => + new GenericACPAgentClient({ + logger, + command: override.command!, + env: override.env, + }), + }); + continue; + } + + const baseProvider = resolvedProviders.get(override.extends); + if (!baseProvider) { + throw new Error( + `Custom provider '${providerId}' extends unknown provider '${override.extends}'`, + ); + } + + const mergedRuntimeSettings = mergeRuntimeSettings( + baseProvider.runtimeSettings, + toRuntimeSettings(override), + ); + const baseDefinition = baseProvider.definition; + const baseFactory = getProviderClientFactory(override.extends); + + resolvedProviders.set(providerId, { + definition: createDerivedDefinition(providerId, baseDefinition, override), + runtimeSettings: mergedRuntimeSettings, + profileModels: override.models ?? [], + enabled: override.enabled !== false, + createBaseClient: (logger) => baseFactory(logger, mergedRuntimeSettings), + }); + } +} + export function buildProviderRegistry( logger: Logger, options?: BuildProviderRegistryOptions, ): Record { const runtimeSettings = options?.runtimeSettings; + const providerOverrides = options?.providerOverrides ?? {}; + const resolvedProviders = buildResolvedBuiltinProviders(providerOverrides, runtimeSettings); + addDerivedProviders(resolvedProviders, providerOverrides); + return Object.fromEntries( - AGENT_PROVIDER_DEFINITIONS.map((definition) => { - const createClient = getProviderClientFactory(definition.id); - const modelClient = createClient(logger, runtimeSettings); - return [ - definition.id, - { - ...definition, - createClient: (providerLogger: Logger) => createClient(providerLogger, runtimeSettings), - fetchModels: (listOptions?: ListModelsOptions) => modelClient.listModels(listOptions), - fetchModes: (listOptions?: ListModesOptions) => - modelClient.listModes - ? modelClient.listModes(listOptions) - : Promise.resolve(definition.modes), - } satisfies ProviderDefinition, - ]; - }), + [...resolvedProviders.entries()] + .filter(([, resolved]) => resolved.enabled) + .map(([provider, resolved]) => [provider, createRegistryEntry(logger, provider, resolved)]), ) as Record; } +export function getProviderIds( + registry: Record, +): AgentProvider[] { + return Object.keys(registry); +} + // Deprecated: Use buildProviderRegistry instead export const PROVIDER_REGISTRY: Record = null as any; diff --git a/packages/server/src/server/agent/provider-snapshot-manager.test.ts b/packages/server/src/server/agent/provider-snapshot-manager.test.ts index 44110d40e..111637d19 100644 --- a/packages/server/src/server/agent/provider-snapshot-manager.test.ts +++ b/packages/server/src/server/agent/provider-snapshot-manager.test.ts @@ -19,6 +19,9 @@ type Deferred = { type MockProviderOptions = { provider: AgentProvider; + label?: string; + description?: string; + defaultModeId?: string | null; isAvailable?: () => Promise; fetchModels?: (cwd?: string) => Promise; fetchModes?: (cwd?: string) => Promise; @@ -58,10 +61,21 @@ describe("ProviderSnapshotManager", () => { const snapshot = manager.getSnapshot("/tmp/project"); - expect(snapshot).toEqual([ - { provider: "claude", status: "loading" }, - { provider: "codex", status: "loading" }, - ]); + expect(snapshot.map((entry) => entry.provider)).toEqual(["codex", "claude"]); + expect(getProviderEntry(snapshot, "claude")).toMatchObject({ + provider: "claude", + status: "loading", + label: "claude", + description: "claude test provider", + defaultModeId: null, + }); + expect(getProviderEntry(snapshot, "codex")).toMatchObject({ + provider: "codex", + status: "loading", + label: "codex", + description: "codex test provider", + defaultModeId: null, + }); await vi.waitFor(() => { expect(handles.claude?.isAvailable).toHaveBeenCalledTimes(1); @@ -101,12 +115,18 @@ describe("ProviderSnapshotManager", () => { status: "ready", models: [createModel("codex", "gpt-5.2")], modes: [createMode("auto")], + label: "codex", + description: "codex test provider", + defaultModeId: null, }); expect(getProviderEntry(snapshot, "claude")).toMatchObject({ provider: "claude", status: "ready", models: [createModel("claude", "sonnet")], modes: [createMode("default")], + label: "claude", + description: "claude test provider", + defaultModeId: null, }); expect(getProviderEntry(snapshot, "codex")?.fetchedAt).toEqual(expect.any(String)); @@ -126,7 +146,13 @@ describe("ProviderSnapshotManager", () => { await vi.waitFor(() => { expect(manager.getSnapshot("/tmp/project")).toEqual([ - { provider: "codex", status: "unavailable" }, + { + provider: "codex", + status: "unavailable", + label: "codex", + description: "codex test provider", + defaultModeId: null, + }, ]); }); @@ -155,6 +181,9 @@ describe("ProviderSnapshotManager", () => { provider: "codex", status: "error", error: "model lookup failed", + label: "codex", + description: "codex test provider", + defaultModeId: null, }, ]); }); @@ -236,7 +265,15 @@ describe("ProviderSnapshotManager", () => { }); manager.refresh("/tmp/project"); - expect(manager.getSnapshot("/tmp/project")).toEqual([{ provider: "codex", status: "loading" }]); + expect(manager.getSnapshot("/tmp/project")).toEqual([ + { + provider: "codex", + status: "loading", + label: "codex", + description: "codex test provider", + defaultModeId: null, + }, + ]); await vi.waitFor(() => { expect(getProviderEntry(manager.getSnapshot("/tmp/project"), "codex")?.models?.[0]?.id).toBe( @@ -265,7 +302,15 @@ describe("ProviderSnapshotManager", () => { manager.refresh("/tmp/project"); - expect(manager.getSnapshot("/tmp/project")).toEqual([{ provider: "codex", status: "loading" }]); + expect(manager.getSnapshot("/tmp/project")).toEqual([ + { + provider: "codex", + status: "loading", + label: "codex", + description: "codex test provider", + defaultModeId: null, + }, + ]); manager.refresh("/tmp/project"); manager.refresh("/tmp/project"); @@ -361,6 +406,89 @@ describe("ProviderSnapshotManager", () => { manager.destroy(); }); + + test("snapshot includes user-defined providers from the registry", async () => { + const { registry } = createRegistry([ + createMockProvider({ provider: "claude" }), + createMockProvider({ + provider: "zai", + label: "ZAI", + description: "Custom Claude profile", + defaultModeId: "default", + fetchModes: async () => [createMode("default")], + }), + ]); + const manager = new ProviderSnapshotManager(registry, createTestLogger()); + + manager.getSnapshot("/tmp/project"); + + await vi.waitFor(() => { + expect(getProviderEntry(manager.getSnapshot("/tmp/project"), "zai")?.status).toBe("ready"); + }); + + expect(getProviderEntry(manager.getSnapshot("/tmp/project"), "zai")).toMatchObject({ + provider: "zai", + status: "ready", + label: "ZAI", + description: "Custom Claude profile", + defaultModeId: "default", + }); + + manager.destroy(); + }); + + test("enabled false providers are omitted when absent from the registry", () => { + const { registry } = createRegistry([createMockProvider({ provider: "claude" })]); + const manager = new ProviderSnapshotManager(registry, createTestLogger()); + + const snapshot = manager.getSnapshot("/tmp/project"); + + expect(snapshot.map((entry) => entry.provider)).toEqual(["claude"]); + expect(getProviderEntry(snapshot, "zai")).toBeUndefined(); + + manager.destroy(); + }); + + test("snapshot entries include label and description from the registry", async () => { + const models = deferred(); + const modes = deferred(); + const { registry } = createRegistry([ + createMockProvider({ + provider: "zai", + label: "ZAI", + description: "Custom Claude profile", + defaultModeId: "plan", + fetchModels: async () => models.promise, + fetchModes: async () => modes.promise, + }), + ]); + const manager = new ProviderSnapshotManager(registry, createTestLogger()); + + expect(manager.getSnapshot("/tmp/project")).toEqual([ + { + provider: "zai", + status: "loading", + label: "ZAI", + description: "Custom Claude profile", + defaultModeId: "plan", + }, + ]); + + models.resolve([createModel("zai", "zai-fast")]); + modes.resolve([createMode("plan")]); + + await vi.waitFor(() => { + expect(getProviderEntry(manager.getSnapshot("/tmp/project"), "zai")).toMatchObject({ + provider: "zai", + status: "ready", + label: "ZAI", + description: "Custom Claude profile", + defaultModeId: "plan", + }); + }); + + manager.destroy(); + }); }); function deferred(): Deferred { @@ -403,9 +531,9 @@ function createMockProvider(options: MockProviderOptions): MockProviderHandle { const definition: ProviderDefinition = { id: options.provider, - label: options.provider, - description: `${options.provider} test provider`, - defaultModeId: null, + label: options.label ?? options.provider, + description: options.description ?? `${options.provider} test provider`, + defaultModeId: options.defaultModeId ?? null, modes: [], createClient: () => ({ diff --git a/packages/server/src/server/agent/provider-snapshot-manager.ts b/packages/server/src/server/agent/provider-snapshot-manager.ts index b23a21d36..150cde42b 100644 --- a/packages/server/src/server/agent/provider-snapshot-manager.ts +++ b/packages/server/src/server/agent/provider-snapshot-manager.ts @@ -5,7 +5,6 @@ import type { Logger } from "pino"; import type { AgentProvider, ProviderSnapshotEntry } from "./agent-sdk-types.js"; import type { ProviderDefinition } from "./provider-registry.js"; -import { AGENT_PROVIDER_IDS } from "./provider-manifest.js"; const DEFAULT_CWD_KEY = "__default__"; @@ -63,9 +62,13 @@ export class ProviderSnapshotManager { private createLoadingEntries(): Map { const entries = new Map(); for (const provider of this.getProviderIds()) { + const definition = this.providerRegistry[provider]; entries.set(provider, { provider, status: "loading", + label: definition?.label, + description: definition?.description, + defaultModeId: definition?.defaultModeId ?? null, }); } return entries; @@ -107,6 +110,9 @@ export class ProviderSnapshotManager { snapshot.set(provider, { provider, status: "loading", + label: definition.label, + description: definition.description, + defaultModeId: definition.defaultModeId, }); try { @@ -116,6 +122,9 @@ export class ProviderSnapshotManager { snapshot.set(provider, { provider, status: "unavailable", + label: definition.label, + description: definition.description, + defaultModeId: definition.defaultModeId, }); this.emitChange(cwdKey); return; @@ -132,6 +141,9 @@ export class ProviderSnapshotManager { models, modes, fetchedAt: new Date().toISOString(), + label: definition.label, + description: definition.description, + defaultModeId: definition.defaultModeId, }); this.emitChange(cwdKey); } catch (error) { @@ -139,6 +151,9 @@ export class ProviderSnapshotManager { provider, status: "error", error: toErrorMessage(error), + label: definition.label, + description: definition.description, + defaultModeId: definition.defaultModeId, }); this.logger.warn( { err: error, provider, cwd: cwdKey }, @@ -180,7 +195,7 @@ export class ProviderSnapshotManager { } private getProviderIds(): AgentProvider[] { - return AGENT_PROVIDER_IDS.filter((provider) => this.providerRegistry[provider]); + return Object.keys(this.providerRegistry) as AgentProvider[]; } } diff --git a/packages/server/src/server/agent/providers/generic-acp-agent.ts b/packages/server/src/server/agent/providers/generic-acp-agent.ts new file mode 100644 index 000000000..ea9423df9 --- /dev/null +++ b/packages/server/src/server/agent/providers/generic-acp-agent.ts @@ -0,0 +1,46 @@ +import type { Logger } from "pino"; + +import { isCommandAvailable } from "../../../utils/executable.js"; +import { ACPAgentClient } from "./acp-agent.js"; + +type GenericACPAgentClientOptions = { + logger: Logger; + command: string[]; + env?: Record; +}; + +export class GenericACPAgentClient extends ACPAgentClient { + private readonly command: [string, ...string[]]; + + constructor(options: GenericACPAgentClientOptions) { + if (options.command.length === 0) { + throw new Error("Generic ACP provider requires a non-empty command"); + } + + super({ + provider: "acp", + logger: options.logger, + runtimeSettings: { + command: { + mode: "replace", + argv: options.command, + }, + env: options.env, + }, + defaultCommand: options.command as [string, ...string[]], + }); + + this.command = options.command as [string, ...string[]]; + } + + protected override async resolveLaunchCommand(): Promise<{ command: string; args: string[] }> { + return { + command: this.command[0], + args: this.command.slice(1), + }; + } + + override async isAvailable(): Promise { + return isCommandAvailable(this.command[0]); + } +} diff --git a/packages/server/src/server/bootstrap.ts b/packages/server/src/server/bootstrap.ts index 61e47c497..a2a6b7223 100644 --- a/packages/server/src/server/bootstrap.ts +++ b/packages/server/src/server/bootstrap.ts @@ -114,7 +114,10 @@ import { startRelayTransport, type RelayTransportController } from "./relay-tran import { getOrCreateServerId } from "./server-id.js"; import { resolveDaemonVersion } from "./daemon-version.js"; import type { AgentClient, AgentProvider } from "./agent/agent-sdk-types.js"; -import type { AgentProviderRuntimeSettingsMap } from "./agent/provider-launch-config.js"; +import type { + AgentProviderRuntimeSettingsMap, + ProviderOverride, +} from "./agent/provider-launch-config.js"; import { isHostAllowed, type AllowedHostsConfig } from "./allowed-hosts.js"; type AgentMcpTransportMap = Map; @@ -177,6 +180,7 @@ export type PaseoDaemonConfig = { dictationFinalTimeoutMs?: number; downloadTokenTtlMs?: number; agentProviderSettings?: AgentProviderRuntimeSettingsMap; + providerOverrides?: Record; onLifecycleIntent?: (intent: DaemonLifecycleIntent) => void; }; @@ -351,6 +355,7 @@ export async function createPaseoDaemon( clients: { ...createAllClients(logger, { runtimeSettings: config.agentProviderSettings, + providerOverrides: config.providerOverrides, }), ...config.agentClients, }, @@ -359,6 +364,7 @@ export async function createPaseoDaemon( }); const providerRegistry = buildProviderRegistry(logger, { runtimeSettings: config.agentProviderSettings, + providerOverrides: config.providerOverrides, }); const terminalManager = createTerminalManager(); @@ -598,6 +604,7 @@ export async function createPaseoDaemon( finalTimeoutMs: config.dictationFinalTimeoutMs, }, config.agentProviderSettings, + config.providerOverrides, daemonVersion, (intent) => { try { @@ -677,6 +684,7 @@ export async function createPaseoDaemon( await agentStorage.flush().catch(() => undefined); await shutdownProviders(logger, { runtimeSettings: config.agentProviderSettings, + providerOverrides: config.providerOverrides, }); terminalManager.killAll(); speechService.stop(); diff --git a/packages/server/src/server/config.ts b/packages/server/src/server/config.ts index 3043c3c17..02fc0cf8e 100644 --- a/packages/server/src/server/config.ts +++ b/packages/server/src/server/config.ts @@ -4,6 +4,11 @@ import { z } from "zod"; import type { PaseoDaemonConfig } from "./bootstrap.js"; import { loadPersistedConfig } from "./persisted-config.js"; import type { AgentProvider } from "./agent/agent-sdk-types.js"; +import type { + AgentProviderRuntimeSettingsMap, + ProviderOverride, +} from "./agent/provider-launch-config.js"; +import { ProviderOverrideSchema } from "./agent/provider-launch-config.js"; import { AgentProviderSchema } from "./agent/provider-manifest.js"; import { resolveSpeechConfig } from "./speech/speech-config-resolver.js"; import { @@ -52,6 +57,55 @@ function parseOptionalVoiceLlmProvider(value: unknown): AgentProvider | null { return parsed.success ? parsed.data : null; } +function extractProviderOverrides( + providers: Record | undefined, +): Record | undefined { + if (!providers) { + return undefined; + } + + const providerOverrides = Object.entries(providers).flatMap(([providerId, provider]) => { + const parsed = ProviderOverrideSchema.safeParse(provider); + return parsed.success ? [[providerId, parsed.data] as const] : []; + }); + + return providerOverrides.length > 0 ? Object.fromEntries(providerOverrides) : undefined; +} + +function extractAgentProviderSettings( + providerOverrides: Record | undefined, +): AgentProviderRuntimeSettingsMap | undefined { + if (!providerOverrides) { + return undefined; + } + + const runtimeSettings = Object.entries(providerOverrides).flatMap(([providerId, provider]) => { + const parsedProviderId = AgentProviderSchema.safeParse(providerId); + if (!parsedProviderId.success || (!provider.command && !provider.env)) { + return []; + } + + return [ + [ + parsedProviderId.data, + { + command: provider.command + ? { + mode: "replace" as const, + argv: provider.command, + } + : undefined, + env: provider.env, + }, + ] as const, + ]; + }); + + return runtimeSettings.length > 0 + ? (Object.fromEntries(runtimeSettings) as AgentProviderRuntimeSettingsMap) + : undefined; +} + export function loadConfig( paseoHome: string, options?: { @@ -117,6 +171,9 @@ export function loadConfig( const voiceLlmProviderExplicit = envVoiceLlmProvider !== null || persistedVoiceLlmProvider !== null; const voiceLlmModel = persisted.features?.voiceMode?.llm?.model ?? null; + const providerOverrides = extractProviderOverrides( + persisted.agents?.providers as Record | undefined, + ); return { listen, @@ -140,6 +197,7 @@ export function loadConfig( voiceLlmProvider, voiceLlmProviderExplicit, voiceLlmModel, - agentProviderSettings: persisted.agents?.providers, + agentProviderSettings: extractAgentProviderSettings(providerOverrides), + providerOverrides, }; } diff --git a/packages/server/src/server/exports.ts b/packages/server/src/server/exports.ts index 002478abc..66e4ff4eb 100644 --- a/packages/server/src/server/exports.ts +++ b/packages/server/src/server/exports.ts @@ -30,7 +30,11 @@ export { } from "./speech/providers/local/sherpa/sherpa-runtime-env.js"; // Provider binary resolution -export { applyProviderEnv } from "./agent/provider-launch-config.js"; +export { + applyProviderEnv, + type ProviderOverride, + type ProviderProfileModel, +} from "./agent/provider-launch-config.js"; export { findExecutable, executableExists, @@ -42,6 +46,7 @@ export { execCommand, spawnProcess } from "../utils/spawn.js"; // Provider manifest (source of truth for provider definitions) export { AGENT_PROVIDER_DEFINITIONS, + BUILTIN_PROVIDER_IDS, type AgentProviderDefinition, } from "./agent/provider-manifest.js"; diff --git a/packages/server/src/server/persisted-config.test.ts b/packages/server/src/server/persisted-config.test.ts index 75b5085ab..cbeea46ad 100644 --- a/packages/server/src/server/persisted-config.test.ts +++ b/packages/server/src/server/persisted-config.test.ts @@ -3,7 +3,7 @@ import { describe, expect, test } from "vitest"; import { PersistedConfigSchema } from "./persisted-config.js"; describe("PersistedConfigSchema agent provider runtime settings", () => { - test("accepts provider command append args and env", () => { + test("legacy append entries are skipped during migration", () => { const parsed = PersistedConfigSchema.parse({ agents: { providers: { @@ -20,8 +20,7 @@ describe("PersistedConfigSchema agent provider runtime settings", () => { }, }); - expect(parsed.agents?.providers?.claude?.command?.mode).toBe("append"); - expect(parsed.agents?.providers?.claude?.env?.FOO).toBe("bar"); + expect(parsed.agents?.providers).toEqual({}); }); test("accepts provider command replace argv", () => { @@ -38,7 +37,12 @@ describe("PersistedConfigSchema agent provider runtime settings", () => { }, }); - expect(parsed.agents?.providers?.codex?.command?.mode).toBe("replace"); + expect(parsed.agents?.providers?.codex?.command).toEqual([ + "docker", + "run", + "--rm", + "my-codex-wrapper", + ]); }); test("rejects replace command without argv", () => { @@ -58,6 +62,286 @@ describe("PersistedConfigSchema agent provider runtime settings", () => { }); }); +describe("provider overrides (new format)", () => { + test("override built-in provider with command and env", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + command: ["/opt/custom/claude"], + env: { + ANTHROPIC_API_KEY: "sk-test", + }, + }, + }, + }, + }); + + expect(parsed.agents?.providers?.claude).toEqual({ + command: ["/opt/custom/claude"], + env: { + ANTHROPIC_API_KEY: "sk-test", + }, + }); + }); + + test("new provider extending claude with label", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + zai: { + extends: "claude", + label: "ZAI", + }, + }, + }, + }); + + expect(parsed.agents?.providers?.zai).toEqual({ + extends: "claude", + label: "ZAI", + }); + }); + + test("new provider extending acp with command", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + "my-agent": { + extends: "acp", + label: "My Agent", + command: ["my-agent", "--acp"], + }, + }, + }, + }); + + expect(parsed.agents?.providers?.["my-agent"]).toEqual({ + extends: "acp", + label: "My Agent", + command: ["my-agent", "--acp"], + }); + }); + + test("enabled: false accepted", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + enabled: false, + }, + }, + }, + }); + + expect(parsed.agents?.providers?.claude?.enabled).toBe(false); + }); + + test("models array accepted", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + zai: { + extends: "claude", + label: "ZAI", + models: [ + { + id: "zai-fast", + label: "ZAI Fast", + isDefault: true, + }, + ], + }, + }, + }, + }); + + expect(parsed.agents?.providers?.zai?.models).toEqual([ + { + id: "zai-fast", + label: "ZAI Fast", + isDefault: true, + }, + ]); + }); + + test("order field accepted", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + order: 1, + }, + }, + }, + }); + + expect(parsed.agents?.providers?.claude?.order).toBe(1); + }); + + test("new provider without extends → error", () => { + const result = PersistedConfigSchema.safeParse({ + agents: { + providers: { + zai: { + label: "ZAI", + }, + }, + }, + }); + + expect(result.success).toBe(false); + }); + + test("new provider without label → error", () => { + const result = PersistedConfigSchema.safeParse({ + agents: { + providers: { + zai: { + extends: "claude", + }, + }, + }, + }); + + expect(result.success).toBe(false); + }); + + test("extends: acp without command → error", () => { + const result = PersistedConfigSchema.safeParse({ + agents: { + providers: { + "my-agent": { + extends: "acp", + label: "My Agent", + }, + }, + }, + }); + + expect(result.success).toBe(false); + }); + + test("extends unknown provider → error", () => { + const result = PersistedConfigSchema.safeParse({ + agents: { + providers: { + zai: { + extends: "unknown", + label: "ZAI", + }, + }, + }, + }); + + expect(result.success).toBe(false); + }); + + test("invalid provider ID format → error", () => { + const result = PersistedConfigSchema.safeParse({ + agents: { + providers: { + ZAI: { + extends: "claude", + label: "ZAI", + }, + }, + }, + }); + + expect(result.success).toBe(false); + }); + + test("old format with mode: replace auto-migrates", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + command: { + mode: "replace", + argv: ["docker", "run", "--rm", "claude"], + }, + }, + }, + }, + }); + + expect(parsed.agents?.providers?.claude).toEqual({ + command: ["docker", "run", "--rm", "claude"], + }); + }); + + test("old format with mode: default auto-migrates", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + command: { + mode: "default", + }, + }, + }, + }, + }); + + expect(parsed.agents?.providers?.claude).toEqual({}); + }); + + test("old format env preserved during migration", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + command: { + mode: "default", + }, + env: { + FOO: "bar", + }, + }, + }, + }, + }); + + expect(parsed.agents?.providers?.claude).toEqual({ + env: { + FOO: "bar", + }, + }); + }); + + test("mixed old and new format entries both work", () => { + const parsed = PersistedConfigSchema.parse({ + agents: { + providers: { + claude: { + command: { + mode: "replace", + argv: ["custom-claude"], + }, + }, + zai: { + extends: "claude", + label: "ZAI", + command: ["zai"], + }, + }, + }, + }); + + expect(parsed.agents?.providers).toEqual({ + claude: { + command: ["custom-claude"], + }, + zai: { + extends: "claude", + label: "ZAI", + command: ["zai"], + }, + }); + }); +}); + describe("PersistedConfigSchema logging config", () => { test("accepts destination-specific logging config", () => { const parsed = PersistedConfigSchema.parse({ diff --git a/packages/server/src/server/persisted-config.ts b/packages/server/src/server/persisted-config.ts index 34b9515e8..5fd0483a2 100644 --- a/packages/server/src/server/persisted-config.ts +++ b/packages/server/src/server/persisted-config.ts @@ -2,7 +2,12 @@ import { existsSync, mkdirSync, readFileSync, writeFileSync } from "node:fs"; import path from "node:path"; import { z } from "zod"; -import { AgentProviderRuntimeSettingsMapSchema } from "./agent/provider-launch-config.js"; +import { + AgentProviderRuntimeSettingsMapSchema, + migrateProviderSettings, + ProviderOverrideSchema, +} from "./agent/provider-launch-config.js"; +import type { AgentProviderRuntimeSettingsMap } from "./agent/provider-launch-config.js"; const LogLevelSchema = z.enum(["trace", "debug", "info", "warn", "error", "fatal"]); const LogFormatSchema = z.enum(["pretty", "json"]); @@ -113,6 +118,107 @@ const FeatureVoiceModeSchema = z }) .strict(); +const BUILTIN_PROVIDER_IDS = ["claude", "codex", "copilot", "opencode", "pi"] as const; +const PROVIDER_ID_PATTERN = /^[a-z][a-z0-9-]*$/; + +const ProviderOverridesSchema = z + .record(z.string(), ProviderOverrideSchema) + .superRefine((providers, ctx) => { + const builtinProviderIdSet = new Set(BUILTIN_PROVIDER_IDS); + const validExtendsValues = new Set([...BUILTIN_PROVIDER_IDS, "acp"]); + + for (const [providerId, provider] of Object.entries(providers)) { + if (!PROVIDER_ID_PATTERN.test(providerId)) { + ctx.addIssue({ + code: z.ZodIssueCode.custom, + path: [providerId], + message: `Provider ID "${providerId}" must match ${PROVIDER_ID_PATTERN}.`, + }); + } + + const isBuiltinProvider = builtinProviderIdSet.has(providerId); + if (!isBuiltinProvider && !provider.extends) { + ctx.addIssue({ + code: z.ZodIssueCode.custom, + path: [providerId, "extends"], + message: `Custom provider "${providerId}" must declare extends.`, + }); + } + + if (!isBuiltinProvider && !provider.label) { + ctx.addIssue({ + code: z.ZodIssueCode.custom, + path: [providerId, "label"], + message: `Custom provider "${providerId}" must declare label.`, + }); + } + + if (provider.extends && !validExtendsValues.has(provider.extends)) { + ctx.addIssue({ + code: z.ZodIssueCode.custom, + path: [providerId, "extends"], + message: `Provider "${providerId}" extends unknown provider "${provider.extends}".`, + }); + } + + if (provider.extends === "acp" && !provider.command) { + ctx.addIssue({ + code: z.ZodIssueCode.custom, + path: [providerId, "command"], + message: `Provider "${providerId}" extending "acp" must declare command.`, + }); + } + } + }); + +function isLegacyProviderEntry(value: unknown): boolean { + if (!value || typeof value !== "object" || Array.isArray(value)) { + return false; + } + + const command = (value as Record).command; + if (!command || typeof command !== "object" || Array.isArray(command)) { + return false; + } + + return typeof (command as Record).mode === "string"; +} + +function normalizeAgentProviders(value: unknown): unknown { + if (!value || typeof value !== "object" || Array.isArray(value)) { + return value; + } + + const rawProviders = value as Record; + const hasLegacyEntries = Object.values(rawProviders).some((entry) => + isLegacyProviderEntry(entry), + ); + if (!hasLegacyEntries) { + return value; + } + + const legacyEntries: Record = {}; + const normalizedEntries: Record = {}; + + for (const [providerId, providerValue] of Object.entries(rawProviders)) { + if (isLegacyProviderEntry(providerValue)) { + legacyEntries[providerId] = providerValue; + continue; + } + normalizedEntries[providerId] = providerValue; + } + + const parsedLegacyEntries = AgentProviderRuntimeSettingsMapSchema.safeParse(legacyEntries); + if (!parsedLegacyEntries.success) { + return value; + } + + return { + ...normalizedEntries, + ...migrateProviderSettings(parsedLegacyEntries.data, [...BUILTIN_PROVIDER_IDS]), + }; +} + export const PersistedConfigSchema = z .object({ // v1 schema marker @@ -158,7 +264,7 @@ export const PersistedConfigSchema = z providers: ProvidersSchema.optional(), agents: z .object({ - providers: AgentProviderRuntimeSettingsMapSchema.optional(), + providers: z.preprocess(normalizeAgentProviders, ProviderOverridesSchema).optional(), }) .strict() .optional(), @@ -174,10 +280,16 @@ export const PersistedConfigSchema = z }) .strict(); -export type PersistedConfig = z.infer; +type PersistedConfigSchemaOutput = z.infer; + +export type PersistedConfig = Omit & { + agents?: Omit, "providers"> & { + providers?: AgentProviderRuntimeSettingsMap; + }; +}; const CONFIG_FILENAME = "config.json"; -const DEFAULT_PERSISTED_CONFIG: PersistedConfig = PersistedConfigSchema.parse({ +const DEFAULT_PERSISTED_CONFIG = PersistedConfigSchema.parse({ version: 1, daemon: { listen: "127.0.0.1:6767", @@ -191,7 +303,7 @@ const DEFAULT_PERSISTED_CONFIG: PersistedConfig = PersistedConfigSchema.parse({ app: { baseUrl: "https://app.paseo.sh", }, -}); +}) as PersistedConfig; type LoggerLike = { child(bindings: Record): LoggerLike; @@ -275,7 +387,7 @@ export function loadPersistedConfig(paseoHome: string, logger?: LoggerLike): Per } log?.info(`Loaded from ${configPath}`); - return result.data; + return result.data as PersistedConfig; } export function savePersistedConfig( diff --git a/packages/server/src/server/persistence-hooks.test.ts b/packages/server/src/server/persistence-hooks.test.ts index 6a65ce3f7..1554b8ade 100644 --- a/packages/server/src/server/persistence-hooks.test.ts +++ b/packages/server/src/server/persistence-hooks.test.ts @@ -16,6 +16,7 @@ import type { const testLogger = { child: () => testLogger, error: vi.fn(), + warn: vi.fn(), } as any; type ManagedAgentOverrides = Omit< @@ -208,4 +209,22 @@ describe("persistence hooks", () => { }, }); }); + + test("buildSessionConfig skips records whose provider is missing from the registry", () => { + const record = createRecord({ + id: "agent-missing-provider", + provider: "zai", + }); + + expect( + buildSessionConfig(record, { + validProviders: ["claude", "codex"], + logger: testLogger, + }), + ).toBeNull(); + expect(testLogger.warn).toHaveBeenCalledWith( + { agentId: "agent-missing-provider", provider: "zai" }, + "Skipping persisted agent with unknown provider 'zai'", + ); + }); }); diff --git a/packages/server/src/server/persistence-hooks.ts b/packages/server/src/server/persistence-hooks.ts index c197a1c2d..c14881b45 100644 --- a/packages/server/src/server/persistence-hooks.ts +++ b/packages/server/src/server/persistence-hooks.ts @@ -1,11 +1,11 @@ import type { AgentManager } from "./agent/agent-manager.js"; -import type { AgentSessionConfig } from "./agent/agent-sdk-types.js"; +import type { AgentProvider, AgentSessionConfig } from "./agent/agent-sdk-types.js"; import type { AgentStorage, StoredAgentRecord } from "./agent/agent-storage.js"; -import { isValidAgentProvider } from "./agent/provider-manifest.js"; type LoggerLike = { child(bindings: Record): LoggerLike; error(...args: any[]): void; + warn(...args: any[]): void; }; function getLogger(logger: LoggerLike): LoggerLike { @@ -15,6 +15,11 @@ function getLogger(logger: LoggerLike): LoggerLike { type AgentStoragePersistence = Pick; type AgentManagerStateSource = Pick; +type BuildSessionConfigOptions = { + validProviders?: Iterable; + logger?: LoggerLike; +}; + /** * Attach AgentStorage persistence to an AgentManager instance so every * agent_state snapshot is flushed to disk. @@ -51,9 +56,18 @@ export function buildConfigOverrides(record: StoredAgentRecord): Partial>(); -const DEFAULT_AGENT_PROVIDER = AGENT_PROVIDER_IDS[0]; +const DEFAULT_AGENT_PROVIDER = "claude"; // TODO: Remove once all app store clients are on >=0.1.45 and understand arbitrary provider strings. // Clients before 0.1.45 validate providers with z.enum(["claude", "codex", "opencode"]) and reject @@ -416,6 +418,7 @@ export type SessionOptions = { getSpeechReadiness?: () => SpeechReadinessSnapshot; }; agentProviderRuntimeSettings?: AgentProviderRuntimeSettingsMap; + providerOverrides?: Record; }; export type SessionLifecycleIntent = @@ -487,8 +490,20 @@ function convertPCMToWavBuffer( return wavBuffer; } -function coerceAgentProvider(logger: pino.Logger, value: string, agentId?: string): AgentProvider { - if (isValidAgentProvider(value)) { +function isRegisteredProvider( + providerRegistry: ReturnType, + value: string, +): boolean { + return Object.prototype.hasOwnProperty.call(providerRegistry, value); +} + +function coerceAgentProvider( + logger: pino.Logger, + providerRegistry: ReturnType, + value: string, + agentId?: string, +): AgentProvider { + if (isRegisteredProvider(providerRegistry, value)) { return value; } logger.warn( @@ -500,13 +515,14 @@ function coerceAgentProvider(logger: pino.Logger, value: string, agentId?: strin function toAgentPersistenceHandle( logger: pino.Logger, + providerRegistry: ReturnType, handle: StoredAgentRecord["persistence"], ): AgentPersistenceHandle | null { if (!handle) { return null; } const provider = handle.provider; - if (!isValidAgentProvider(provider)) { + if (!isRegisteredProvider(providerRegistry, provider)) { logger.warn({ provider }, `Ignoring persistence handle with unknown provider '${provider}'`); return null; } @@ -619,6 +635,7 @@ export class Session { private readonly unregisterVoiceCallerContext?: (agentId: string) => void; private readonly getSpeechReadiness?: () => SpeechReadinessSnapshot; private readonly agentProviderRuntimeSettings: AgentProviderRuntimeSettingsMap | undefined; + private readonly providerOverrides: Record | undefined; private voiceModeAgentId: string | null = null; private voiceModeBaseConfig: VoiceModeBaseConfig | null = null; @@ -652,6 +669,7 @@ export class Session { voiceBridge, dictation, agentProviderRuntimeSettings, + providerOverrides, } = options; this.clientId = clientId; this.appVersion = appVersion; @@ -707,6 +725,7 @@ export class Session { this.unregisterVoiceCallerContext = voiceBridge?.unregisterVoiceCallerContext; this.getSpeechReadiness = dictation?.getSpeechReadiness; this.agentProviderRuntimeSettings = agentProviderRuntimeSettings; + this.providerOverrides = providerOverrides; this.abortController = new AbortController(); this.sessionLogger = logger.child({ module: "session", @@ -715,6 +734,7 @@ export class Session { }); this.providerRegistry = buildProviderRegistry(this.sessionLogger, { runtimeSettings: this.agentProviderRuntimeSettings, + providerOverrides: this.providerOverrides, }); // Initialize per-session managers @@ -1071,10 +1091,20 @@ export class Session { const updatedAt = new Date(this.resolveStoredAgentPayloadUpdatedAt(record)); const lastUserMessageAt = record.lastUserMessageAt ? new Date(record.lastUserMessageAt) : null; - const provider = coerceAgentProvider(this.sessionLogger, record.provider, record.id); + const provider = coerceAgentProvider( + this.sessionLogger, + this.providerRegistry, + record.provider, + record.id, + ); const runtimeInfo = record.runtimeInfo ? { - provider: coerceAgentProvider(this.sessionLogger, record.runtimeInfo.provider, record.id), + provider: coerceAgentProvider( + this.sessionLogger, + this.providerRegistry, + record.runtimeInfo.provider, + record.id, + ), sessionId: record.runtimeInfo.sessionId, ...(Object.prototype.hasOwnProperty.call(record.runtimeInfo, "model") ? { model: record.runtimeInfo.model ?? null } @@ -1107,7 +1137,11 @@ export class Session { currentModeId: record.lastModeId ?? null, availableModes: [], pendingPermissions: [], - persistence: toAgentPersistenceHandle(this.sessionLogger, record.persistence), + persistence: toAgentPersistenceHandle( + this.sessionLogger, + this.providerRegistry, + record.persistence, + ), lastUsage: undefined, lastError: undefined, title: record.title ?? record.config?.title ?? null, @@ -1153,7 +1187,11 @@ export class Session { throw new Error(`Agent not found: ${agentId}`); } - const handle = toAgentPersistenceHandle(this.sessionLogger, record.persistence); + const handle = toAgentPersistenceHandle( + this.sessionLogger, + this.providerRegistry, + record.persistence, + ); let snapshot: ManagedAgent; if (handle) { snapshot = await this.agentManager.resumeAgentFromPersistence( @@ -1167,7 +1205,13 @@ export class Session { "Agent resumed from persistence", ); } else { - const config = buildSessionConfig(record); + const config = buildSessionConfig(record, { + validProviders: Object.keys(this.providerRegistry), + logger: this.sessionLogger, + }); + if (!config) { + throw new Error(`Agent ${agentId} references unavailable provider '${record.provider}'`); + } snapshot = await this.agentManager.createAgent(config, agentId, { labels: record.labels }); this.sessionLogger.info( { agentId, provider: record.provider }, @@ -3065,7 +3109,11 @@ export class Session { if (!record) { throw new Error(`Agent not found: ${agentId}`); } - const handle = toAgentPersistenceHandle(this.sessionLogger, record.persistence); + const handle = toAgentPersistenceHandle( + this.sessionLogger, + this.providerRegistry, + record.persistence, + ); if (!handle) { throw new Error(`Agent ${agentId} cannot be refreshed because it lacks persistence`); } diff --git a/packages/server/src/server/websocket-server.notifications.test.ts b/packages/server/src/server/websocket-server.notifications.test.ts index bf7f84a50..848881f0c 100644 --- a/packages/server/src/server/websocket-server.notifications.test.ts +++ b/packages/server/src/server/websocket-server.notifications.test.ts @@ -84,6 +84,7 @@ function createServer(agentManagerOverrides?: Record) { undefined, undefined, undefined, + undefined, "1.2.3-test", undefined, undefined, diff --git a/packages/server/src/server/websocket-server.relay-reconnect.test.ts b/packages/server/src/server/websocket-server.relay-reconnect.test.ts index a6c70e40a..84ce0b535 100644 --- a/packages/server/src/server/websocket-server.relay-reconnect.test.ts +++ b/packages/server/src/server/websocket-server.relay-reconnect.test.ts @@ -180,6 +180,7 @@ function createServer(options?: { speechReadiness?: SpeechReadinessSnapshot | nu undefined, undefined, undefined, + undefined, TEST_DAEMON_VERSION, undefined, undefined, diff --git a/packages/server/src/server/websocket-server.ts b/packages/server/src/server/websocket-server.ts index 259c9966d..411c8dc7e 100644 --- a/packages/server/src/server/websocket-server.ts +++ b/packages/server/src/server/websocket-server.ts @@ -27,7 +27,10 @@ import type { AllowedHostsConfig } from "./allowed-hosts.js"; import { isHostAllowed } from "./allowed-hosts.js"; import { Session, type SessionLifecycleIntent, type SessionRuntimeMetrics } from "./session.js"; import type { AgentProvider } from "./agent/agent-sdk-types.js"; -import type { AgentProviderRuntimeSettingsMap } from "./agent/provider-launch-config.js"; +import type { + AgentProviderRuntimeSettingsMap, + ProviderOverride, +} from "./agent/provider-launch-config.js"; import { ProviderSnapshotManager } from "./agent/provider-snapshot-manager.js"; import { buildProviderRegistry } from "./agent/provider-registry.js"; import { WorkspaceGitServiceImpl } from "./workspace-git-service.js"; @@ -249,6 +252,7 @@ export class VoiceAssistantWebSocketServer { private readonly voiceSpeakHandlers = new Map(); private readonly voiceCallerContexts = new Map(); private readonly agentProviderRuntimeSettings: AgentProviderRuntimeSettingsMap | undefined; + private readonly providerOverrides: Record | undefined; private readonly providerSnapshotManager: ProviderSnapshotManager; private readonly onLifecycleIntent: ((intent: SessionLifecycleIntent) => void) | null; private serverCapabilities: ServerCapabilities | undefined; @@ -294,6 +298,7 @@ export class VoiceAssistantWebSocketServer { finalTimeoutMs?: number; }, agentProviderRuntimeSettings?: AgentProviderRuntimeSettingsMap, + providerOverrides?: Record, daemonVersion?: string, onLifecycleIntent?: (intent: SessionLifecycleIntent) => void, projectRegistry?: ProjectRegistry, @@ -341,10 +346,12 @@ export class VoiceAssistantWebSocketServer { this.terminalManager = terminalManager ?? null; this.dictation = dictation ?? null; this.agentProviderRuntimeSettings = agentProviderRuntimeSettings; + this.providerOverrides = providerOverrides; const providerSnapshotLogger = this.logger.child({ module: "provider-snapshot-manager" }); this.providerSnapshotManager = new ProviderSnapshotManager( buildProviderRegistry(providerSnapshotLogger, { runtimeSettings: this.agentProviderRuntimeSettings, + providerOverrides: this.providerOverrides, }), providerSnapshotLogger, ); @@ -672,6 +679,7 @@ export class VoiceAssistantWebSocketServer { } : undefined, agentProviderRuntimeSettings: this.agentProviderRuntimeSettings, + providerOverrides: this.providerOverrides, }); connection = { diff --git a/packages/server/src/shared/messages.ts b/packages/server/src/shared/messages.ts index 207523510..63fdd7e6b 100644 --- a/packages/server/src/shared/messages.ts +++ b/packages/server/src/shared/messages.ts @@ -93,6 +93,8 @@ const AgentModeSchema: z.ZodType = z.object({ id: z.string(), label: z.string(), description: z.string().optional(), + icon: z.string().optional(), + colorTier: z.string().optional(), }); const ProviderStatusSchema: z.ZodType = z.enum([ @@ -154,6 +156,9 @@ const ProviderSnapshotEntrySchema: z.ZodType = z.object({ models: z.array(AgentModelDefinitionSchema).optional(), modes: z.array(AgentModeSchema).optional(), fetchedAt: z.string().optional(), + label: z.string().optional(), + description: z.string().optional(), + defaultModeId: z.string().nullable().optional(), }); const AgentCapabilityFlagsSchema: z.ZodType = z.object({