feat(coding-agent): expose dynamic provider refresh
This commit is contained in:
@@ -1,35 +1,38 @@
|
|||||||
import type { Api, Model } from "./types.ts";
|
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. */
|
/** Persistent model catalogs keyed by provider ID. */
|
||||||
export interface ModelsStore {
|
export interface ModelsStore {
|
||||||
read(providerId: string): Promise<readonly Model<Api>[] | undefined>;
|
read(providerId: string): Promise<ModelsStoreEntry | undefined>;
|
||||||
write(providerId: string, models: readonly Model<Api>[]): Promise<void>;
|
write(providerId: string, entry: ModelsStoreEntry): Promise<void>;
|
||||||
delete(providerId: string): Promise<void>;
|
delete(providerId: string): Promise<void>;
|
||||||
}
|
}
|
||||||
|
|
||||||
/** ModelsStore scoped to one provider. Providers cannot access other providers' catalogs. */
|
/** ModelsStore scoped to one provider. Providers cannot access other providers' catalogs. */
|
||||||
export interface ProviderModelsStore {
|
export interface ProviderModelsStore {
|
||||||
read(): Promise<readonly Model<Api>[] | undefined>;
|
read(): Promise<ModelsStoreEntry | undefined>;
|
||||||
write(models: readonly Model<Api>[]): Promise<void>;
|
write(entry: ModelsStoreEntry): Promise<void>;
|
||||||
delete(): Promise<void>;
|
delete(): Promise<void>;
|
||||||
}
|
}
|
||||||
|
|
||||||
export class InMemoryModelsStore implements ModelsStore {
|
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> {
|
async read(providerId: string): Promise<ModelsStoreEntry | undefined> {
|
||||||
const models = this.models.get(providerId);
|
const entry = this.entries.get(providerId);
|
||||||
return models?.map((model) => structuredClone(model));
|
return entry ? structuredClone(entry) : undefined;
|
||||||
}
|
}
|
||||||
|
|
||||||
async write(providerId: string, models: readonly Model<Api>[]): Promise<void> {
|
async write(providerId: string, entry: ModelsStoreEntry): Promise<void> {
|
||||||
this.models.set(
|
this.entries.set(providerId, structuredClone(entry));
|
||||||
providerId,
|
|
||||||
models.map((model) => structuredClone(model)),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
async delete(providerId: string): Promise<void> {
|
async delete(providerId: string): Promise<void> {
|
||||||
this.models.delete(providerId);
|
this.entries.delete(providerId);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -282,7 +282,7 @@ class ModelsImpl implements MutableModels {
|
|||||||
if (options.signal?.aborted) return;
|
if (options.signal?.aborted) return;
|
||||||
const store: ProviderModelsStore = {
|
const store: ProviderModelsStore = {
|
||||||
read: () => this.modelsStore.read(provider.id),
|
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),
|
delete: () => this.modelsStore.delete(provider.id),
|
||||||
};
|
};
|
||||||
let stored: Credential | undefined;
|
let stored: Credential | undefined;
|
||||||
@@ -589,7 +589,7 @@ export function createProvider<TApi extends Api = Api>(input: CreateProviderOpti
|
|||||||
try {
|
try {
|
||||||
const stored = await context.store.read();
|
const stored = await context.store.read();
|
||||||
if (stored) {
|
if (stored) {
|
||||||
dynamicModels = stored
|
dynamicModels = stored.models
|
||||||
.filter((model) => model.provider === input.id)
|
.filter((model) => model.provider === input.id)
|
||||||
.map((model) => model as Model<TApi>);
|
.map((model) => model as Model<TApi>);
|
||||||
}
|
}
|
||||||
@@ -597,7 +597,7 @@ export function createProvider<TApi extends Api = Api>(input: CreateProviderOpti
|
|||||||
const refreshed = await fetchModels(context);
|
const refreshed = await fetchModels(context);
|
||||||
if (context.signal?.aborted) return;
|
if (context.signal?.aborted) return;
|
||||||
dynamicModels = refreshed;
|
dynamicModels = refreshed;
|
||||||
await context.store.write(refreshed);
|
await context.store.write({ models: refreshed, checkedAt: Date.now() });
|
||||||
} finally {
|
} finally {
|
||||||
inflightRefresh = undefined;
|
inflightRefresh = undefined;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -37,14 +37,14 @@ export function radiusProvider(options: RadiusProviderOptions = {}): Provider<"p
|
|||||||
inflightRefresh ??= (async () => {
|
inflightRefresh ??= (async () => {
|
||||||
try {
|
try {
|
||||||
const stored = await context.store.read();
|
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.
|
// Import catalogs cached by the pre-ModelsStore Radius implementation.
|
||||||
if (!stored && context.credential?.type === "oauth") {
|
if (!stored && context.credential?.type === "oauth") {
|
||||||
const legacy = getRadiusModels(id, context.credential);
|
const legacy = getRadiusModels(id, context.credential);
|
||||||
if (legacy.length > 0) {
|
if (legacy.length > 0) {
|
||||||
models = legacy;
|
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);
|
const config = await loadRadiusGatewayConfig(gateway, apiKey, context.signal);
|
||||||
if (context.signal?.aborted) return;
|
if (context.signal?.aborted) return;
|
||||||
models = getRadiusModelsFromConfig(id, config);
|
models = getRadiusModelsFromConfig(id, config);
|
||||||
await context.store.write(models);
|
await context.store.write({ models, checkedAt: Date.now() });
|
||||||
} finally {
|
} finally {
|
||||||
inflightRefresh = undefined;
|
inflightRefresh = undefined;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import { describe, expect, it } from "vitest";
|
|||||||
import { envApiKeyAuth } from "../src/auth/helpers.ts";
|
import { envApiKeyAuth } from "../src/auth/helpers.ts";
|
||||||
import type { AuthContext, AuthEvent } from "../src/auth/types.ts";
|
import type { AuthContext, AuthEvent } from "../src/auth/types.ts";
|
||||||
import { createModels, createProvider } from "../src/models.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 { builtinModels, builtinProviders } from "../src/providers/all.ts";
|
||||||
import { amazonBedrockProvider } from "../src/providers/amazon-bedrock.ts";
|
import { amazonBedrockProvider } from "../src/providers/amazon-bedrock.ts";
|
||||||
import { anthropicProvider } from "../src/providers/anthropic.ts";
|
import { anthropicProvider } from "../src/providers/anthropic.ts";
|
||||||
@@ -360,7 +360,7 @@ describe("createProvider", () => {
|
|||||||
credential: { type: "api_key" as const },
|
credential: { type: "api_key" as const },
|
||||||
store: {
|
store: {
|
||||||
read: () => store.read("dynamic"),
|
read: () => store.read("dynamic"),
|
||||||
write: (listed: readonly Model<Api>[]) => store.write("dynamic", listed),
|
write: (entry: ModelsStoreEntry) => store.write("dynamic", entry),
|
||||||
delete: () => store.delete("dynamic"),
|
delete: () => store.delete("dynamic"),
|
||||||
},
|
},
|
||||||
allowNetwork: true,
|
allowNetwork: true,
|
||||||
|
|||||||
@@ -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`.
|
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
|
```typescript
|
||||||
// Register a new provider with custom models
|
// 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)
|
// Override baseUrl for an existing provider (keeps all models)
|
||||||
pi.registerProvider("anthropic", {
|
pi.registerProvider("anthropic", {
|
||||||
baseUrl: "https://proxy.example.com"
|
baseUrl: "https://proxy.example.com"
|
||||||
@@ -1736,6 +1756,7 @@ pi.registerProvider("corporate-ai", {
|
|||||||
- `headers` - Custom headers to include in requests.
|
- `headers` - Custom headers to include in requests.
|
||||||
- `authHeader` - If true, adds `Authorization: Bearer` header automatically.
|
- `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.
|
- `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.
|
- `oauth` - OAuth provider config for `/login` support. When provided, the provider appears in the login menu.
|
||||||
- `streamSimple` - Custom streaming implementation for non-standard APIs.
|
- `streamSimple` - Custom streaming implementation for non-standard APIs.
|
||||||
|
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ import type {
|
|||||||
OAuthCredentials,
|
OAuthCredentials,
|
||||||
OAuthLoginCallbacks,
|
OAuthLoginCallbacks,
|
||||||
ProviderHeaders,
|
ProviderHeaders,
|
||||||
|
RefreshModelsContext,
|
||||||
SimpleStreamOptions,
|
SimpleStreamOptions,
|
||||||
TextContent,
|
TextContent,
|
||||||
ToolResultMessage,
|
ToolResultMessage,
|
||||||
@@ -1420,6 +1421,11 @@ export interface ProviderConfig {
|
|||||||
authHeader?: boolean;
|
authHeader?: boolean;
|
||||||
/** Models to register. If provided, replaces all existing models for this provider. */
|
/** Models to register. If provided, replaces all existing models for this provider. */
|
||||||
models?: ProviderModelConfig[];
|
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 provider for /login support. The `id` is set automatically from the provider name. */
|
||||||
oauth?: {
|
oauth?: {
|
||||||
/** Display name for the provider in login UI. */
|
/** Display name for the provider in login UI. */
|
||||||
|
|||||||
@@ -1,23 +1,23 @@
|
|||||||
import { join } from "node:path";
|
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 { getAgentDir } from "../config.ts";
|
||||||
import { type AuthStorageBackend, FileAuthStorageBackend } from "./auth-storage.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 {
|
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> {
|
async read(providerId: string): Promise<ModelsStoreEntry | undefined> {
|
||||||
return this.models.get(providerId);
|
return this.entries.get(providerId);
|
||||||
}
|
}
|
||||||
|
|
||||||
async write(providerId: string, models: readonly Model<Api>[]): Promise<void> {
|
async write(providerId: string, entry: ModelsStoreEntry): Promise<void> {
|
||||||
this.models.set(providerId, models);
|
this.entries.set(providerId, entry);
|
||||||
}
|
}
|
||||||
|
|
||||||
async delete(providerId: string): Promise<void> {
|
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) : {};
|
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) => ({
|
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) => {
|
await this.storage.withLockAsync(async (content) => {
|
||||||
const current = this.parse(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) };
|
return { result: undefined, next: JSON.stringify(current, null, 2) };
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import {
|
|||||||
type OAuthLoginCallbacks,
|
type OAuthLoginCallbacks,
|
||||||
type Provider,
|
type Provider,
|
||||||
type ProviderHeaders,
|
type ProviderHeaders,
|
||||||
|
type RefreshModelsContext,
|
||||||
type SimpleStreamOptions,
|
type SimpleStreamOptions,
|
||||||
type StreamOptions,
|
type StreamOptions,
|
||||||
} from "@earendil-works/pi-ai";
|
} from "@earendil-works/pi-ai";
|
||||||
@@ -63,6 +64,7 @@ export interface ProviderConfigInput {
|
|||||||
headers?: Record<string, string>;
|
headers?: Record<string, string>;
|
||||||
compat?: Model<Api>["compat"];
|
compat?: Model<Api>["compat"];
|
||||||
}>;
|
}>;
|
||||||
|
refreshModels?(context: RefreshModelsContext): Promise<NonNullable<ProviderConfigInput["models"]>>;
|
||||||
}
|
}
|
||||||
|
|
||||||
export type AuthStatus = {
|
export type AuthStatus = {
|
||||||
@@ -415,10 +417,17 @@ export function composeModelProvider(
|
|||||||
): Provider {
|
): Provider {
|
||||||
const config = modelConfig.getProvider(providerId);
|
const config = modelConfig.getProvider(providerId);
|
||||||
let extensionOAuthCredential: OAuthCredentials | undefined;
|
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,
|
// models.json modelOverrides are the topmost user-config layer: they apply once,
|
||||||
// after custom-model upserts, extension model replacement, and legacy OAuth projection.
|
// after custom-model upserts, extension model replacement, and legacy OAuth projection.
|
||||||
const getModels = () => {
|
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) {
|
if (extensionOAuthCredential && extension?.oauth?.modifyModels) {
|
||||||
models = extension.oauth.modifyModels(models, extensionOAuthCredential);
|
models = extension.oauth.modifyModels(models, extensionOAuthCredential);
|
||||||
}
|
}
|
||||||
@@ -464,9 +473,20 @@ export function composeModelProvider(
|
|||||||
auth: { ...(apiKey ? { apiKey } : {}), ...(oauth ? { oauth } : {}) },
|
auth: { ...(apiKey ? { apiKey } : {}), ...(oauth ? { oauth } : {}) },
|
||||||
getModels,
|
getModels,
|
||||||
refreshModels:
|
refreshModels:
|
||||||
base?.refreshModels || extension?.oauth?.modifyModels
|
base?.refreshModels || extension?.refreshModels || extension?.oauth?.modifyModels
|
||||||
? async (context) => {
|
? async (context) => {
|
||||||
await base?.refreshModels?.(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;
|
extensionOAuthCredential = context.credential?.type === "oauth" ? context.credential : undefined;
|
||||||
}
|
}
|
||||||
: undefined,
|
: undefined,
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
import type { Api, Model, Provider } from "@earendil-works/pi-ai";
|
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";
|
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>[] {
|
function mergeModels(baseline: readonly Model<Api>[], dynamic: readonly Model<Api>[]): Model<Api>[] {
|
||||||
const merged = [...baseline];
|
const merged = [...baseline];
|
||||||
@@ -38,22 +41,37 @@ export function withRemoteCatalog(provider: Provider, catalogBaseUrl: string = D
|
|||||||
inflightRefresh ??= (async () => {
|
inflightRefresh ??= (async () => {
|
||||||
try {
|
try {
|
||||||
const stored = await context.store.read();
|
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 (!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 url = new URL(`/api/models/providers/${encodeURIComponent(provider.id)}`, catalogBaseUrl);
|
||||||
const response = await fetch(url, {
|
const response = await fetch(url, {
|
||||||
headers: { accept: "application/json" },
|
headers: {
|
||||||
|
accept: "application/json",
|
||||||
|
"User-Agent": getPiUserAgent(VERSION),
|
||||||
|
},
|
||||||
signal: context.signal,
|
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) {
|
if (!response.ok) {
|
||||||
|
await context.store.write({ models: dynamicModels, 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());
|
||||||
if (context.signal?.aborted) return;
|
if (context.signal?.aborted) return;
|
||||||
dynamicModels = refreshed;
|
dynamicModels = refreshed;
|
||||||
await context.store.write(refreshed);
|
await context.store.write({ models: refreshed, checkedAt });
|
||||||
} finally {
|
} finally {
|
||||||
inflightRefresh = undefined;
|
inflightRefresh = undefined;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -18,8 +18,34 @@ function model(id: string): Model<"openai-completions"> {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
describe("legacy extension OAuth modifyModels", () => {
|
describe("extension provider model lifecycle", () => {
|
||||||
it("applies the synchronous projection after async credential initialization", async () => {
|
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({
|
const runtime = await ModelRuntime.create({
|
||||||
credentials: AuthStorage.inMemory({
|
credentials: AuthStorage.inMemory({
|
||||||
"extension-oauth": {
|
"extension-oauth": {
|
||||||
|
|||||||
@@ -36,15 +36,16 @@ describe("FileModelsStore", () => {
|
|||||||
const path = join(dir, "models-store.json");
|
const path = join(dir, "models-store.json");
|
||||||
const store = new FileModelsStore(path);
|
const store = new FileModelsStore(path);
|
||||||
|
|
||||||
await store.write("one", [model("one", "m1")]);
|
await store.write("one", { models: [model("one", "m1")], checkedAt: 100 });
|
||||||
await store.write("two", [model("two", "m2")]);
|
await store.write("two", { models: [model("two", "m2")], checkedAt: 200 });
|
||||||
|
|
||||||
const reloaded = new FileModelsStore(path);
|
const reloaded = new FileModelsStore(path);
|
||||||
expect((await reloaded.read("one"))?.map((entry) => entry.id)).toEqual(["m1"]);
|
expect((await reloaded.read("one"))?.models.map((entry) => entry.id)).toEqual(["m1"]);
|
||||||
expect((await reloaded.read("two"))?.map((entry) => entry.id)).toEqual(["m2"]);
|
expect((await reloaded.read("one"))?.checkedAt).toBe(100);
|
||||||
|
expect((await reloaded.read("two"))?.models.map((entry) => entry.id)).toEqual(["m2"]);
|
||||||
|
|
||||||
await reloaded.delete("one");
|
await reloaded.delete("one");
|
||||||
expect(await reloaded.read("one")).toBeUndefined();
|
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"]);
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -87,7 +87,7 @@ describe("Radius provider", () => {
|
|||||||
});
|
});
|
||||||
|
|
||||||
expect(runtime.getModel(RADIUS_PROVIDER_ID, "auto")).toBeDefined();
|
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" });
|
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 { createProvider, InMemoryModelsStore, type Model } 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 { withRemoteCatalog } from "../src/core/remote-catalog-provider.ts";
|
import { withRemoteCatalog } from "../src/core/remote-catalog-provider.ts";
|
||||||
|
|
||||||
function model(id: string): Model<"openai-completions"> {
|
function model(id: string): Model<"openai-completions"> {
|
||||||
@@ -20,8 +21,8 @@ function model(id: string): Model<"openai-completions"> {
|
|||||||
afterEach(() => vi.restoreAllMocks());
|
afterEach(() => vi.restoreAllMocks());
|
||||||
|
|
||||||
describe("remote catalog provider", () => {
|
describe("remote catalog provider", () => {
|
||||||
it("parses catalogs keyed by model ID", async () => {
|
it("parses keyed catalogs, sends version headers, and observes the refresh TTL", async () => {
|
||||||
vi.spyOn(globalThis, "fetch").mockResolvedValue(
|
const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue(
|
||||||
new Response(JSON.stringify({ dynamic: model("dynamic") }), {
|
new Response(JSON.stringify({ dynamic: model("dynamic") }), {
|
||||||
status: 200,
|
status: 200,
|
||||||
headers: { "content-type": "application/json" },
|
headers: { "content-type": "application/json" },
|
||||||
@@ -47,14 +48,27 @@ describe("remote catalog provider", () => {
|
|||||||
credential: { type: "api_key" },
|
credential: { type: "api_key" },
|
||||||
store: {
|
store: {
|
||||||
read: () => store.read(provider.id),
|
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),
|
delete: () => store.delete(provider.id),
|
||||||
},
|
},
|
||||||
allowNetwork: true,
|
allowNetwork: 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))?.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 () => {
|
it("treats unimplemented pi.dev catalog routes as an unavailable overlay", async () => {
|
||||||
@@ -81,13 +95,13 @@ describe("remote catalog provider", () => {
|
|||||||
credential: { type: "api_key" },
|
credential: { type: "api_key" },
|
||||||
store: {
|
store: {
|
||||||
read: () => store.read(provider.id),
|
read: () => store.read(provider.id),
|
||||||
write: (models) => store.write(provider.id, models),
|
write: (entry) => store.write(provider.id, entry),
|
||||||
delete: () => store.delete(provider.id),
|
delete: () => store.delete(provider.id),
|
||||||
},
|
},
|
||||||
allowNetwork: true,
|
allowNetwork: true,
|
||||||
}),
|
}),
|
||||||
).resolves.toBeUndefined();
|
).resolves.toBeUndefined();
|
||||||
expect(provider.getModels().map((entry) => entry.id)).toEqual(["static"]);
|
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) });
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
Reference in New Issue
Block a user