diff --git a/packages/ai/src/models-store.ts b/packages/ai/src/models-store.ts index 020edb1f..3bb2b57e 100644 --- a/packages/ai/src/models-store.ts +++ b/packages/ai/src/models-store.ts @@ -2,6 +2,8 @@ import type { Api, Model } from "./types.ts"; export interface ModelsStoreEntry { models: readonly Model[]; + /** Unix timestamp from the remote catalog's Last-Modified header. */ + lastModified?: number; /** Unix timestamp of the last completed remote check. */ checkedAt?: number; } diff --git a/packages/ai/src/providers/all.ts b/packages/ai/src/providers/all.ts index fd7bea58..549880e3 100644 --- a/packages/ai/src/providers/all.ts +++ b/packages/ai/src/providers/all.ts @@ -67,6 +67,11 @@ export function getBuiltinProviders(): BuiltinProvider[] { return Object.keys(MODELS) as BuiltinProvider[]; } +/** URL of a generated provider catalog, used to compare its mtime with remote catalogs during development. */ +export function getBuiltinModelDataUrl(provider: BuiltinProvider): URL { + return new URL(`./data/${provider}.json`, import.meta.url); +} + export function getBuiltinModels( provider: TProvider, ): Model>[] { diff --git a/packages/coding-agent/src/core/model-runtime.ts b/packages/coding-agent/src/core/model-runtime.ts index 64501f73..2cd85b8d 100644 --- a/packages/coding-agent/src/core/model-runtime.ts +++ b/packages/coding-agent/src/core/model-runtime.ts @@ -143,7 +143,15 @@ export class ModelRuntime implements Models { const providers = builtinProviderCatalog .builtinProviders() .map((provider) => - provider.id === "radius" ? provider : withRemoteCatalog(provider, options.catalogBaseUrl), + provider.id === "radius" + ? provider + : withRemoteCatalog( + provider, + options.catalogBaseUrl, + builtinProviderCatalog.getBuiltinModelDataUrl( + provider.id as builtinProviderCatalog.BuiltinProvider, + ), + ), ); const runtime = new ModelRuntime( credentials, diff --git a/packages/coding-agent/src/core/remote-catalog-provider.ts b/packages/coding-agent/src/core/remote-catalog-provider.ts index 6667916d..b823cf7d 100644 --- a/packages/coding-agent/src/core/remote-catalog-provider.ts +++ b/packages/coding-agent/src/core/remote-catalog-provider.ts @@ -1,4 +1,5 @@ -import type { Api, Model, Provider } from "@earendil-works/pi-ai"; +import { stat } from "node:fs/promises"; +import type { Api, Model, ModelsStoreEntry, Provider } from "@earendil-works/pi-ai"; import { VERSION } from "../config.ts"; import { getPiUserAgent } from "../utils/pi-user-agent.ts"; @@ -29,8 +30,26 @@ function parseCatalog(providerId: string, value: unknown): Model[] { .map((model) => ({ ...model, provider: providerId })); } +function remoteModels( + entry: ModelsStoreEntry | undefined, + localLastModified: number | undefined, +): readonly Model[] { + if (!entry) return []; + if ( + localLastModified !== undefined && + (entry.lastModified === undefined || entry.lastModified <= localLastModified) + ) { + return []; + } + return entry.models; +} + /** Add a persisted pi.dev catalog overlay to a static built-in provider. */ -export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = DEFAULT_CATALOG_BASE_URL): Provider { +export function withRemoteCatalog( + provider: Provider, + catalogBaseUrl: string = DEFAULT_CATALOG_BASE_URL, + localCatalogUrl?: URL, +): Provider { let dynamicModels: readonly Model[] = []; let inflightRefresh: Promise | undefined; @@ -40,12 +59,21 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D refreshModels: (context) => { inflightRefresh ??= (async () => { try { + const localLastModified = localCatalogUrl + ? await stat(localCatalogUrl).then( + (value) => value.mtimeMs, + () => undefined, + ) + : undefined; const stored = await context.store.read(); - if (stored) dynamicModels = stored.models.filter((model) => model.provider === provider.id); + dynamicModels = remoteModels(stored, localLastModified).filter( + (model) => model.provider === provider.id, + ); if (!context.allowNetwork || context.signal?.aborted) return; if ( !context.force && stored?.checkedAt !== undefined && + stored.lastModified !== undefined && Date.now() - stored.checkedAt < REMOTE_CATALOG_REFRESH_INTERVAL_MS ) { return; @@ -62,17 +90,23 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D if (context.signal?.aborted) return; const checkedAt = Date.now(); if (response.status === 404 || response.status === 501) { - await context.store.write({ models: dynamicModels, checkedAt }); + await context.store.write({ ...(stored ?? { models: [] }), checkedAt, lastModified: 0 }); return; } if (!response.ok) { - await context.store.write({ models: dynamicModels, checkedAt }); + await context.store.write({ ...(stored ?? { models: [] }), checkedAt }); throw new Error(`Model catalog request failed for ${provider.id}: ${response.status}`); } const refreshed = parseCatalog(provider.id, await response.json()); + const lastModified = Date.parse(response.headers.get("last-modified") ?? ""); if (context.signal?.aborted) return; - dynamicModels = refreshed; - await context.store.write({ models: refreshed, checkedAt }); + const entry = { + models: refreshed, + checkedAt, + lastModified: Number.isNaN(lastModified) ? 0 : lastModified, + }; + dynamicModels = remoteModels(entry, localLastModified); + await context.store.write(entry); } finally { inflightRefresh = undefined; } diff --git a/packages/coding-agent/test/remote-catalog-provider.test.ts b/packages/coding-agent/test/remote-catalog-provider.test.ts index c8c5f869..8b43a4f3 100644 --- a/packages/coding-agent/test/remote-catalog-provider.test.ts +++ b/packages/coding-agent/test/remote-catalog-provider.test.ts @@ -1,4 +1,11 @@ -import { createProvider, InMemoryModelsStore, type Model } from "@earendil-works/pi-ai"; +import { statSync } from "node:fs"; +import { + createProvider, + InMemoryModelsStore, + type Model, + type ModelsStoreEntry, + type ProviderModelsStore, +} 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"; @@ -18,6 +25,34 @@ function model(id: string): Model<"openai-completions"> { }; } +function testProvider(localCatalogUrl?: URL) { + return withRemoteCatalog( + createProvider({ + id: "test-provider", + auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, + models: [model("static")], + api: { + stream: () => { + throw new Error("not used"); + }, + streamSimple: () => { + throw new Error("not used"); + }, + }, + }), + "https://pi.dev", + localCatalogUrl, + ); +} + +function scopedStore(store: InMemoryModelsStore): ProviderModelsStore { + return { + read: () => store.read("test-provider"), + write: (entry: ModelsStoreEntry) => store.write("test-provider", entry), + delete: () => store.delete("test-provider"), + }; +} + afterEach(() => vi.restoreAllMocks()); describe("remote catalog provider", () => { @@ -29,50 +64,12 @@ describe("remote catalog provider", () => { headers: { "content-type": "application/json" }, }), ); - const provider = withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { - throw new Error("not used"); - }, - streamSimple: () => { - throw new Error("not used"); - }, - }, - }), - ); + const provider = testProvider(); const store = new InMemoryModelsStore(); - 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, - }); - 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, - }); - 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, - force: true, - }); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; + await provider.refreshModels?.(refresh); + await provider.refreshModels?.(refresh); + await provider.refreshModels?.({ ...refresh, force: true }); expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]); expect((await store.read(provider.id))?.models.map((entry) => entry.id)).toEqual(["dynamic"]); @@ -82,33 +79,40 @@ describe("remote catalog provider", () => { }); }); + it("prefers the newer of the generated and remote catalogs", async () => { + const localCatalogUrl = new URL(import.meta.url); + const localMtime = statSync(localCatalogUrl).mtimeMs; + const newerHeader = new Date(localMtime + 60_000).toUTCString(); + const responses = [ + new Response(JSON.stringify({ old: model("old") }), { + headers: { "last-modified": new Date(localMtime - 60_000).toUTCString() }, + }), + new Response(JSON.stringify({ newer: model("newer") }), { + headers: { "last-modified": newerHeader }, + }), + ]; + vi.spyOn(globalThis, "fetch").mockImplementation(async () => responses.shift() as Response); + const provider = testProvider(localCatalogUrl); + const store = new InMemoryModelsStore(); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; + + await provider.refreshModels?.(refresh); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static"]); + + await provider.refreshModels?.({ ...refresh, force: true }); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "newer"]); + expect(await store.read(provider.id)).toMatchObject({ lastModified: Date.parse(newerHeader) }); + }); + it("treats unimplemented pi.dev catalog routes as an unavailable overlay", async () => { vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response("not implemented", { status: 501 })); - const provider = withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { - throw new Error("not used"); - }, - streamSimple: () => { - throw new Error("not used"); - }, - }, - }), - ); + const provider = testProvider(); const store = new InMemoryModelsStore(); await expect( provider.refreshModels?.({ credential: { type: "api_key" }, - store: { - read: () => store.read(provider.id), - write: (entry) => store.write(provider.id, entry), - delete: () => store.delete(provider.id), - }, + store: scopedStore(store), allowNetwork: true, }), ).resolves.toBeUndefined();