From bd9e09db441f4c4dcad2f8a8446c8818303c7134 Mon Sep 17 00:00:00 2001 From: Mario Zechner Date: Wed, 15 Jul 2026 12:56:50 +0200 Subject: [PATCH] feat(coding-agent): expose dynamic provider refresh --- packages/ai/src/models-store.ts | 31 ++++++++++--------- packages/ai/src/models.ts | 6 ++-- packages/ai/src/providers/radius.ts | 6 ++-- packages/ai/test/providers.test.ts | 4 +-- packages/coding-agent/docs/extensions.md | 23 +++++++++++++- .../coding-agent/src/core/extensions/types.ts | 6 ++++ .../coding-agent/src/core/models-store.ts | 24 +++++++------- .../src/core/provider-composer.ts | 24 ++++++++++++-- .../src/core/remote-catalog-provider.ts | 26 +++++++++++++--- ...model-runtime-modify-models-compat.test.ts | 30 ++++++++++++++++-- .../coding-agent/test/models-store.test.ts | 11 ++++--- packages/coding-agent/test/radius.test.ts | 2 +- .../test/remote-catalog-provider.test.ts | 26 ++++++++++++---- 13 files changed, 164 insertions(+), 55 deletions(-) diff --git a/packages/ai/src/models-store.ts b/packages/ai/src/models-store.ts index 40670c5c..020edb1f 100644 --- a/packages/ai/src/models-store.ts +++ b/packages/ai/src/models-store.ts @@ -1,35 +1,38 @@ import type { Api, Model } from "./types.ts"; +export interface ModelsStoreEntry { + models: readonly Model[]; + /** Unix timestamp of the last completed remote check. */ + checkedAt?: number; +} + /** Persistent model catalogs keyed by provider ID. */ export interface ModelsStore { - read(providerId: string): Promise[] | undefined>; - write(providerId: string, models: readonly Model[]): Promise; + read(providerId: string): Promise; + write(providerId: string, entry: ModelsStoreEntry): Promise; delete(providerId: string): Promise; } /** ModelsStore scoped to one provider. Providers cannot access other providers' catalogs. */ export interface ProviderModelsStore { - read(): Promise[] | undefined>; - write(models: readonly Model[]): Promise; + read(): Promise; + write(entry: ModelsStoreEntry): Promise; delete(): Promise; } export class InMemoryModelsStore implements ModelsStore { - private readonly models = new Map[]>(); + private readonly entries = new Map(); - async read(providerId: string): Promise[] | undefined> { - const models = this.models.get(providerId); - return models?.map((model) => structuredClone(model)); + async read(providerId: string): Promise { + const entry = this.entries.get(providerId); + return entry ? structuredClone(entry) : undefined; } - async write(providerId: string, models: readonly Model[]): Promise { - this.models.set( - providerId, - models.map((model) => structuredClone(model)), - ); + async write(providerId: string, entry: ModelsStoreEntry): Promise { + this.entries.set(providerId, structuredClone(entry)); } async delete(providerId: string): Promise { - this.models.delete(providerId); + this.entries.delete(providerId); } } diff --git a/packages/ai/src/models.ts b/packages/ai/src/models.ts index 2acab681..7938ef12 100644 --- a/packages/ai/src/models.ts +++ b/packages/ai/src/models.ts @@ -282,7 +282,7 @@ class ModelsImpl implements MutableModels { if (options.signal?.aborted) return; const store: ProviderModelsStore = { read: () => this.modelsStore.read(provider.id), - write: (models) => this.modelsStore.write(provider.id, models), + write: (entry) => this.modelsStore.write(provider.id, entry), delete: () => this.modelsStore.delete(provider.id), }; let stored: Credential | undefined; @@ -589,7 +589,7 @@ export function createProvider(input: CreateProviderOpti try { const stored = await context.store.read(); if (stored) { - dynamicModels = stored + dynamicModels = stored.models .filter((model) => model.provider === input.id) .map((model) => model as Model); } @@ -597,7 +597,7 @@ export function createProvider(input: CreateProviderOpti const refreshed = await fetchModels(context); if (context.signal?.aborted) return; dynamicModels = refreshed; - await context.store.write(refreshed); + await context.store.write({ models: refreshed, checkedAt: Date.now() }); } finally { inflightRefresh = undefined; } diff --git a/packages/ai/src/providers/radius.ts b/packages/ai/src/providers/radius.ts index 19afd8d8..1321cd18 100644 --- a/packages/ai/src/providers/radius.ts +++ b/packages/ai/src/providers/radius.ts @@ -37,14 +37,14 @@ export function radiusProvider(options: RadiusProviderOptions = {}): Provider<"p inflightRefresh ??= (async () => { try { const stored = await context.store.read(); - if (stored) models = stored.filter((model) => model.provider === id) as typeof models; + if (stored) models = stored.models.filter((model) => model.provider === id) as typeof models; // Import catalogs cached by the pre-ModelsStore Radius implementation. if (!stored && context.credential?.type === "oauth") { const legacy = getRadiusModels(id, context.credential); if (legacy.length > 0) { models = legacy; - await context.store.write(legacy); + await context.store.write({ models: legacy, checkedAt: Date.now() }); } } @@ -54,7 +54,7 @@ export function radiusProvider(options: RadiusProviderOptions = {}): Provider<"p const config = await loadRadiusGatewayConfig(gateway, apiKey, context.signal); if (context.signal?.aborted) return; models = getRadiusModelsFromConfig(id, config); - await context.store.write(models); + await context.store.write({ models, checkedAt: Date.now() }); } finally { inflightRefresh = undefined; } diff --git a/packages/ai/test/providers.test.ts b/packages/ai/test/providers.test.ts index bbccd5a3..37c359c5 100644 --- a/packages/ai/test/providers.test.ts +++ b/packages/ai/test/providers.test.ts @@ -2,7 +2,7 @@ import { describe, expect, it } from "vitest"; import { envApiKeyAuth } from "../src/auth/helpers.ts"; import type { AuthContext, AuthEvent } from "../src/auth/types.ts"; import { createModels, createProvider } from "../src/models.ts"; -import { InMemoryModelsStore } from "../src/models-store.ts"; +import { InMemoryModelsStore, type ModelsStoreEntry } from "../src/models-store.ts"; import { builtinModels, builtinProviders } from "../src/providers/all.ts"; import { amazonBedrockProvider } from "../src/providers/amazon-bedrock.ts"; import { anthropicProvider } from "../src/providers/anthropic.ts"; @@ -360,7 +360,7 @@ describe("createProvider", () => { credential: { type: "api_key" as const }, store: { read: () => store.read("dynamic"), - write: (listed: readonly Model[]) => store.write("dynamic", listed), + write: (entry: ModelsStoreEntry) => store.write("dynamic", entry), delete: () => store.delete("dynamic"), }, allowNetwork: true, diff --git a/packages/coding-agent/docs/extensions.md b/packages/coding-agent/docs/extensions.md index c74eecbc..ec1ea00a 100644 --- a/packages/coding-agent/docs/extensions.md +++ b/packages/coding-agent/docs/extensions.md @@ -1677,7 +1677,7 @@ Register or override a model provider dynamically. Useful for proxies, custom en Calls made during the extension factory function are queued and applied once the runner initialises. Calls made after that — for example from a command handler following a user setup flow — take effect immediately without requiring a `/reload`. -If you need to discover models from a remote endpoint, prefer an async extension factory over deferring the fetch to `session_start`. pi waits for the factory before startup continues, so the registered models are available immediately, including to `pi --list-models`. +Dynamic providers can implement `refreshModels`. Pi calls it during model refresh, publishes the returned list synchronously through the provider, and passes the canonical credential/store/network/signal context. The extension decides whether to persist the catalog through `context.store`; live servers such as llama.cpp can ignore it. ```typescript // Register a new provider with custom models @@ -1699,6 +1699,26 @@ pi.registerProvider("my-proxy", { ] }); +// Register a live llama.cpp catalog without persisting discovered models +pi.registerProvider("llama.cpp", { + baseUrl: "http://localhost:8080/v1", + apiKey: "local", + api: "openai-completions", + async refreshModels({ signal }) { + const response = await fetch("http://localhost:8080/v1/models", { signal }); + const { data } = await response.json(); + return data.map(({ id }) => ({ + id, + name: id, + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128000, + maxTokens: 16384 + })); + } +}); + // Override baseUrl for an existing provider (keeps all models) pi.registerProvider("anthropic", { baseUrl: "https://proxy.example.com" @@ -1736,6 +1756,7 @@ pi.registerProvider("corporate-ai", { - `headers` - Custom headers to include in requests. - `authHeader` - If true, adds `Authorization: Bearer` header automatically. - `models` - Array of model definitions. If provided, replaces all existing models for this provider. Model definitions can set `baseUrl` to override the provider endpoint for that model. +- `refreshModels` - Async dynamic discovery callback. Its returned models replace extension-provided models. Use the scoped `context.store` only when results should persist. - `oauth` - OAuth provider config for `/login` support. When provided, the provider appears in the login menu. - `streamSimple` - Custom streaming implementation for non-standard APIs. diff --git a/packages/coding-agent/src/core/extensions/types.ts b/packages/coding-agent/src/core/extensions/types.ts index 2f41c54c..9ce96938 100644 --- a/packages/coding-agent/src/core/extensions/types.ts +++ b/packages/coding-agent/src/core/extensions/types.ts @@ -25,6 +25,7 @@ import type { OAuthCredentials, OAuthLoginCallbacks, ProviderHeaders, + RefreshModelsContext, SimpleStreamOptions, TextContent, ToolResultMessage, @@ -1420,6 +1421,11 @@ export interface ProviderConfig { authHeader?: boolean; /** Models to register. If provided, replaces all existing models for this provider. */ models?: ProviderModelConfig[]; + /** + * Refresh this provider's model list. The returned list replaces extension-provided models. + * Use context.store explicitly when the catalog should persist across sessions. + */ + refreshModels?(context: RefreshModelsContext): Promise; /** OAuth provider for /login support. The `id` is set automatically from the provider name. */ oauth?: { /** Display name for the provider in login UI. */ diff --git a/packages/coding-agent/src/core/models-store.ts b/packages/coding-agent/src/core/models-store.ts index 83edbca0..dca57995 100644 --- a/packages/coding-agent/src/core/models-store.ts +++ b/packages/coding-agent/src/core/models-store.ts @@ -1,23 +1,23 @@ import { join } from "node:path"; -import type { Api, Model, ModelsStore } from "@earendil-works/pi-ai"; +import type { ModelsStore, ModelsStoreEntry } from "@earendil-works/pi-ai"; import { getAgentDir } from "../config.ts"; import { type AuthStorageBackend, FileAuthStorageBackend } from "./auth-storage.ts"; -type StoredModels = Record[]>; +type StoredModels = Record; export class InMemoryCodingAgentModelsStore implements ModelsStore { - private readonly models = new Map[]>(); + private readonly entries = new Map(); - async read(providerId: string): Promise[] | undefined> { - return this.models.get(providerId); + async read(providerId: string): Promise { + return this.entries.get(providerId); } - async write(providerId: string, models: readonly Model[]): Promise { - this.models.set(providerId, models); + async write(providerId: string, entry: ModelsStoreEntry): Promise { + this.entries.set(providerId, entry); } async delete(providerId: string): Promise { - this.models.delete(providerId); + this.entries.delete(providerId); } } @@ -33,16 +33,16 @@ export class FileModelsStore implements ModelsStore { return content ? (JSON.parse(content) as StoredModels) : {}; } - async read(providerId: string): Promise[] | undefined> { + async read(providerId: string): Promise { return this.storage.withLock((content) => ({ - result: this.parse(content)[providerId]?.map((model) => structuredClone(model)), + result: structuredClone(this.parse(content)[providerId]), })); } - async write(providerId: string, models: readonly Model[]): Promise { + async write(providerId: string, entry: ModelsStoreEntry): Promise { await this.storage.withLockAsync(async (content) => { const current = this.parse(content); - current[providerId] = models.map((model) => structuredClone(model)); + current[providerId] = structuredClone(entry); return { result: undefined, next: JSON.stringify(current, null, 2) }; }); } diff --git a/packages/coding-agent/src/core/provider-composer.ts b/packages/coding-agent/src/core/provider-composer.ts index 28a77f6e..a29c3762 100644 --- a/packages/coding-agent/src/core/provider-composer.ts +++ b/packages/coding-agent/src/core/provider-composer.ts @@ -15,6 +15,7 @@ import { type OAuthLoginCallbacks, type Provider, type ProviderHeaders, + type RefreshModelsContext, type SimpleStreamOptions, type StreamOptions, } from "@earendil-works/pi-ai"; @@ -63,6 +64,7 @@ export interface ProviderConfigInput { headers?: Record; compat?: Model["compat"]; }>; + refreshModels?(context: RefreshModelsContext): Promise>; } export type AuthStatus = { @@ -415,10 +417,17 @@ export function composeModelProvider( ): Provider { const config = modelConfig.getProvider(providerId); let extensionOAuthCredential: OAuthCredentials | undefined; + let refreshedExtensionModels: ProviderConfigInput["models"]; + const currentExtension = (): ProviderConfigInput | undefined => + extension && refreshedExtensionModels ? { ...extension, models: refreshedExtensionModels } : extension; // models.json modelOverrides are the topmost user-config layer: they apply once, // after custom-model upserts, extension model replacement, and legacy OAuth projection. const getModels = () => { - let models = applyExtension(providerId, applyModelsJson(providerId, base?.getModels() ?? [], config), extension); + let models = applyExtension( + providerId, + applyModelsJson(providerId, base?.getModels() ?? [], config), + currentExtension(), + ); if (extensionOAuthCredential && extension?.oauth?.modifyModels) { models = extension.oauth.modifyModels(models, extensionOAuthCredential); } @@ -464,9 +473,20 @@ export function composeModelProvider( auth: { ...(apiKey ? { apiKey } : {}), ...(oauth ? { oauth } : {}) }, getModels, refreshModels: - base?.refreshModels || extension?.oauth?.modifyModels + base?.refreshModels || extension?.refreshModels || extension?.oauth?.modifyModels ? async (context) => { await base?.refreshModels?.(context); + if (extension?.refreshModels) { + const refreshed = await extension.refreshModels(context); + if (!context.signal?.aborted) { + // Validate before publishing the new synchronous list. + applyExtension(providerId, applyModelsJson(providerId, base?.getModels() ?? [], config), { + ...extension, + models: refreshed, + }); + refreshedExtensionModels = refreshed; + } + } extensionOAuthCredential = context.credential?.type === "oauth" ? context.credential : undefined; } : undefined, diff --git a/packages/coding-agent/src/core/remote-catalog-provider.ts b/packages/coding-agent/src/core/remote-catalog-provider.ts index b9a46d3f..411f2122 100644 --- a/packages/coding-agent/src/core/remote-catalog-provider.ts +++ b/packages/coding-agent/src/core/remote-catalog-provider.ts @@ -1,6 +1,9 @@ import type { Api, Model, Provider } from "@earendil-works/pi-ai"; +import { VERSION } from "../config.ts"; +import { getPiUserAgent } from "../utils/pi-user-agent.ts"; const DEFAULT_CATALOG_BASE_URL = "https://pi.dev"; +export const REMOTE_CATALOG_REFRESH_INTERVAL_MS = 4 * 60 * 60 * 1000; function mergeModels(baseline: readonly Model[], dynamic: readonly Model[]): Model[] { const merged = [...baseline]; @@ -38,22 +41,37 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D inflightRefresh ??= (async () => { try { const stored = await context.store.read(); - if (stored) dynamicModels = stored.filter((model) => model.provider === provider.id); + if (stored) dynamicModels = stored.models.filter((model) => model.provider === provider.id); if (!context.allowNetwork || context.signal?.aborted) return; + if ( + stored?.checkedAt !== undefined && + Date.now() - stored.checkedAt < REMOTE_CATALOG_REFRESH_INTERVAL_MS + ) { + return; + } const url = new URL(`/api/models/providers/${encodeURIComponent(provider.id)}`, catalogBaseUrl); const response = await fetch(url, { - headers: { accept: "application/json" }, + headers: { + accept: "application/json", + "User-Agent": getPiUserAgent(VERSION), + }, signal: context.signal, }); - if (response.status === 404 || response.status === 501) return; + if (context.signal?.aborted) return; + const checkedAt = Date.now(); + if (response.status === 404 || response.status === 501) { + await context.store.write({ models: dynamicModels, checkedAt }); + return; + } if (!response.ok) { + await context.store.write({ models: dynamicModels, checkedAt }); throw new Error(`Model catalog request failed for ${provider.id}: ${response.status}`); } const refreshed = parseCatalog(provider.id, await response.json()); if (context.signal?.aborted) return; dynamicModels = refreshed; - await context.store.write(refreshed); + await context.store.write({ models: refreshed, checkedAt }); } finally { inflightRefresh = undefined; } diff --git a/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts b/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts index f4f1fc3a..f8555a67 100644 --- a/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts +++ b/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts @@ -18,8 +18,34 @@ function model(id: string): Model<"openai-completions"> { }; } -describe("legacy extension OAuth modifyModels", () => { - it("applies the synchronous projection after async credential initialization", async () => { +describe("extension provider model lifecycle", () => { + it("publishes refreshModels results without forcing ModelsStore persistence", async () => { + const modelsStore = new InMemoryModelsStore(); + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory(), + modelsStore, + modelsPath: null, + allowModelNetwork: false, + }); + runtime.registerProvider("extension-dynamic", { + baseUrl: "http://localhost:8080/v1", + apiKey: "local", + api: "openai-completions", + refreshModels: async () => [ + { + ...model("live"), + provider: "extension-dynamic", + baseUrl: "http://localhost:8080/v1", + }, + ], + }); + + await runtime.refresh({ allowNetwork: false }); + expect(runtime.getModel("extension-dynamic", "live")).toBeDefined(); + expect(await modelsStore.read("extension-dynamic")).toBeUndefined(); + }); + + it("applies legacy OAuth modifyModels after async credential initialization", async () => { const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory({ "extension-oauth": { diff --git a/packages/coding-agent/test/models-store.test.ts b/packages/coding-agent/test/models-store.test.ts index 351b7ac0..48315397 100644 --- a/packages/coding-agent/test/models-store.test.ts +++ b/packages/coding-agent/test/models-store.test.ts @@ -36,15 +36,16 @@ describe("FileModelsStore", () => { const path = join(dir, "models-store.json"); const store = new FileModelsStore(path); - await store.write("one", [model("one", "m1")]); - await store.write("two", [model("two", "m2")]); + await store.write("one", { models: [model("one", "m1")], checkedAt: 100 }); + await store.write("two", { models: [model("two", "m2")], checkedAt: 200 }); const reloaded = new FileModelsStore(path); - expect((await reloaded.read("one"))?.map((entry) => entry.id)).toEqual(["m1"]); - expect((await reloaded.read("two"))?.map((entry) => entry.id)).toEqual(["m2"]); + expect((await reloaded.read("one"))?.models.map((entry) => entry.id)).toEqual(["m1"]); + expect((await reloaded.read("one"))?.checkedAt).toBe(100); + expect((await reloaded.read("two"))?.models.map((entry) => entry.id)).toEqual(["m2"]); await reloaded.delete("one"); expect(await reloaded.read("one")).toBeUndefined(); - expect((await reloaded.read("two"))?.map((entry) => entry.id)).toEqual(["m2"]); + expect((await reloaded.read("two"))?.models.map((entry) => entry.id)).toEqual(["m2"]); }); }); diff --git a/packages/coding-agent/test/radius.test.ts b/packages/coding-agent/test/radius.test.ts index 16d16088..de383672 100644 --- a/packages/coding-agent/test/radius.test.ts +++ b/packages/coding-agent/test/radius.test.ts @@ -87,7 +87,7 @@ describe("Radius provider", () => { }); expect(runtime.getModel(RADIUS_PROVIDER_ID, "auto")).toBeDefined(); - expect(await modelsStore.read(RADIUS_PROVIDER_ID)).toHaveLength(1); + expect((await modelsStore.read(RADIUS_PROVIDER_ID))?.models).toHaveLength(1); expect(vi.mocked(fetch).mock.calls[0]?.[1]?.headers).toMatchObject({ authorization: "Bearer access-token" }); }); diff --git a/packages/coding-agent/test/remote-catalog-provider.test.ts b/packages/coding-agent/test/remote-catalog-provider.test.ts index ff4f729b..cf4d68fa 100644 --- a/packages/coding-agent/test/remote-catalog-provider.test.ts +++ b/packages/coding-agent/test/remote-catalog-provider.test.ts @@ -1,5 +1,6 @@ import { createProvider, InMemoryModelsStore, type Model } from "@earendil-works/pi-ai"; import { afterEach, describe, expect, it, vi } from "vitest"; +import { VERSION } from "../src/config.ts"; import { withRemoteCatalog } from "../src/core/remote-catalog-provider.ts"; function model(id: string): Model<"openai-completions"> { @@ -20,8 +21,8 @@ function model(id: string): Model<"openai-completions"> { afterEach(() => vi.restoreAllMocks()); describe("remote catalog provider", () => { - it("parses catalogs keyed by model ID", async () => { - vi.spyOn(globalThis, "fetch").mockResolvedValue( + it("parses keyed catalogs, sends version headers, and observes the refresh TTL", async () => { + const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue( new Response(JSON.stringify({ dynamic: model("dynamic") }), { status: 200, headers: { "content-type": "application/json" }, @@ -47,14 +48,27 @@ describe("remote catalog provider", () => { credential: { type: "api_key" }, store: { read: () => store.read(provider.id), - write: (models) => store.write(provider.id, models), + write: (entry) => store.write(provider.id, entry), + delete: () => store.delete(provider.id), + }, + allowNetwork: true, + }); + await provider.refreshModels?.({ + credential: { type: "api_key" }, + store: { + read: () => store.read(provider.id), + write: (entry) => store.write(provider.id, entry), delete: () => store.delete(provider.id), }, allowNetwork: true, }); expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]); - expect((await store.read(provider.id))?.map((entry) => entry.id)).toEqual(["dynamic"]); + expect((await store.read(provider.id))?.models.map((entry) => entry.id)).toEqual(["dynamic"]); + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(fetchSpy.mock.calls[0]?.[1]?.headers).toMatchObject({ + "User-Agent": expect.stringContaining(`pi/${VERSION}`), + }); }); it("treats unimplemented pi.dev catalog routes as an unavailable overlay", async () => { @@ -81,13 +95,13 @@ describe("remote catalog provider", () => { credential: { type: "api_key" }, store: { read: () => store.read(provider.id), - write: (models) => store.write(provider.id, models), + write: (entry) => store.write(provider.id, entry), delete: () => store.delete(provider.id), }, allowNetwork: true, }), ).resolves.toBeUndefined(); expect(provider.getModels().map((entry) => entry.id)).toEqual(["static"]); - expect(await store.read(provider.id)).toBeUndefined(); + expect(await store.read(provider.id)).toMatchObject({ models: [], checkedAt: expect.any(Number) }); }); });