diff --git a/apps/cli/src/commands/daemon/start.ts b/apps/cli/src/commands/daemon/start.ts index cd2bb9bc8..6e1584a79 100644 --- a/apps/cli/src/commands/daemon/start.ts +++ b/apps/cli/src/commands/daemon/start.ts @@ -8,6 +8,7 @@ import { configureClientLoggerForService, createLogger, discoverClaudeCodeSkills, + discoverProviderModels, flushClientSentry, initClientSentry, } from "@first-tree/client"; @@ -366,6 +367,28 @@ export function registerDaemonStartCommand(daemon: Command): void { }).finally(() => capabilityRefresher.endInteractive(command.provider)); }); + // Host-local model catalog: web opens Model settings → server asks this + // daemon → we discover from the real provider and reply on the WS. + runtime.onProviderModelsList((command) => { + void (async () => { + try { + const catalog = await discoverProviderModels(command.provider); + runtime.sendProviderModelsResult(command.ref, catalog); + } catch (err) { + const message = err instanceof Error ? err.message : String(err); + writeStatus("⚠️", `provider-models:list failed for ${command.provider}: ${message}`); + runtime.sendProviderModelsResult(command.ref, { + provider: command.provider, + models: [], + defaultModelId: null, + fetchedAt: new Date().toISOString(), + source: "unavailable", + error: message, + }); + } + })(); + }); + await runtime.start(); // Post-register capabilities upload + arm the background poll — the diff --git a/apps/cli/src/core/client-runtime.ts b/apps/cli/src/core/client-runtime.ts index e437fefaf..36c0bffb9 100644 --- a/apps/cli/src/core/client-runtime.ts +++ b/apps/cli/src/core/client-runtime.ts @@ -8,6 +8,7 @@ import { getChildProcessRegistry, getHandlerFactory, hasHandler, + type ProviderModelsListCommand, type RuntimeAuthCommand, registerBuiltinHandlers, type UpdateHooks, @@ -17,6 +18,7 @@ import { AGENT_BIND_REJECT_REASONS, type AgentPinnedMessage, type ClientPausedReason, + type ProviderModelCatalog, type RuntimeProvider, runtimeProviderSchema, } from "@first-tree/shared"; @@ -284,6 +286,20 @@ export class ClientRuntime { this.connection.on("runtime-auth:start", callback); } + /** + * Register a handler for the server→client `provider-models:list` command. + * The daemon discovers models from the host-local provider and replies with + * `provider-models:result` on the same connection. + */ + onProviderModelsList(callback: (command: ProviderModelsListCommand) => void): void { + this.connection.on("provider-models:list", callback); + } + + /** Reply to a correlated `provider-models:list` with the discovered catalog. */ + sendProviderModelsResult(ref: string, catalog: ProviderModelCatalog): void { + this.connection.sendProviderModelsResult(ref, catalog); + } + addAgent(name: string, config: AgentConfig): void { if (this.agentNames.has(name)) return; // The runtime provider is a valid enum value, but this client build may not diff --git a/packages/client/package.json b/packages/client/package.json index d85e429e2..1891db3d3 100644 --- a/packages/client/package.json +++ b/packages/client/package.json @@ -37,6 +37,7 @@ "mdast-util-from-markdown": "^2.0.3", "pino": "^9.5.0", "semver": "^7.6.3", + "smol-toml": "^1.7.0", "ws": "^8.20.0", "yaml": "^2.8.3", "zod": "^4.0.0" diff --git a/packages/client/src/__tests__/discover-models.test.ts b/packages/client/src/__tests__/discover-models.test.ts new file mode 100644 index 000000000..efd9b5235 --- /dev/null +++ b/packages/client/src/__tests__/discover-models.test.ts @@ -0,0 +1,87 @@ +import { describe, expect, it } from "vitest"; +import { + parseCursorModelsOutput, + parseKimiConfigModels, + resolveKimiConfigPath, +} from "../runtime/capabilities/discover-models.js"; + +describe("parseCursorModelsOutput", () => { + it("parses id/label rows and marks the default", () => { + const parsed = parseCursorModelsOutput(`Available models + +auto - Auto (default) +gpt-5.2 - GPT-5.2 +composer-2.5 - Composer 2.5 +`); + expect(parsed.defaultModelId).toBe("auto"); + expect(parsed.models).toEqual([ + { id: "auto", label: "Auto", isDefault: true, hint: "default" }, + { id: "gpt-5.2", label: "GPT-5.2" }, + { id: "composer-2.5", label: "Composer 2.5" }, + ]); + }); +}); + +describe("resolveKimiConfigPath", () => { + it("defaults to ~/.kimi-code/config.toml", () => { + expect(resolveKimiConfigPath({}, "/home/user")).toBe("/home/user/.kimi-code/config.toml"); + }); + + it("honors KIMI_CODE_HOME relocation", () => { + expect(resolveKimiConfigPath({ KIMI_CODE_HOME: "/opt/kimi" }, "/home/user")).toBe("/opt/kimi/config.toml"); + }); + + it("trims KIMI_CODE_HOME and ignores empty", () => { + expect(resolveKimiConfigPath({ KIMI_CODE_HOME: " /custom/kimi " }, "/home/user")).toBe( + "/custom/kimi/config.toml", + ); + expect(resolveKimiConfigPath({ KIMI_CODE_HOME: " " }, "/home/user")).toBe("/home/user/.kimi-code/config.toml"); + }); +}); + +describe("parseKimiConfigModels", () => { + it('reads default_model and quoted [models."."] sections', () => { + const parsed = parseKimiConfigModels(` +default_model = "kimi-code/k3" + +[models."kimi-code/kimi-for-coding"] +provider = "managed:kimi-code" +model = "kimi-for-coding" +display_name = "K2.7 Coding" + +[models."kimi-code/k3"] +provider = "managed:kimi-code" +model = "k3" +display_name = "K3" +`); + expect(parsed.defaultModelId).toBe("kimi-code/k3"); + expect(parsed.models).toEqual([ + { id: "kimi-code/kimi-for-coding", label: "K2.7 Coding" }, + { id: "kimi-code/k3", label: "K3", isDefault: true, hint: "default" }, + ]); + }); + + it("accepts bare model aliases and marks default_model", () => { + const parsed = parseKimiConfigModels(` +default_model = "gemini-3-pro-preview" + +[models.gemini-3-pro-preview] +provider = "openai" +model = "gemini-3-pro-preview" +display_name = "Gemini 3 Pro" + +[models."kimi-code/k3"] +provider = "managed:kimi-code" +model = "k3" +display_name = "K3" +`); + expect(parsed.defaultModelId).toBe("gemini-3-pro-preview"); + expect(parsed.models).toEqual( + expect.arrayContaining([ + { id: "gemini-3-pro-preview", label: "Gemini 3 Pro", isDefault: true, hint: "default" }, + { id: "kimi-code/k3", label: "K3" }, + ]), + ); + expect(parsed.models).toHaveLength(2); + }); +}); diff --git a/packages/client/src/client-connection.ts b/packages/client/src/client-connection.ts index bef230b8e..0a51c0d52 100644 --- a/packages/client/src/client-connection.ts +++ b/packages/client/src/client-connection.ts @@ -16,6 +16,10 @@ import { inboxDeliverFrameSchema, inboxRecoverAcceptedFrameSchema, inboxRecoverRejectedFrameSchema, + PROVIDER_MODELS_LIST_TYPE, + PROVIDER_MODELS_RESULT_TYPE, + type ProviderModelCatalog, + providerModelsListCommandSchema, RUNTIME_AUTH_START_TYPE, type RuntimeAuthMethod, type RuntimeProvider, @@ -164,6 +168,12 @@ export type RuntimeAuthCommand = { ref: string; }; +/** Server→client command to discover host-local provider models. */ +export type ProviderModelsListCommand = { + provider: RuntimeProvider; + ref: string; +}; + /** * Welcome frame received after `auth:ok`. `isReconnect` is true for every * occurrence after the first welcome in the lifetime of this `ClientConnection` @@ -199,6 +209,7 @@ type ClientConnectionEvents = { "agent:pinned": [message: AgentPinnedMessage]; "session:command": [command: SessionCommand]; "runtime-auth:start": [command: RuntimeAuthCommand]; + "provider-models:list": [command: ProviderModelsListCommand]; "session:reconcile:result": [result: SessionReconcileResult]; "auth:expired": []; /** @@ -1082,6 +1093,12 @@ export class ClientConnection extends EventEmitter { this.ws.send(JSON.stringify({ type: "session:reconcile", agentId, chatIds })); } + /** Reply to a `provider-models:list` reverse command with the host catalog. */ + sendProviderModelsResult(ref: string, catalog: ProviderModelCatalog): void { + if (!this.ws || this.ws.readyState !== WebSocket.OPEN) return; + this.ws.send(JSON.stringify({ type: PROVIDER_MODELS_RESULT_TYPE, ref, catalog })); + } + async disconnect(): Promise { this.closing = true; this.connectAbort?.abort(); @@ -1579,6 +1596,15 @@ export class ClientConnection extends EventEmitter { return; } + if (type === PROVIDER_MODELS_LIST_TYPE) { + const parsed = providerModelsListCommandSchema.safeParse(msg); + if (parsed.success) { + const { provider, ref } = parsed.data; + this.emit("provider-models:list", { provider, ref }); + } + return; + } + if (type === "session:reconcile:result") { const agentId = msg.agentId as string; const staleChatIds = Array.isArray(msg.staleChatIds) ? (msg.staleChatIds as string[]) : null; diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts index 9afefb6fb..8e8af2381 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -1,6 +1,7 @@ export type { BoundAgent, ClientConnectionConfig, + ProviderModelsListCommand, RuntimeAuthCommand, ServerWelcome, SessionCommand, @@ -42,6 +43,12 @@ export { resolveCodexRuntimeBinary, } from "./runtime/capabilities/codex.js"; export { probeCursorCapability } from "./runtime/capabilities/cursor.js"; +export { + type DiscoverModelsDeps, + discoverProviderModels, + parseCursorModelsOutput, + parseKimiConfigModels, +} from "./runtime/capabilities/discover-models.js"; export { CAPABILITY_REFRESH_BASE_MS, CAPABILITY_REFRESH_MAX_MS, @@ -54,6 +61,7 @@ export { revalidateCapabilities, shouldFullReprobe, } from "./runtime/capabilities/index.js"; + export type { AdoptOptions, ChildCategory, diff --git a/packages/client/src/runtime/capabilities/discover-models.ts b/packages/client/src/runtime/capabilities/discover-models.ts new file mode 100644 index 000000000..199b9407a --- /dev/null +++ b/packages/client/src/runtime/capabilities/discover-models.ts @@ -0,0 +1,216 @@ +import { readFile } from "node:fs/promises"; +import { homedir } from "node:os"; +import { join } from "node:path"; +import type { ProviderModelCatalog, ProviderModelOption, RuntimeProvider } from "@first-tree/shared"; +import { parse as parseToml } from "smol-toml"; +import { findCursorExecutableOnPath } from "../cursor-binary.js"; +import { runCommand } from "./launch-probe.js"; + +/** Ceiling for `agent models` — account catalog fetch can be network-bound. */ +const CURSOR_MODELS_TIMEOUT_MS = 20_000; + +export type DiscoverModelsDeps = { + env?: NodeJS.ProcessEnv; + now?: () => Date; + findCursorBinary?: (env?: Record) => string | null; + runCursorModels?: ( + binary: string, + env: NodeJS.ProcessEnv, + ) => Promise<{ ok: boolean; stdout: string; stderr: string }>; + readKimiConfig?: () => Promise; + kimiConfigPath?: string; +}; + +function fetchedAt(deps: DiscoverModelsDeps): string { + return (deps.now ?? (() => new Date()))().toISOString(); +} + +function unavailable(provider: RuntimeProvider, error: string, deps: DiscoverModelsDeps): ProviderModelCatalog { + return { + provider, + models: [], + defaultModelId: null, + fetchedAt: fetchedAt(deps), + source: "unavailable", + error, + }; +} + +/** + * Parse `agent models` / `agent --list-models` text: + * Available models + * auto - Auto (default) + * gpt-5.2 - GPT-5.2 + * + * Uses indexOf/slice instead of `\s+` / `.+` regexes so CodeQL does not flag + * polynomial-time matching on CLI stdout. + */ +export function parseCursorModelsOutput(stdout: string): { + models: ProviderModelOption[]; + defaultModelId: string | null; +} { + const models: ProviderModelOption[] = []; + let defaultModelId: string | null = null; + for (const rawLine of stdout.split(/\r?\n/)) { + const line = rawLine.trim(); + if (!line) continue; + if (line.toLowerCase() === "available models") continue; + const sep = line.indexOf(" - "); + if (sep <= 0) continue; + const id = line.slice(0, sep); + // Model ids are single tokens (`auto`, `gpt-5.2`); reject spaced ids. + if (!id || id.includes(" ") || id.includes("\t")) continue; + let label = line.slice(sep + 3).trim(); + const defaultMarker = "(default)"; + const defaultAt = label.toLowerCase().indexOf(defaultMarker); + const isDefault = defaultAt >= 0; + if (isDefault) { + defaultModelId = id; + label = `${label.slice(0, defaultAt)}${label.slice(defaultAt + defaultMarker.length)}`.trim(); + } + models.push({ + id, + label: label || id, + ...(isDefault ? { isDefault: true, hint: "default" } : {}), + }); + } + return { models, defaultModelId }; +} + +/** + * Parse Kimi Code `config.toml` model tables via a real TOML parser so we + * accept both quoted headers (`[models."kimi-code/k3"]`) and bare aliases + * (`[models.gemini-3-pro-preview]`) documented by Kimi. + */ +export function parseKimiConfigModels(toml: string): { + models: ProviderModelOption[]; + defaultModelId: string | null; +} { + let data: Record; + try { + data = parseToml(toml) as Record; + } catch { + return { models: [], defaultModelId: null }; + } + + const defaultModelId = typeof data.default_model === "string" ? data.default_model : null; + const modelsRaw = data.models; + if (!modelsRaw || typeof modelsRaw !== "object" || Array.isArray(modelsRaw)) { + return { models: [], defaultModelId }; + } + + const models: ProviderModelOption[] = []; + for (const [id, value] of Object.entries(modelsRaw as Record)) { + if (!id || typeof value !== "object" || value === null || Array.isArray(value)) continue; + const row = value as Record; + const displayName = typeof row.display_name === "string" ? row.display_name : undefined; + const isDefault = defaultModelId === id; + models.push({ + id, + ...(displayName ? { label: displayName } : {}), + ...(isDefault ? { isDefault: true, hint: "default" } : {}), + }); + } + return { models, defaultModelId }; +} + +/** Effective Kimi config path: `$KIMI_CODE_HOME/config.toml` or `~/.kimi-code/config.toml`. */ +export function resolveKimiConfigPath(env: NodeJS.ProcessEnv = process.env, home: string = homedir()): string { + const custom = env.KIMI_CODE_HOME?.trim(); + const root = custom && custom.length > 0 ? custom : join(home, ".kimi-code"); + return join(root, "config.toml"); +} + +async function discoverCursorModels(deps: DiscoverModelsDeps): Promise { + const env = deps.env ?? process.env; + const findBinary = deps.findCursorBinary ?? findCursorExecutableOnPath; + const binary = findBinary(env); + if (!binary) { + return unavailable("cursor", "cursor-agent / agent binary not found on this host", deps); + } + const run = + deps.runCursorModels ?? + (async (bin, processEnv) => { + const result = await runCommand(bin, ["models"], { timeoutMs: CURSOR_MODELS_TIMEOUT_MS, env: processEnv }); + return { ok: result.ok, stdout: result.stdout, stderr: result.stderr }; + }); + const result = await run(binary, env); + if (!result.ok) { + const detail = (result.stderr || result.stdout || "agent models failed").trim(); + return unavailable("cursor", detail.slice(0, 500), deps); + } + const parsed = parseCursorModelsOutput(result.stdout); + if (parsed.models.length === 0) { + return unavailable("cursor", "agent models returned no parseable model rows", deps); + } + return { + provider: "cursor", + models: parsed.models, + defaultModelId: parsed.defaultModelId, + fetchedAt: fetchedAt(deps), + source: "provider-cli", + error: null, + }; +} + +async function discoverKimiModels(deps: DiscoverModelsDeps): Promise { + const env = deps.env ?? process.env; + const path = deps.kimiConfigPath ?? resolveKimiConfigPath(env); + const read = + deps.readKimiConfig ?? + (async () => { + try { + return await readFile(path, "utf8"); + } catch (err) { + const code = err && typeof err === "object" && "code" in err ? String((err as { code: unknown }).code) : ""; + if (code === "ENOENT") return null; + throw err; + } + }); + let toml: string | null; + try { + toml = await read(); + } catch (err) { + return unavailable("kimi-code", err instanceof Error ? err.message : String(err), deps); + } + if (toml == null) { + return unavailable("kimi-code", `Kimi config not found at ${path}`, deps); + } + const parsed = parseKimiConfigModels(toml); + if (parsed.models.length === 0) { + return unavailable("kimi-code", "Kimi config has no [models.*] entries", deps); + } + return { + provider: "kimi-code", + models: parsed.models, + defaultModelId: parsed.defaultModelId, + fetchedAt: fetchedAt(deps), + source: "provider-config", + error: null, + }; +} + +/** + * Discover the model catalog for a runtime provider from the host-local + * provider. Phase 1 implements Cursor + Kimi; other providers return + * `source: "unavailable"` so the web can keep its curated/fallback UI. + */ +export async function discoverProviderModels( + provider: RuntimeProvider, + deps: DiscoverModelsDeps = {}, +): Promise { + switch (provider) { + case "cursor": + return discoverCursorModels(deps); + case "kimi-code": + return discoverKimiModels(deps); + case "claude-code": + case "claude-code-tui": + case "codex": + return unavailable(provider, `Host-local model discovery for ${provider} lands in a later phase`, deps); + default: { + const _exhaustive: never = provider; + return unavailable(_exhaustive, `Unknown provider: ${String(provider)}`, deps); + } + } +} diff --git a/packages/server/src/__tests__/clients-provider-models.test.ts b/packages/server/src/__tests__/clients-provider-models.test.ts new file mode 100644 index 000000000..beaae76ee --- /dev/null +++ b/packages/server/src/__tests__/clients-provider-models.test.ts @@ -0,0 +1,278 @@ +import crypto from "node:crypto"; +import { eq } from "drizzle-orm"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import type { WebSocket } from "ws"; +import { clients } from "../db/schema/clients.js"; +import * as clientService from "../services/client.js"; +import { + removeClientConnection, + resolveClientReply, + setClientConnection, + setClientReplyTimeoutMsForTests, + waitForClientReply, +} from "../services/connection-manager.js"; +import { readModelCatalogRpcResult, storeModelCatalogRpcResult } from "../services/provider-models-rpc.js"; +import { createAdminContext, useTestApp } from "./helpers.js"; + +/** + * `GET /api/v1/clients/:clientId/providers/:provider/models` asks the connected + * daemon for a host-local model catalog and waits for the correlated reply. + */ +describe("GET /clients/:clientId/providers/:provider/models", () => { + const getApp = useTestApp(); + + afterEach(() => { + setClientReplyTimeoutMsForTests(null); + }); + + async function markClientOnInstance(app: ReturnType, clientId: string, instanceId: string) { + await app.db.update(clients).set({ status: "connected", instanceId }).where(eq(clients.id, clientId)); + } + + it("forwards provider-models:list and returns the daemon catalog", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + await markClientOnInstance(app, admin.clientId, app.config.instanceId); + const ws = { readyState: 1, send: vi.fn(), close: vi.fn() }; + setClientConnection(admin.clientId, ws as unknown as WebSocket); + try { + const pending = app.inject({ + method: "GET", + url: `/api/v1/clients/${admin.clientId}/providers/cursor/models`, + headers: { authorization: `Bearer ${admin.accessToken}` }, + }); + + await vi.waitFor(() => { + expect(ws.send).toHaveBeenCalled(); + }); + const frame = JSON.parse(String(ws.send.mock.calls[0]?.[0])); + expect(frame).toMatchObject({ type: "provider-models:list", provider: "cursor" }); + expect(typeof frame.ref).toBe("string"); + + const catalog = { + provider: "cursor" as const, + models: [{ id: "auto", label: "Auto", isDefault: true }], + defaultModelId: "auto", + fetchedAt: new Date().toISOString(), + source: "provider-cli" as const, + error: null, + }; + expect(resolveClientReply(admin.clientId, frame.ref, catalog)).toBe(true); + + const res = await pending; + expect(res.statusCode).toBe(200); + expect(res.json()).toMatchObject(catalog); + } finally { + removeClientConnection(admin.clientId, ws as unknown as WebSocket); + } + }); + + it("returns 503 when the daemon is not connected", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + const res = await app.inject({ + method: "GET", + url: `/api/v1/clients/${admin.clientId}/providers/kimi-code/models`, + headers: { authorization: `Bearer ${admin.accessToken}` }, + }); + expect(res.statusCode).toBe(503); + }); + + it("fans the reverse command to the DB-authoritative instance when remote", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + await markClientOnInstance(app, admin.clientId, "replica-other"); + + const notifyCommand = vi.spyOn(app.notifier, "notifyDaemonClientCommand").mockResolvedValue(); + + const pending = app.inject({ + method: "GET", + url: `/api/v1/clients/${admin.clientId}/providers/cursor/models`, + headers: { authorization: `Bearer ${admin.accessToken}` }, + }); + + await vi.waitFor(() => { + expect(notifyCommand).toHaveBeenCalled(); + }); + const command = notifyCommand.mock.calls[0]?.[0]; + expect(command).toMatchObject({ + type: "provider-models:list", + clientId: admin.clientId, + provider: "cursor", + targetInstanceId: "replica-other", + }); + expect(typeof command?.ref).toBe("string"); + const ref = command?.ref; + if (!ref) throw new Error("expected notifyDaemonClientCommand ref"); + + const catalog = { + provider: "cursor" as const, + models: [{ id: "auto", label: "Auto", isDefault: true }], + defaultModelId: "auto", + fetchedAt: new Date().toISOString(), + source: "provider-cli" as const, + error: null, + }; + await storeModelCatalogRpcResult(app.db, admin.clientId, ref, catalog); + expect(await readModelCatalogRpcResult(app.db, admin.clientId, ref)).toMatchObject(catalog); + expect(resolveClientReply(admin.clientId, ref, catalog)).toBe(true); + + const res = await pending; + expect(res.statusCode).toBe(200); + expect(res.json()).toMatchObject(catalog); + notifyCommand.mockRestore(); + }); + + it("does not deliver on a stale local socket after instance takeover", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + // DB says another replica owns the connection, but this process still has a socket. + await markClientOnInstance(app, admin.clientId, "replica-other"); + const staleWs = { readyState: 1, send: vi.fn(), close: vi.fn() }; + setClientConnection(admin.clientId, staleWs as unknown as WebSocket); + + const notifyCommand = vi.spyOn(app.notifier, "notifyDaemonClientCommand").mockResolvedValue(); + try { + const pending = app.inject({ + method: "GET", + url: `/api/v1/clients/${admin.clientId}/providers/cursor/models`, + headers: { authorization: `Bearer ${admin.accessToken}` }, + }); + + await vi.waitFor(() => { + expect(notifyCommand).toHaveBeenCalled(); + }); + expect(staleWs.send).not.toHaveBeenCalled(); + expect(notifyCommand.mock.calls[0]?.[0]).toMatchObject({ + targetInstanceId: "replica-other", + clientId: admin.clientId, + }); + + const ref = notifyCommand.mock.calls[0]?.[0]?.ref; + if (!ref) throw new Error("expected ref"); + const catalog = { + provider: "cursor" as const, + models: [{ id: "auto", label: "Auto" }], + defaultModelId: "auto", + fetchedAt: new Date().toISOString(), + source: "provider-cli" as const, + error: null, + }; + expect(resolveClientReply(admin.clientId, ref, catalog)).toBe(true); + const res = await pending; + expect(res.statusCode).toBe(200); + } finally { + notifyCommand.mockRestore(); + removeClientConnection(admin.clientId, staleWs as unknown as WebSocket); + } + }); + + it("resolves a remote waiter from metadata after a result wake", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + const ref = crypto.randomUUID(); + const catalog = { + provider: "kimi-code" as const, + models: [{ id: "gemini-3-pro-preview", label: "Gemini 3 Pro", isDefault: true }], + defaultModelId: "gemini-3-pro-preview", + fetchedAt: new Date().toISOString(), + source: "provider-config" as const, + error: null, + }; + + const replyPromise = waitForClientReply(admin.clientId, ref); + await storeModelCatalogRpcResult(app.db, admin.clientId, ref, catalog); + await app.notifier.notifyDaemonClientCommandResult({ clientId: admin.clientId, ref }); + + await expect(replyPromise).resolves.toMatchObject(catalog); + }); + + it("returns a stored catalog when the result wake is lost", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + await markClientOnInstance(app, admin.clientId, app.config.instanceId); + setClientReplyTimeoutMsForTests(80); + + const ws = { readyState: 1, send: vi.fn(), close: vi.fn() }; + setClientConnection(admin.clientId, ws as unknown as WebSocket); + try { + const pending = app.inject({ + method: "GET", + url: `/api/v1/clients/${admin.clientId}/providers/cursor/models`, + headers: { authorization: `Bearer ${admin.accessToken}` }, + }); + + await vi.waitFor(() => { + expect(ws.send).toHaveBeenCalled(); + }); + const frame = JSON.parse(String(ws.send.mock.calls[0]?.[0])); + const catalog = { + provider: "cursor" as const, + models: [{ id: "auto", label: "Auto", isDefault: true }], + defaultModelId: "auto", + fetchedAt: new Date().toISOString(), + source: "provider-cli" as const, + error: null, + }; + // Durable store arrives, but no resolveClientReply / result NOTIFY (lost wake). + await storeModelCatalogRpcResult(app.db, admin.clientId, frame.ref, catalog); + + const res = await pending; + expect(res.statusCode).toBe(200); + expect(res.json()).toMatchObject(catalog); + } finally { + removeClientConnection(admin.clientId, ws as unknown as WebSocket); + } + }); + + it("stores concurrent refs without clobbering sibling metadata", async () => { + const app = getApp(); + const admin = await createAdminContext(app, { username: `pm-${crypto.randomUUID().slice(0, 6)}` }); + + const detectedAt = new Date().toISOString(); + await clientService.updateClientCapabilities(app.db, admin.clientId, { + cursor: { state: "ok", available: true, detectedAt, sdkVersion: "1.0.0" }, + }); + + const ref1 = crypto.randomUUID(); + const ref2 = crypto.randomUUID(); + const cat1 = { + provider: "cursor" as const, + models: [{ id: "auto", label: "Auto" }], + defaultModelId: "auto", + fetchedAt: new Date().toISOString(), + source: "provider-cli" as const, + error: null, + }; + const cat2 = { + provider: "kimi-code" as const, + models: [{ id: "k3", label: "K3" }], + defaultModelId: "k3", + fetchedAt: new Date().toISOString(), + source: "provider-config" as const, + error: null, + }; + + await Promise.all([ + storeModelCatalogRpcResult(app.db, admin.clientId, ref1, cat1), + storeModelCatalogRpcResult(app.db, admin.clientId, ref2, cat2), + ]); + + expect(await readModelCatalogRpcResult(app.db, admin.clientId, ref1)).toMatchObject(cat1); + expect(await readModelCatalogRpcResult(app.db, admin.clientId, ref2)).toMatchObject(cat2); + + await clientService.updateClientCapabilities(app.db, admin.clientId, { + cursor: { state: "ok", available: true, detectedAt, sdkVersion: "1.0.1" }, + "kimi-code": { state: "ok", available: true, detectedAt }, + }); + + expect(await readModelCatalogRpcResult(app.db, admin.clientId, ref1)).toMatchObject(cat1); + expect(await readModelCatalogRpcResult(app.db, admin.clientId, ref2)).toMatchObject(cat2); + + const row = await clientService.getClient(app.db, admin.clientId); + expect(clientService.extractCapabilities(row?.metadata)).toMatchObject({ + cursor: { available: true, sdkVersion: "1.0.1" }, + "kimi-code": { available: true }, + }); + }); +}); diff --git a/packages/server/src/__tests__/notifier-extra.test.ts b/packages/server/src/__tests__/notifier-extra.test.ts index 5a20be671..9f7b0d4c6 100644 --- a/packages/server/src/__tests__/notifier-extra.test.ts +++ b/packages/server/src/__tests__/notifier-extra.test.ts @@ -102,8 +102,16 @@ describe("createNotifier", () => { await notifier.notifyChatAudience("chat_1"); await notifier.notifyChatUpdated("chat_1"); await notifier.notifyAgentRouteChange(payload); + await notifier.notifyDaemonClientCommand({ + type: "provider-models:list", + clientId: "client_1", + provider: "cursor", + ref: "ref_1", + targetInstanceId: "instance_1", + }); + await notifier.notifyDaemonClientCommandResult({ clientId: "client_1", ref: "ref_1" }); - expect(ok.calls).toHaveLength(10); + expect(ok.calls).toHaveLength(12); expect(ok.calls.map((values) => values[0])).toEqual([ "inbox_notifications", "config_changes", @@ -115,6 +123,8 @@ describe("createNotifier", () => { "chat_audience_events", "chat_updated_events", "agent_route_events", + "daemon_client_commands", + "daemon_client_command_results", ]); const failing = createNotifier(makeListenClient(true).client as never); @@ -128,6 +138,18 @@ describe("createNotifier", () => { await expect(failing.notifyChatAudience("chat_1")).resolves.toBeUndefined(); await expect(failing.notifyChatUpdated("chat_1")).resolves.toBeUndefined(); await expect(failing.notifyAgentRouteChange(payload)).resolves.toBeUndefined(); + await expect( + failing.notifyDaemonClientCommand({ + type: "provider-models:list", + clientId: "client_1", + provider: "cursor", + ref: "ref_1", + targetInstanceId: "instance_1", + }), + ).resolves.toBeUndefined(); + await expect( + failing.notifyDaemonClientCommandResult({ clientId: "client_1", ref: "ref_1" }), + ).resolves.toBeUndefined(); }); it("parses LISTEN payloads, ignores malformed data, swallows handler errors, and stops idempotently", async () => { @@ -160,6 +182,14 @@ describe("createNotifier", () => { throw new Error("consumer failed"); }); const agentRouteSecond = vi.fn(); + const daemonCommand = vi.fn(() => { + throw new Error("consumer failed"); + }); + const daemonCommandSecond = vi.fn(); + const daemonResult = vi.fn(() => { + throw new Error("consumer failed"); + }); + const daemonResultSecond = vi.fn(); notifier.onConfigChange(config); notifier.onSessionStateChange(sessionState); @@ -176,6 +206,10 @@ describe("createNotifier", () => { notifier.onMeChatsChanged(meChatsChanged); notifier.onAgentRouteChange(agentRoute); notifier.onAgentRouteChange(agentRouteSecond); + notifier.onDaemonClientCommand(daemonCommand); + notifier.onDaemonClientCommand(daemonCommandSecond); + notifier.onDaemonClientCommandResult(daemonResult); + notifier.onDaemonClientCommandResult(daemonResultSecond); await notifier.start(); listeners.get("config_changes")?.("agent"); @@ -206,6 +240,20 @@ describe("createNotifier", () => { ); listeners.get("agent_route_events")?.("{not json"); listeners.get("agent_route_events")?.(JSON.stringify({ agentId: "agent_1" })); + listeners.get("daemon_client_commands")?.( + JSON.stringify({ + type: "provider-models:list", + clientId: "client_1", + provider: "cursor", + ref: "ref_1", + targetInstanceId: "instance_1", + }), + ); + listeners.get("daemon_client_commands")?.("{not json"); + listeners.get("daemon_client_commands")?.(JSON.stringify({ type: "provider-models:list" })); + listeners.get("daemon_client_command_results")?.(JSON.stringify({ clientId: "client_1", ref: "ref_1" })); + listeners.get("daemon_client_command_results")?.("{not json"); + listeners.get("daemon_client_command_results")?.(JSON.stringify({ clientId: "client_1" })); expect(config).toHaveBeenCalledWith("agent"); expect(sessionState).toHaveBeenCalledWith({ @@ -243,10 +291,18 @@ describe("createNotifier", () => { runtimeProvider: "codex", targetClientId: "client_1", }); + expect(daemonCommandSecond).toHaveBeenCalledWith({ + type: "provider-models:list", + clientId: "client_1", + provider: "cursor", + ref: "ref_1", + targetInstanceId: "instance_1", + }); + expect(daemonResultSecond).toHaveBeenCalledWith({ clientId: "client_1", ref: "ref_1" }); await notifier.stop(); await notifier.stop(); - expect(unlisteners).toHaveLength(11); + expect(unlisteners).toHaveLength(13); for (const unlisten of unlisteners) { expect(unlisten).toHaveBeenCalledTimes(1); } diff --git a/packages/server/src/api/agent/ws-client.ts b/packages/server/src/api/agent/ws-client.ts index ccb1e2275..5fa6533dc 100644 --- a/packages/server/src/api/agent/ws-client.ts +++ b/packages/server/src/api/agent/ws-client.ts @@ -12,6 +12,9 @@ import { inboxAckFrameSchema, inboxDeliverFrameSchema, inboxRecoverFrameSchema, + PROVIDER_MODELS_LIST_TYPE, + PROVIDER_MODELS_RESULT_TYPE, + providerModelsResultFrameSchema, runtimeStateMessageSchema, sessionEventMessageSchema, sessionEventRejectedReasonSchema, @@ -54,6 +57,7 @@ import * as landingCampaignChatStateService from "../../services/landing-campaig import * as notificationService from "../../services/notification.js"; import type { InboxPushHandler, Notifier } from "../../services/notifier.js"; import * as presenceService from "../../services/presence.js"; +import { readModelCatalogRpcResult, storeModelCatalogRpcResult } from "../../services/provider-models-rpc.js"; import * as runtimeLivenessService from "../../services/runtime-liveness.js"; import * as sessionEventService from "../../services/session-event.js"; @@ -293,6 +297,31 @@ export function clientWsRoutes(notifier: Notifier, instanceId: string) { connectionManager.sendToClient(payload.targetClientId, frame.data); }); + // Cross-replica reverse commands: only the DB-authoritative instance may + // deliver. A stale open socket on a previous replica must not receive the + // same ref after reconnect/takeover. + notifier.onDaemonClientCommand((payload) => { + if (payload.type !== PROVIDER_MODELS_LIST_TYPE) return; + if (payload.targetInstanceId !== instanceId) return; + connectionManager.sendToClient(payload.clientId, { + type: PROVIDER_MODELS_LIST_TYPE, + provider: payload.provider, + ref: payload.ref, + }); + }); + + // Cross-replica result wake: catalog is in clients.metadata; resolve any + // local HTTP waiter that registered waitForClientReply for this ref. + notifier.onDaemonClientCommandResult((payload) => { + void (async () => { + const catalog = await readModelCatalogRpcResult(app.db, payload.clientId, payload.ref); + if (!catalog) return; + connectionManager.resolveClientReply(payload.clientId, payload.ref, catalog); + })().catch((err) => { + app.log.debug({ err, clientId: payload.clientId, ref: payload.ref }, "provider-models result wake failed"); + }); + }); + // WS upgrade is excluded from HTTP tracing in app.ts via the autotelic // plugin's `ignoreRoutes` — fastify hijacks the reply on upgrade, so a // `onResponse`-terminated HTTP span would never end. The connection's @@ -1757,6 +1786,41 @@ export function clientWsRoutes(notifier: Notifier, instanceId: string) { await reconcilePinnedAgentsForClient(); } socket.send(JSON.stringify({ type: "heartbeat:ack" })); + } else if (type === PROVIDER_MODELS_RESULT_TYPE) { + if (!clientId) { + socket.send(JSON.stringify({ type: "error", message: "Must register client first" })); + return; + } + const result = providerModelsResultFrameSchema.safeParse(msg); + if (!result.success) { + socket.send(JSON.stringify({ type: "error", message: "Malformed provider-models:result frame" })); + return; + } + // Only the DB-authoritative instance accepts results — a replaced + // connection on a prior replica must not win the rendezvous. + const [owner] = await app.db + .select({ instanceId: clients.instanceId }) + .from(clients) + .where(eq(clients.id, clientId)) + .limit(1); + if (!owner || owner.instanceId !== instanceId) { + app.log.debug( + { clientId, ref: result.data.ref, instanceId, ownerInstanceId: owner?.instanceId ?? null }, + "ignoring provider-models:result from non-authoritative instance", + ); + return; + } + // Durable rendezvous first so another replica's HTTP waiter can + // load the catalog after the tiny result-wake NOTIFY. + await storeModelCatalogRpcResult(app.db, clientId, result.data.ref, result.data.catalog); + const resolved = connectionManager.resolveClientReply(clientId, result.data.ref, result.data.catalog); + await notifier.notifyDaemonClientCommandResult({ clientId, ref: result.data.ref }); + if (!resolved) { + app.log.debug( + { clientId, ref: result.data.ref }, + "provider-models:result matched no pending HTTP waiter on this replica", + ); + } } } catch (err) { const message = err instanceof Error ? err.message : "Internal error"; diff --git a/packages/server/src/api/clients.ts b/packages/server/src/api/clients.ts index a2baf7c83..30a8d4020 100644 --- a/packages/server/src/api/clients.ts +++ b/packages/server/src/api/clients.ts @@ -1,7 +1,10 @@ import { randomUUID } from "node:crypto"; import { + PROVIDER_MODELS_LIST_TYPE, + providerModelCatalogSchema, RUNTIME_AUTH_START_TYPE, runtimeAuthStartRequestSchema, + runtimeProviderSchema, updateClientCapabilitiesSchema, } from "@first-tree/shared"; import { getChannelConfig } from "@first-tree/shared/channel"; @@ -11,7 +14,13 @@ import { stampClientResource } from "../observability/request-context.js"; import { requireUser } from "../scope/require-user.js"; import { expiryToSeconds } from "../services/auth.js"; import * as clientService from "../services/client.js"; -import { forceDisconnectClient, sendToClient } from "../services/connection-manager.js"; +import { + forceDisconnectClient, + rejectPendingRepliesForClient, + sendToClient, + waitForClientReply, +} from "../services/connection-manager.js"; +import { isClientConnectedSomewhere, readModelCatalogRpcResult } from "../services/provider-models-rpc.js"; import { serializeDate } from "../utils.js"; import { clientCommandVersionHint } from "./client-command-version.js"; @@ -91,6 +100,69 @@ export async function clientRoutes(app: FastifyInstance): Promise { return { ref, started: true as const }; }); + // Host-local model catalog: ask the connected daemon to discover models from + // the real provider on that computer, wait for the correlated reply, and + // return the catalog to the web. Delivery is scoped to the DB-authoritative + // `clients.instance_id` (local send or PG NOTIFY fan-out). Results are stored + // in clients.metadata; on waiter timeout we still read that durable copy so a + // lost NOTIFY does not false-503. Hard 503 only when offline / truly missing. + app.get<{ Params: { clientId: string; provider: string } }>( + "/:clientId/providers/:provider/models", + async (request) => { + const { userId } = requireUser(request); + const { clientId, provider: rawProvider } = request.params; + stampClientResource(request, clientId); + await clientService.assertClientOwner(app.db, clientId, { userId }); + await clientService.assertClientNotRetired(app.db, clientId); + const provider = runtimeProviderSchema.parse(rawProvider); + const client = await clientService.getClient(app.db, clientId); + if (!client || !isClientConnectedSomewhere(client) || !client.instanceId) { + throw new ServiceUnavailableError( + "Could not list models because this computer is not connected. Make sure the daemon is running, then retry.", + ); + } + const targetInstanceId = client.instanceId; + const ref = randomUUID(); + const replyPromise = waitForClientReply(clientId, ref); + const daemonFrame = { + type: PROVIDER_MODELS_LIST_TYPE, + provider, + ref, + }; + if (targetInstanceId === app.config.instanceId) { + const delivered = sendToClient(clientId, daemonFrame); + if (!delivered) { + rejectPendingRepliesForClient(clientId, new Error("Computer not connected")); + await replyPromise.catch(() => undefined); + throw new ServiceUnavailableError( + "Could not list models because this computer is not connected. Make sure the daemon is running, then retry.", + ); + } + } else { + await app.notifier.notifyDaemonClientCommand({ + type: PROVIDER_MODELS_LIST_TYPE, + clientId, + provider, + ref, + targetInstanceId, + }); + } + try { + const raw = await replyPromise; + return providerModelCatalogSchema.parse(raw); + } catch (err) { + // Race-safe fallback: catalog may already be durable while the wake was lost. + const stored = await readModelCatalogRpcResult(app.db, clientId, ref); + if (stored) return stored; + throw new ServiceUnavailableError( + err instanceof Error + ? err.message + : "Could not list models from this computer. Retry after the daemon is connected.", + ); + } + }, + ); + app.post<{ Params: { clientId: string } }>("/:clientId/disconnect", async (request) => { const { userId } = requireUser(request); const { clientId } = request.params; diff --git a/packages/server/src/services/client.ts b/packages/server/src/services/client.ts index 96acf94cc..ba52efe7c 100644 --- a/packages/server/src/services/client.ts +++ b/packages/server/src/services/client.ts @@ -346,9 +346,20 @@ export async function updateClientCapabilities( if (existingCapabilities.success && stableJson(existingCapabilities.data) === stableJson(parsed.data)) { return; } - const merged = { ...baseMetadata, capabilities: parsed.data }; - await db.update(clients).set({ metadata: merged }).where(eq(clients.id, clientId)); + // Atomic key update so concurrent modelCatalogRpc refs (and other metadata + // writers) are not erased by a whole-object metadata replace. + await db + .update(clients) + .set({ + metadata: sql`jsonb_set( + COALESCE(${clients.metadata}, '{}'::jsonb), + '{capabilities}', + ${JSON.stringify(parsed.data)}::jsonb, + true + )`, + }) + .where(eq(clients.id, clientId)); } /** diff --git a/packages/server/src/services/connection-manager.ts b/packages/server/src/services/connection-manager.ts index f5924a505..7b100923a 100644 --- a/packages/server/src/services/connection-manager.ts +++ b/packages/server/src/services/connection-manager.ts @@ -168,6 +168,7 @@ export function removeClientConnection(clientId: string, ws: WebSocket): string[ activeConnections.delete(agentId); } clientConnections.delete(clientId); + rejectPendingRepliesForClient(clientId, new Error("Client disconnected")); return agentIds; } @@ -222,5 +223,73 @@ export function forceDisconnectClient(clientId: string): string[] { activeConnections.delete(agentId); } clientConnections.delete(clientId); + rejectPendingRepliesForClient(clientId, new Error("Client disconnected")); return agentIds; } + +/** + * HTTP→daemon request/response correlation. Used by host-local discovery + * (provider model catalogs) where the web needs a synchronous answer from the + * connected computer rather than fire-and-forget + poll. + */ +type PendingClientReply = { + clientId: string; + resolve: (value: unknown) => void; + reject: (error: Error) => void; + timer: ReturnType; +}; + +const pendingClientReplies = new Map(); + +export const DEFAULT_CLIENT_REPLY_TIMEOUT_MS = 25_000; + +let clientReplyTimeoutMsForTests: number | null = null; + +/** Test seam: shorten HTTP↔daemon reply waits without changing production default. */ +export function setClientReplyTimeoutMsForTests(timeoutMs: number | null): void { + clientReplyTimeoutMsForTests = timeoutMs; +} + +export function waitForClientReply( + clientId: string, + ref: string, + timeoutMs: number = clientReplyTimeoutMsForTests ?? DEFAULT_CLIENT_REPLY_TIMEOUT_MS, +): Promise { + return new Promise((resolve, reject) => { + if (pendingClientReplies.has(ref)) { + reject(new Error(`Duplicate pending client reply ref: ${ref}`)); + return; + } + const timer = setTimeout(() => { + pendingClientReplies.delete(ref); + reject(new Error("Timed out waiting for the computer to reply")); + }, timeoutMs); + pendingClientReplies.set(ref, { clientId, resolve, reject, timer }); + }); +} + +export function resolveClientReply(clientId: string, ref: string, value: unknown): boolean { + const pending = pendingClientReplies.get(ref); + if (!pending || pending.clientId !== clientId) return false; + clearTimeout(pending.timer); + pendingClientReplies.delete(ref); + pending.resolve(value); + return true; +} + +export function rejectPendingRepliesForClient(clientId: string, error: Error): void { + for (const [ref, pending] of pendingClientReplies) { + if (pending.clientId !== clientId) continue; + clearTimeout(pending.timer); + pendingClientReplies.delete(ref); + pending.reject(error); + } +} + +/** Test seam: drop every pending reply without resolving. */ +export function clearPendingClientRepliesForTests(): void { + for (const pending of pendingClientReplies.values()) { + clearTimeout(pending.timer); + } + pendingClientReplies.clear(); +} diff --git a/packages/server/src/services/notifier.ts b/packages/server/src/services/notifier.ts index 28f244351..d6a744c02 100644 --- a/packages/server/src/services/notifier.ts +++ b/packages/server/src/services/notifier.ts @@ -38,6 +38,18 @@ const CHAT_AUDIENCE_CHANNEL = "chat_audience_events"; */ const CHAT_UPDATED_CHANNEL = "chat_updated_events"; const AGENT_ROUTE_CHANNEL = "agent_route_events"; +/** + * Cross-replica reverse command to a connected daemon (e.g. provider-models:list). + * Payload is small JSON: `{ type, clientId, provider, ref }`. The replica that + * owns the client's WebSocket delivers it via `sendToClient`; others no-op. + */ +const DAEMON_CLIENT_COMMAND_CHANNEL = "daemon_client_commands"; +/** + * Cross-replica wake that a daemon command result is ready. Payload is + * `{ clientId, ref }` only — the catalog lives in `clients.metadata` so large + * Cursor lists stay under the PG NOTIFY 8KB limit. + */ +const DAEMON_CLIENT_COMMAND_RESULT_CHANNEL = "daemon_client_command_results"; /** * A viewer's PRIVATE me-chats projection changed (currently: they pinned or * unpinned a chat). Carries `:` so the WS layer @@ -97,6 +109,24 @@ export type AgentRouteChangePayload = { reason: string; }; export type AgentRouteChangeHandler = (payload: AgentRouteChangePayload) => void; + +/** Small reverse-command frame fan-out for host-local daemon RPCs. */ +export type DaemonClientCommandPayload = { + type: string; + clientId: string; + provider: string; + ref: string; + /** DB-authoritative `clients.instance_id` — only that replica may deliver. */ + targetInstanceId: string; +}; +export type DaemonClientCommandHandler = (payload: DaemonClientCommandPayload) => void; + +/** Wake waiters that a correlated daemon RPC result is stored in client metadata. */ +export type DaemonClientCommandResultPayload = { + clientId: string; + ref: string; +}; +export type DaemonClientCommandResultHandler = (payload: DaemonClientCommandResultPayload) => void; export type MeChatsChangedHandler = (payload: { humanAgentId: string; organizationId: string }) => void; /** @@ -146,6 +176,17 @@ export type Notifier = { notifyMeChatsChanged(humanAgentId: string, organizationId: string): Promise; /** Agent runtime route changed: fan local WS detach/pin handling to every server replica. */ notifyAgentRouteChange(payload: AgentRouteChangePayload): Promise; + /** + * Fan a small reverse-command frame to every replica so the process that + * owns the daemon WebSocket can `sendToClient`. Payload must stay tiny + * (no catalog bodies). + */ + notifyDaemonClientCommand(payload: DaemonClientCommandPayload): Promise; + /** + * Wake waiters that a correlated daemon RPC result is durable in + * `clients.metadata` (catalog bodies are too large for NOTIFY). + */ + notifyDaemonClientCommandResult(payload: DaemonClientCommandResultPayload): Promise; /** * Push a raw JSON frame to every socket currently subscribed to `inboxId` * on **this server instance only**. Unlike `notify`, does not fan out @@ -174,6 +215,10 @@ export type Notifier = { onMeChatsChanged(handler: MeChatsChangedHandler): void; /** Register a handler for agent runtime route changes. */ onAgentRouteChange(handler: AgentRouteChangeHandler): void; + /** Register a handler for cross-replica daemon reverse commands. */ + onDaemonClientCommand(handler: DaemonClientCommandHandler): void; + /** Register a handler for cross-replica daemon RPC result wakes. */ + onDaemonClientCommandResult(handler: DaemonClientCommandResultHandler): void; /** Start listening for PG notifications */ start(): Promise; /** Stop listening */ @@ -192,6 +237,8 @@ export function createNotifier(listenClient: postgres.Sql): Notifier { const chatUpdatedHandlers: ChatUpdatedChangeHandler[] = []; const meChatsChangedHandlers: MeChatsChangedHandler[] = []; const agentRouteHandlers: AgentRouteChangeHandler[] = []; + const daemonClientCommandHandlers: DaemonClientCommandHandler[] = []; + const daemonClientCommandResultHandlers: DaemonClientCommandResultHandler[] = []; let unlistenInboxFn: (() => Promise) | null = null; let unlistenConfigFn: (() => Promise) | null = null; let unlistenSessionStateFn: (() => Promise) | null = null; @@ -203,6 +250,8 @@ export function createNotifier(listenClient: postgres.Sql): Notifier { let unlistenChatUpdatedFn: (() => Promise) | null = null; let unlistenMeChatsChangedFn: (() => Promise) | null = null; let unlistenAgentRouteFn: (() => Promise) | null = null; + let unlistenDaemonClientCommandFn: (() => Promise) | null = null; + let unlistenDaemonClientCommandResultFn: (() => Promise) | null = null; function handleNotification(payload: string) { // payload format: "inboxId:messageId" @@ -344,6 +393,22 @@ export function createNotifier(listenClient: postgres.Sql): Notifier { } }, + async notifyDaemonClientCommand(payload: DaemonClientCommandPayload) { + try { + await listenClient`SELECT pg_notify(${DAEMON_CLIENT_COMMAND_CHANNEL}, ${JSON.stringify(payload)})`; + } catch { + // fire-and-forget — HTTP waiter timeout is the durable fallback. + } + }, + + async notifyDaemonClientCommandResult(payload: DaemonClientCommandResultPayload) { + try { + await listenClient`SELECT pg_notify(${DAEMON_CLIENT_COMMAND_RESULT_CHANNEL}, ${JSON.stringify(payload)})`; + } catch { + // fire-and-forget — HTTP waiter timeout is the durable fallback. + } + }, + async pushFrameToInbox(inboxId: string, frame: string): Promise { const map = subscriptions.get(inboxId); if (!map) return 0; @@ -404,6 +469,14 @@ export function createNotifier(listenClient: postgres.Sql): Notifier { agentRouteHandlers.push(handler); }, + onDaemonClientCommand(handler: DaemonClientCommandHandler) { + daemonClientCommandHandlers.push(handler); + }, + + onDaemonClientCommandResult(handler: DaemonClientCommandResultHandler) { + daemonClientCommandResultHandlers.push(handler); + }, + async start() { const inboxResult = await listenClient.listen(INBOX_CHANNEL, (payload) => { if (payload) handleNotification(payload); @@ -595,6 +668,55 @@ export function createNotifier(listenClient: postgres.Sql): Notifier { } }); unlistenAgentRouteFn = agentRouteResult.unlisten; + + const daemonClientCommandResult = await listenClient.listen(DAEMON_CLIENT_COMMAND_CHANNEL, (payload) => { + if (!payload) return; + try { + const parsed = JSON.parse(payload) as Partial; + if ( + typeof parsed.type !== "string" || + typeof parsed.clientId !== "string" || + typeof parsed.provider !== "string" || + typeof parsed.ref !== "string" || + typeof parsed.targetInstanceId !== "string" + ) { + return; + } + for (const handler of daemonClientCommandHandlers) { + try { + handler(parsed as DaemonClientCommandPayload); + } catch { + // swallow — handler errors must not poison fan-out + } + } + } catch { + // ignore malformed payloads + } + }); + unlistenDaemonClientCommandFn = daemonClientCommandResult.unlisten; + + const daemonClientCommandResultWake = await listenClient.listen( + DAEMON_CLIENT_COMMAND_RESULT_CHANNEL, + (payload) => { + if (!payload) return; + try { + const parsed = JSON.parse(payload) as Partial; + if (typeof parsed.clientId !== "string" || typeof parsed.ref !== "string") { + return; + } + for (const handler of daemonClientCommandResultHandlers) { + try { + handler(parsed as DaemonClientCommandResultPayload); + } catch { + // swallow — handler errors must not poison fan-out + } + } + } catch { + // ignore malformed payloads + } + }, + ); + unlistenDaemonClientCommandResultFn = daemonClientCommandResultWake.unlisten; }, async stop() { @@ -642,6 +764,14 @@ export function createNotifier(listenClient: postgres.Sql): Notifier { await unlistenAgentRouteFn(); unlistenAgentRouteFn = null; } + if (unlistenDaemonClientCommandFn) { + await unlistenDaemonClientCommandFn(); + unlistenDaemonClientCommandFn = null; + } + if (unlistenDaemonClientCommandResultFn) { + await unlistenDaemonClientCommandResultFn(); + unlistenDaemonClientCommandResultFn = null; + } }, }; } diff --git a/packages/server/src/services/provider-models-rpc.ts b/packages/server/src/services/provider-models-rpc.ts new file mode 100644 index 000000000..aa340aa06 --- /dev/null +++ b/packages/server/src/services/provider-models-rpc.ts @@ -0,0 +1,97 @@ +import { type ProviderModelCatalog, providerModelCatalogSchema } from "@first-tree/shared"; +import { eq, sql } from "drizzle-orm"; +import type { Database } from "../db/connection.js"; +import { clients } from "../db/schema/clients.js"; + +/** + * Durable rendezvous for host-local model-catalog RPC. + * + * PG NOTIFY payloads must stay small (≈8KB). Cursor catalogs can exceed that, + * so the socket-owning replica stores the catalog under + * `clients.metadata.modelCatalogRpc[ref]` with an atomic top-level `jsonb_set` + * (sibling keys like `capabilities` stay intact) and a nested `||` merge for + * the ref (concurrent UPDATEs on the same client row serialize under the row + * lock and re-read the latest map). A tiny `{ clientId, ref }` wake fans out + * after the durable write. + */ + +const RPC_METADATA_KEY = "modelCatalogRpc"; +/** Ignore durable entries older than this (logical TTL; keys are not rewritten). */ +const MAX_AGE_MS = 120_000; +const REF_RE = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; + +type RpcEntry = { + catalog: ProviderModelCatalog; + storedAt: string; +}; + +function asRpcEntry(raw: unknown): RpcEntry | null { + if (!raw || typeof raw !== "object" || Array.isArray(raw)) return null; + const row = raw as Record; + const parsed = providerModelCatalogSchema.safeParse(row.catalog); + if (!parsed.success || typeof row.storedAt !== "string") return null; + const storedMs = Date.parse(row.storedAt); + if (!Number.isFinite(storedMs) || Date.now() - storedMs >= MAX_AGE_MS) return null; + return { catalog: parsed.data, storedAt: row.storedAt }; +} + +/** + * Persist one catalog ref without replacing the whole `clients.metadata` object. + * Concurrent refs and sibling metadata writers (capabilities) are preserved. + */ +export async function storeModelCatalogRpcResult( + db: Database, + clientId: string, + ref: string, + catalog: ProviderModelCatalog, +): Promise { + if (!REF_RE.test(ref)) { + throw new Error(`Invalid model-catalog RPC ref: ${ref}`); + } + const entry = { + catalog, + storedAt: new Date().toISOString(), + }; + // Top-level jsonb_set keeps capabilities / lastUpdateAttempt. Nested || merges + // one ref; concurrent UPDATEs on this row serialize and re-evaluate against + // the latest map under READ COMMITTED. + await db + .update(clients) + .set({ + metadata: sql`jsonb_set( + COALESCE(${clients.metadata}, '{}'::jsonb), + '{modelCatalogRpc}', + COALESCE(${clients.metadata} -> 'modelCatalogRpc', '{}'::jsonb) + || jsonb_build_object(${ref}::text, ${JSON.stringify(entry)}::jsonb), + true + )`, + }) + .where(eq(clients.id, clientId)); +} + +/** Load a previously stored catalog when still within the logical TTL. */ +export async function readModelCatalogRpcResult( + db: Database, + clientId: string, + ref: string, +): Promise { + const [client] = await db + .select({ metadata: clients.metadata }) + .from(clients) + .where(eq(clients.id, clientId)) + .limit(1); + if (!client) return null; + const base = (client.metadata ?? {}) as Record; + const map = base[RPC_METADATA_KEY]; + if (!map || typeof map !== "object" || Array.isArray(map)) return null; + const entry = asRpcEntry((map as Record)[ref]); + return entry?.catalog ?? null; +} + +/** + * True when the DB says a daemon WebSocket is live somewhere (this process or + * another replica). Used to decide between cross-replica fan-out and a hard 503. + */ +export function isClientConnectedSomewhere(client: { status: string; instanceId: string | null }): boolean { + return client.status === "connected" && client.instanceId != null; +} diff --git a/packages/shared/src/index.ts b/packages/shared/src/index.ts index 2f012a0d0..4802b4d9a 100644 --- a/packages/shared/src/index.ts +++ b/packages/shared/src/index.ts @@ -886,6 +886,20 @@ export { sessionStateMessageSchema, sessionStateSchema, } from "./schemas/presence.js"; +export { + PROVIDER_MODELS_LIST_TYPE, + PROVIDER_MODELS_RESULT_TYPE, + type ProviderModelCatalog, + type ProviderModelCatalogSource, + type ProviderModelOption, + type ProviderModelsListCommand, + type ProviderModelsResultFrame, + providerModelCatalogSchema, + providerModelCatalogSourceSchema, + providerModelOptionSchema, + providerModelsListCommandSchema, + providerModelsResultFrameSchema, +} from "./schemas/provider-models.js"; export { type AgentStatusReason, agentStatusReasonSchema, diff --git a/packages/shared/src/schemas/provider-models.ts b/packages/shared/src/schemas/provider-models.ts new file mode 100644 index 000000000..c5c8bf244 --- /dev/null +++ b/packages/shared/src/schemas/provider-models.ts @@ -0,0 +1,69 @@ +import { z } from "zod"; +import { runtimeProviderSchema } from "./runtime-provider.js"; + +/** + * Host-local model catalog discovery: the daemon asks the real provider on the + * computer for the models the operator can pick, and the web renders that list. + * First Tree never maintains a global model matrix and never proxies provider + * credentials — discovery runs only on the daemon host. + * + * Wire shape mirrors runtime-auth: + * - Web → Server HTTP `GET /clients/:clientId/providers/:provider/models` + * - Server → daemon reverse command `provider-models:list` (ref-correlated) + * - Daemon → Server reply frame `provider-models:result` + * - Server resolves the pending HTTP request with the catalog + */ + +/** Server→client command: discover models for one runtime provider. */ +export const PROVIDER_MODELS_LIST_TYPE = "provider-models:list" as const; + +/** Client→server reply carrying the discovered catalog (or an unavailable stub). */ +export const PROVIDER_MODELS_RESULT_TYPE = "provider-models:result" as const; + +export const providerModelOptionSchema = z.object({ + /** Exact value written to `config.payload.model`. */ + id: z.string().min(1), + /** Human-facing label; defaults to `id` in the UI when absent. */ + label: z.string().optional(), + /** Secondary hint (e.g. "default", "flagship"). */ + hint: z.string().optional(), + /** Provider-reported default; informational — Web still uses empty string for DEFAULT. */ + isDefault: z.boolean().optional(), +}); +export type ProviderModelOption = z.infer; + +export const providerModelCatalogSourceSchema = z.enum([ + "provider-cli", + "provider-config", + "provider-rpc", + "provider-cache", + "unavailable", +]); +export type ProviderModelCatalogSource = z.infer; + +export const providerModelCatalogSchema = z.object({ + provider: runtimeProviderSchema, + models: z.array(providerModelOptionSchema), + /** Provider's own default model id when known (Kimi `default_model`, Cursor `auto`, …). */ + defaultModelId: z.string().nullable().optional(), + fetchedAt: z.string(), + source: providerModelCatalogSourceSchema, + /** Present when discovery failed or the provider is not supported yet. */ + error: z.string().nullable().optional(), +}); +export type ProviderModelCatalog = z.infer; + +export const providerModelsListCommandSchema = z.object({ + type: z.literal(PROVIDER_MODELS_LIST_TYPE), + provider: runtimeProviderSchema, + /** Correlation id tying command → result → HTTP response. */ + ref: z.string().min(1), +}); +export type ProviderModelsListCommand = z.infer; + +export const providerModelsResultFrameSchema = z.object({ + type: z.literal(PROVIDER_MODELS_RESULT_TYPE), + ref: z.string().min(1), + catalog: providerModelCatalogSchema, +}); +export type ProviderModelsResultFrame = z.infer; diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index fcb9d1774..7e2fffb35 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -126,6 +126,9 @@ importers: semver: specifier: ^7.6.3 version: 7.7.4 + smol-toml: + specifier: ^1.7.0 + version: 1.7.0 ws: specifier: ^8.20.0 version: 8.20.0 @@ -7657,6 +7660,22 @@ snapshots: chai: 5.3.3 tinyrainbow: 2.0.0 + '@vitest/mocker@3.2.4(vite@6.4.1(@types/node@22.19.15)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.2))': + dependencies: + '@vitest/spy': 3.2.4 + estree-walker: 3.0.3 + magic-string: 0.30.21 + optionalDependencies: + vite: 6.4.1(@types/node@22.19.15)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.2) + + '@vitest/mocker@3.2.4(vite@6.4.1(@types/node@22.19.15)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.3))': + dependencies: + '@vitest/spy': 3.2.4 + estree-walker: 3.0.3 + magic-string: 0.30.21 + optionalDependencies: + vite: 6.4.1(@types/node@22.19.15)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.3) + '@vitest/mocker@3.2.4(vite@6.4.1(@types/node@22.19.19)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.3))': dependencies: '@vitest/spy': 3.2.4 @@ -10323,7 +10342,7 @@ snapshots: dependencies: '@types/chai': 5.2.3 '@vitest/expect': 3.2.4 - '@vitest/mocker': 3.2.4(vite@6.4.1(@types/node@22.19.19)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.3)) + '@vitest/mocker': 3.2.4(vite@6.4.1(@types/node@22.19.15)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.2)) '@vitest/pretty-format': 3.2.4 '@vitest/runner': 3.2.4 '@vitest/snapshot': 3.2.4 @@ -10366,7 +10385,7 @@ snapshots: dependencies: '@types/chai': 5.2.3 '@vitest/expect': 3.2.4 - '@vitest/mocker': 3.2.4(vite@6.4.1(@types/node@22.19.19)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.3)) + '@vitest/mocker': 3.2.4(vite@6.4.1(@types/node@22.19.15)(jiti@2.6.1)(lightningcss@1.32.0)(tsx@4.21.0)(yaml@2.8.3)) '@vitest/pretty-format': 3.2.4 '@vitest/runner': 3.2.4 '@vitest/snapshot': 3.2.4