fix(coding-agent): prefer newer generated model catalogs

This commit is contained in:
Armin Ronacher
2026-07-21 11:11:08 +02:00
parent 890b3547af
commit 54fad505b9
5 changed files with 125 additions and 72 deletions
+2
View File
@@ -2,6 +2,8 @@ import type { Api, Model } from "./types.ts";
export interface ModelsStoreEntry { export interface ModelsStoreEntry {
models: readonly Model<Api>[]; models: readonly Model<Api>[];
/** Unix timestamp from the remote catalog's Last-Modified header. */
lastModified?: number;
/** Unix timestamp of the last completed remote check. */ /** Unix timestamp of the last completed remote check. */
checkedAt?: number; checkedAt?: number;
} }
+5
View File
@@ -67,6 +67,11 @@ export function getBuiltinProviders(): BuiltinProvider[] {
return Object.keys(MODELS) as 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<TProvider extends BuiltinProvider>( export function getBuiltinModels<TProvider extends BuiltinProvider>(
provider: TProvider, provider: TProvider,
): Model<BuiltinModelApi<TProvider, keyof (typeof MODELS)[TProvider]>>[] { ): Model<BuiltinModelApi<TProvider, keyof (typeof MODELS)[TProvider]>>[] {
@@ -143,7 +143,15 @@ export class ModelRuntime implements Models {
const providers = builtinProviderCatalog const providers = builtinProviderCatalog
.builtinProviders() .builtinProviders()
.map((provider) => .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( const runtime = new ModelRuntime(
credentials, credentials,
@@ -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 { VERSION } from "../config.ts";
import { getPiUserAgent } from "../utils/pi-user-agent.ts"; import { getPiUserAgent } from "../utils/pi-user-agent.ts";
@@ -29,8 +30,26 @@ function parseCatalog(providerId: string, value: unknown): Model<Api>[] {
.map((model) => ({ ...model, provider: providerId })); .map((model) => ({ ...model, provider: providerId }));
} }
function remoteModels(
entry: ModelsStoreEntry | undefined,
localLastModified: number | undefined,
): readonly Model<Api>[] {
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. */ /** 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<Api>[] = []; let dynamicModels: readonly Model<Api>[] = [];
let inflightRefresh: Promise<void> | undefined; let inflightRefresh: Promise<void> | undefined;
@@ -40,12 +59,21 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D
refreshModels: (context) => { refreshModels: (context) => {
inflightRefresh ??= (async () => { inflightRefresh ??= (async () => {
try { try {
const localLastModified = localCatalogUrl
? await stat(localCatalogUrl).then(
(value) => value.mtimeMs,
() => undefined,
)
: undefined;
const stored = await context.store.read(); 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.allowNetwork || context.signal?.aborted) return;
if ( if (
!context.force && !context.force &&
stored?.checkedAt !== undefined && stored?.checkedAt !== undefined &&
stored.lastModified !== undefined &&
Date.now() - stored.checkedAt < REMOTE_CATALOG_REFRESH_INTERVAL_MS Date.now() - stored.checkedAt < REMOTE_CATALOG_REFRESH_INTERVAL_MS
) { ) {
return; return;
@@ -62,17 +90,23 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D
if (context.signal?.aborted) return; if (context.signal?.aborted) return;
const checkedAt = Date.now(); const checkedAt = Date.now();
if (response.status === 404 || response.status === 501) { if (response.status === 404 || response.status === 501) {
await context.store.write({ models: dynamicModels, checkedAt }); await context.store.write({ ...(stored ?? { models: [] }), checkedAt, lastModified: 0 });
return; return;
} }
if (!response.ok) { 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}`); throw new Error(`Model catalog request failed for ${provider.id}: ${response.status}`);
} }
const refreshed = parseCatalog(provider.id, await response.json()); const refreshed = parseCatalog(provider.id, await response.json());
const lastModified = Date.parse(response.headers.get("last-modified") ?? "");
if (context.signal?.aborted) return; if (context.signal?.aborted) return;
dynamicModels = refreshed; const entry = {
await context.store.write({ models: refreshed, checkedAt }); models: refreshed,
checkedAt,
lastModified: Number.isNaN(lastModified) ? 0 : lastModified,
};
dynamicModels = remoteModels(entry, localLastModified);
await context.store.write(entry);
} finally { } finally {
inflightRefresh = undefined; inflightRefresh = undefined;
} }
@@ -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 { afterEach, describe, expect, it, vi } from "vitest";
import { VERSION } from "../src/config.ts"; import { VERSION } from "../src/config.ts";
import { withRemoteCatalog } from "../src/core/remote-catalog-provider.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()); afterEach(() => vi.restoreAllMocks());
describe("remote catalog provider", () => { describe("remote catalog provider", () => {
@@ -29,50 +64,12 @@ describe("remote catalog provider", () => {
headers: { "content-type": "application/json" }, headers: { "content-type": "application/json" },
}), }),
); );
const provider = withRemoteCatalog( const provider = testProvider();
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 store = new InMemoryModelsStore(); const store = new InMemoryModelsStore();
await provider.refreshModels?.({ const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true };
credential: { type: "api_key" }, await provider.refreshModels?.(refresh);
store: { await provider.refreshModels?.(refresh);
read: () => store.read(provider.id), await provider.refreshModels?.({ ...refresh, force: true });
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,
});
expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]); expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]);
expect((await store.read(provider.id))?.models.map((entry) => entry.id)).toEqual(["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 () => { it("treats unimplemented pi.dev catalog routes as an unavailable overlay", async () => {
vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response("not implemented", { status: 501 })); vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response("not implemented", { status: 501 }));
const provider = withRemoteCatalog( const provider = testProvider();
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 store = new InMemoryModelsStore(); const store = new InMemoryModelsStore();
await expect( await expect(
provider.refreshModels?.({ provider.refreshModels?.({
credential: { type: "api_key" }, credential: { type: "api_key" },
store: { store: scopedStore(store),
read: () => store.read(provider.id),
write: (entry) => store.write(provider.id, entry),
delete: () => store.delete(provider.id),
},
allowNetwork: true, allowNetwork: true,
}), }),
).resolves.toBeUndefined(); ).resolves.toBeUndefined();