feat(coding-agent): expose dynamic provider refresh

This commit is contained in:
Mario Zechner
2026-07-15 12:56:50 +02:00
parent 45203abfa0
commit bd9e09db44
13 changed files with 164 additions and 55 deletions
+17 -14
View File
@@ -1,35 +1,38 @@
import type { Api, Model } from "./types.ts";
export interface ModelsStoreEntry {
models: readonly Model<Api>[];
/** Unix timestamp of the last completed remote check. */
checkedAt?: number;
}
/** Persistent model catalogs keyed by provider ID. */
export interface ModelsStore {
read(providerId: string): Promise<readonly Model<Api>[] | undefined>;
write(providerId: string, models: readonly Model<Api>[]): Promise<void>;
read(providerId: string): Promise<ModelsStoreEntry | undefined>;
write(providerId: string, entry: ModelsStoreEntry): Promise<void>;
delete(providerId: string): Promise<void>;
}
/** ModelsStore scoped to one provider. Providers cannot access other providers' catalogs. */
export interface ProviderModelsStore {
read(): Promise<readonly Model<Api>[] | undefined>;
write(models: readonly Model<Api>[]): Promise<void>;
read(): Promise<ModelsStoreEntry | undefined>;
write(entry: ModelsStoreEntry): Promise<void>;
delete(): Promise<void>;
}
export class InMemoryModelsStore implements ModelsStore {
private readonly models = new Map<string, readonly Model<Api>[]>();
private readonly entries = new Map<string, ModelsStoreEntry>();
async read(providerId: string): Promise<readonly Model<Api>[] | undefined> {
const models = this.models.get(providerId);
return models?.map((model) => structuredClone(model));
async read(providerId: string): Promise<ModelsStoreEntry | undefined> {
const entry = this.entries.get(providerId);
return entry ? structuredClone(entry) : undefined;
}
async write(providerId: string, models: readonly Model<Api>[]): Promise<void> {
this.models.set(
providerId,
models.map((model) => structuredClone(model)),
);
async write(providerId: string, entry: ModelsStoreEntry): Promise<void> {
this.entries.set(providerId, structuredClone(entry));
}
async delete(providerId: string): Promise<void> {
this.models.delete(providerId);
this.entries.delete(providerId);
}
}
+3 -3
View File
@@ -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<TApi extends Api = Api>(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<TApi>);
}
@@ -597,7 +597,7 @@ export function createProvider<TApi extends Api = Api>(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;
}
+3 -3
View File
@@ -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;
}
+2 -2
View File
@@ -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<Api>[]) => store.write("dynamic", listed),
write: (entry: ModelsStoreEntry) => store.write("dynamic", entry),
delete: () => store.delete("dynamic"),
},
allowNetwork: true,
+22 -1
View File
@@ -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.
@@ -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<ProviderModelConfig[]>;
/** OAuth provider for /login support. The `id` is set automatically from the provider name. */
oauth?: {
/** Display name for the provider in login UI. */
+12 -12
View File
@@ -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<string, Model<Api>[]>;
type StoredModels = Record<string, ModelsStoreEntry>;
export class InMemoryCodingAgentModelsStore implements ModelsStore {
private readonly models = new Map<string, readonly Model<Api>[]>();
private readonly entries = new Map<string, ModelsStoreEntry>();
async read(providerId: string): Promise<readonly Model<Api>[] | undefined> {
return this.models.get(providerId);
async read(providerId: string): Promise<ModelsStoreEntry | undefined> {
return this.entries.get(providerId);
}
async write(providerId: string, models: readonly Model<Api>[]): Promise<void> {
this.models.set(providerId, models);
async write(providerId: string, entry: ModelsStoreEntry): Promise<void> {
this.entries.set(providerId, entry);
}
async delete(providerId: string): Promise<void> {
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<readonly Model<Api>[] | undefined> {
async read(providerId: string): Promise<ModelsStoreEntry | undefined> {
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<Api>[]): Promise<void> {
async write(providerId: string, entry: ModelsStoreEntry): Promise<void> {
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) };
});
}
@@ -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<string, string>;
compat?: Model<Api>["compat"];
}>;
refreshModels?(context: RefreshModelsContext): Promise<NonNullable<ProviderConfigInput["models"]>>;
}
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,
@@ -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<Api>[], dynamic: readonly Model<Api>[]): Model<Api>[] {
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;
}
@@ -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": {
@@ -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"]);
});
});
+1 -1
View File
@@ -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" });
});
@@ -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) });
});
});