fix(coding-agent): prefer newer generated model catalogs
This commit is contained in:
@@ -2,6 +2,8 @@ import type { Api, Model } from "./types.ts";
|
||||
|
||||
export interface ModelsStoreEntry {
|
||||
models: readonly Model<Api>[];
|
||||
/** Unix timestamp from the remote catalog's Last-Modified header. */
|
||||
lastModified?: number;
|
||||
/** Unix timestamp of the last completed remote check. */
|
||||
checkedAt?: number;
|
||||
}
|
||||
|
||||
@@ -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<TProvider extends BuiltinProvider>(
|
||||
provider: TProvider,
|
||||
): Model<BuiltinModelApi<TProvider, keyof (typeof MODELS)[TProvider]>>[] {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<Api>[] {
|
||||
.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. */
|
||||
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 inflightRefresh: Promise<void> | 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;
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user