feat(coding-agent): replace model registry with model runtime

Move provider auth and OAuth flows onto pi-ai Models, compose models.json and extension overlays through ModelRuntime, and retain ModelRegistry as an extension compatibility facade.
This commit is contained in:
Mario Zechner
2026-07-14 17:48:45 +02:00
parent 6731a0ba9e
commit 9993c96907
133 changed files with 5103 additions and 4340 deletions
@@ -3,9 +3,8 @@ import type { ThinkingLevel } from "@earendil-works/pi-agent-core";
import type { Model } from "@earendil-works/pi-ai";
import { getAgentDir } from "../config.ts";
import { resolvePath } from "../utils/paths.ts";
import { AuthStorage } from "./auth-storage.ts";
import type { SessionStartEvent, ToolDefinition } from "./extensions/index.ts";
import { ModelRegistry } from "./model-registry.ts";
import { ModelRuntime } from "./model-runtime.ts";
import {
DefaultResourceLoader,
type DefaultResourceLoaderOptions,
@@ -38,9 +37,8 @@ export interface AgentSessionRuntimeDiagnostic {
export interface CreateAgentSessionServicesOptions {
cwd: string;
agentDir?: string;
authStorage?: AuthStorage;
settingsManager?: SettingsManager;
modelRegistry?: ModelRegistry;
modelRuntime?: ModelRuntime;
extensionFlagValues?: Map<string, boolean | string>;
resourceLoaderOptions?: Omit<DefaultResourceLoaderOptions, "cwd" | "agentDir" | "settingsManager">;
resourceLoaderReloadOptions?: ResourceLoaderReloadOptions;
@@ -74,9 +72,8 @@ export interface CreateAgentSessionFromServicesOptions {
export interface AgentSessionServices {
cwd: string;
agentDir: string;
authStorage: AuthStorage;
modelRuntime: ModelRuntime;
settingsManager: SettingsManager;
modelRegistry: ModelRegistry;
resourceLoader: ResourceLoader;
diagnostics: AgentSessionRuntimeDiagnostic[];
}
@@ -139,9 +136,13 @@ export async function createAgentSessionServices(
): Promise<AgentSessionServices> {
const cwd = resolvePath(options.cwd);
const agentDir = options.agentDir ? resolvePath(options.agentDir) : getAgentDir();
const authStorage = options.authStorage ?? AuthStorage.create(join(agentDir, "auth.json"));
const modelRuntime =
options.modelRuntime ??
(await ModelRuntime.create({
authPath: join(agentDir, "auth.json"),
modelsPath: join(agentDir, "models.json"),
}));
const settingsManager = options.settingsManager ?? SettingsManager.create(cwd, agentDir);
const modelRegistry = options.modelRegistry ?? ModelRegistry.create(authStorage, join(agentDir, "models.json"));
const resourceLoader = new DefaultResourceLoader({
...(options.resourceLoaderOptions ?? {}),
cwd,
@@ -154,7 +155,7 @@ export async function createAgentSessionServices(
const extensionsResult = resourceLoader.getExtensions();
for (const { name, config, extensionPath } of extensionsResult.runtime.pendingProviderRegistrations) {
try {
modelRegistry.registerProvider(name, config);
modelRuntime.registerProvider(name, config);
} catch (error) {
const message = error instanceof Error ? error.message : String(error);
diagnostics.push({
@@ -169,9 +170,8 @@ export async function createAgentSessionServices(
return {
cwd,
agentDir,
authStorage,
modelRuntime,
settingsManager,
modelRegistry,
resourceLoader,
diagnostics,
};
@@ -190,9 +190,8 @@ export async function createAgentSessionFromServices(
return createAgentSession({
cwd: options.services.cwd,
agentDir: options.services.agentDir,
authStorage: options.services.authStorage,
modelRuntime: options.services.modelRuntime,
settingsManager: options.services.settingsManager,
modelRegistry: options.services.modelRegistry,
resourceLoader: options.services.resourceLoader,
sessionManager: options.sessionManager,
model: options.model,
+68 -35
View File
@@ -24,7 +24,15 @@ import type {
PrepareNextTurnContext,
ThinkingLevel,
} from "@earendil-works/pi-agent-core";
import type { AssistantMessage, ImageContent, Message, Model, TextContent } from "@earendil-works/pi-ai/compat";
import type {
AssistantMessage,
AuthResult,
ImageContent,
Message,
Model,
ProviderHeaders,
TextContent,
} from "@earendil-works/pi-ai/compat";
import {
clampThinkingLevel,
cleanupSessionResources,
@@ -83,7 +91,8 @@ import {
} from "./extensions/index.ts";
import { emitSessionShutdownEvent } from "./extensions/runner.ts";
import type { BashExecutionMessage, CustomMessage } from "./messages.ts";
import type { ModelRegistry } from "./model-registry.ts";
import { ModelRegistry } from "./model-registry.ts";
import type { ModelRuntime } from "./model-runtime.ts";
import { expandPromptTemplate, type PromptTemplate } from "./prompt-templates.ts";
import type { ResourceExtensionPaths, ResourceLoader } from "./resource-loader.ts";
import type { BranchSummaryEntry, CompactionEntry, SessionEntry, SessionManager } from "./session-manager.ts";
@@ -159,6 +168,12 @@ export type AgentSessionEventListener = (event: AgentSessionEvent) => void;
// Types
// ============================================================================
function withoutDeletedHeaders(headers: ProviderHeaders | undefined): Record<string, string> | undefined {
return headers
? Object.fromEntries(Object.entries(headers).filter((entry): entry is [string, string] => entry[1] !== null))
: undefined;
}
export interface AgentSessionConfig {
agent: Agent;
sessionManager: SessionManager;
@@ -170,8 +185,8 @@ export interface AgentSessionConfig {
resourceLoader: ResourceLoader;
/** SDK custom tools registered outside extensions */
customTools?: ToolDefinition[];
/** Model registry for API key resolution and model discovery */
modelRegistry: ModelRegistry;
/** Canonical model/auth runtime used by coding-agent internals. */
modelRuntime: ModelRuntime;
/** Initial active built-in tool names. Default: [read, bash, edit, write] */
initialActiveToolNames?: string[];
/** Optional allowlist of tool names. When provided, only these tool names are exposed. */
@@ -325,8 +340,7 @@ export class AgentSession {
private _extensionErrorListener?: ExtensionErrorListener;
private _extensionErrorUnsubscriber?: () => void;
// Model registry for API key resolution
private _modelRegistry: ModelRegistry;
private _modelRuntime: ModelRuntime;
// Tool registry for extension getTools/setTools
private _toolRegistry: Map<string, AgentTool> = new Map();
@@ -347,7 +361,7 @@ export class AgentSession {
this._resourceLoader = config.resourceLoader;
this._customTools = config.customTools ?? [];
this._cwd = config.cwd;
this._modelRegistry = config.modelRegistry;
this._modelRuntime = config.modelRuntime;
this._extensionRunnerRef = config.extensionRunnerRef;
this._initialActiveToolNames = config.initialActiveToolNames;
this._allowedToolNames = config.allowedToolNames ? new Set(config.allowedToolNames) : undefined;
@@ -367,9 +381,8 @@ export class AgentSession {
});
}
/** Model registry for API key resolution and model discovery */
get modelRegistry(): ModelRegistry {
return this._modelRegistry;
get modelRuntime(): ModelRuntime {
return this._modelRuntime;
}
private async _getRequiredRequestAuth(model: Model<any>): Promise<{
@@ -377,18 +390,25 @@ export class AgentSession {
headers?: Record<string, string>;
env?: Record<string, string>;
}> {
const result = await this._modelRegistry.getApiKeyAndHeaders(model);
if (!result.ok) {
if (result.error.startsWith("No API key found")) {
let result: AuthResult | undefined;
try {
result = await this._modelRuntime.getAuth(model);
} catch (error) {
const cause = error instanceof Error ? error.cause : undefined;
if (cause instanceof Error && cause.message === "authHeader requires a resolved API key") {
throw new Error(formatNoApiKeyFoundMessage(model.provider));
}
throw new Error(result.error);
throw error;
}
if (result.apiKey) {
return { apiKey: result.apiKey, headers: result.headers, env: result.env };
if (result?.auth.apiKey) {
return {
apiKey: result.auth.apiKey,
headers: withoutDeletedHeaders(result.auth.headers),
env: result.env,
};
}
const isOAuth = this._modelRegistry.isUsingOAuth(model);
const isOAuth = this._modelRuntime.isUsingOAuth(model.provider);
if (isOAuth) {
throw new Error(
`Authentication failed for "${model.provider}". ` +
@@ -408,8 +428,14 @@ export class AgentSession {
return this._getRequiredRequestAuth(model);
}
const result = await this._modelRegistry.getApiKeyAndHeaders(model);
return result.ok ? { apiKey: result.apiKey, headers: result.headers, env: result.env } : {};
try {
const result = await this._modelRuntime.getAuth(model);
return result
? { apiKey: result.auth.apiKey, headers: withoutDeletedHeaders(result.auth.headers), env: result.env }
: {};
} catch {
return {};
}
}
/**
@@ -1141,8 +1167,11 @@ export class AgentSession {
throw new Error(formatNoModelSelectedMessage());
}
if (!this._modelRegistry.hasConfiguredAuth(this.model)) {
const isOAuth = this._modelRegistry.isUsingOAuth(this.model);
const hasConfiguredAuth =
this._modelRuntime.hasConfiguredAuth(this.model.provider) ||
(await this._modelRuntime.checkAuth(this.model.provider)) !== undefined;
if (!hasConfiguredAuth) {
const isOAuth = this._modelRuntime.isUsingOAuth(this.model.provider);
if (isOAuth) {
throw new Error(
`Authentication failed for "${this.model.provider}". ` +
@@ -1535,7 +1564,7 @@ export class AgentSession {
* @throws Error if no auth is configured for the model
*/
async setModel(model: Model<any>): Promise<void> {
if (!this._modelRegistry.hasConfiguredAuth(model)) {
if (!(await this._modelRuntime.checkAuth(model.provider))) {
throw new Error(`No API key for ${model.provider}/${model.id}`);
}
@@ -1565,7 +1594,13 @@ export class AgentSession {
}
private async _cycleScopedModel(direction: "forward" | "backward"): Promise<ModelCycleResult | undefined> {
const scopedModels = this._scopedModels.filter((scoped) => this._modelRegistry.hasConfiguredAuth(scoped.model));
const checks = await Promise.all(
this._scopedModels.map(async (scoped) => ({
scoped,
auth: await this._modelRuntime.checkAuth(scoped.model.provider),
})),
);
const scopedModels = checks.filter(({ auth }) => auth !== undefined).map(({ scoped }) => scoped);
if (scopedModels.length <= 1) return undefined;
const currentModel = this.model;
@@ -1594,7 +1629,7 @@ export class AgentSession {
}
private async _cycleAvailableModel(direction: "forward" | "backward"): Promise<ModelCycleResult | undefined> {
const availableModels = await this._modelRegistry.getAvailable();
const availableModels = await this._modelRuntime.getAvailable();
if (availableModels.length <= 1) return undefined;
const currentModel = this.model;
@@ -2004,12 +2039,10 @@ export class AgentSession {
let headers: Record<string, string> | undefined;
let env: Record<string, string> | undefined;
if (this.agent.streamFn === streamSimple) {
const authResult = await this._modelRegistry.getApiKeyAndHeaders(this.model);
if (!authResult.ok || !authResult.apiKey) {
return false;
}
apiKey = authResult.apiKey;
headers = authResult.headers;
const authResult = await this._modelRuntime.getAuth(this.model);
if (!authResult?.auth.apiKey) return false;
apiKey = authResult.auth.apiKey;
headers = withoutDeletedHeaders(authResult.auth.headers);
env = authResult.env;
} else {
({ apiKey, headers, env } = await this._getCompactionRequestAuth(this.model));
@@ -2267,7 +2300,7 @@ export class AgentSession {
return;
}
const refreshedModel = this._modelRegistry.find(currentModel.provider, currentModel.id);
const refreshedModel = this._modelRuntime.getModel(currentModel.provider, currentModel.id);
if (!refreshedModel || refreshedModel === currentModel) {
return;
}
@@ -2343,7 +2376,7 @@ export class AgentSession {
refreshTools: () => this._refreshToolRegistry(),
getCommands,
setModel: async (model) => {
if (!this.modelRegistry.hasConfiguredAuth(model)) return false;
if (!this._modelRuntime.hasConfiguredAuth(model.provider)) return false;
await this.setModel(model);
return true;
},
@@ -2383,11 +2416,11 @@ export class AgentSession {
},
{
registerProvider: (name, config) => {
this._modelRegistry.registerProvider(name, config);
this._modelRuntime.registerProvider(name, config);
this._refreshCurrentModelFromRegistry();
},
unregisterProvider: (name) => {
this._modelRegistry.unregisterProvider(name);
this._modelRuntime.unregisterProvider(name);
this._refreshCurrentModelFromRegistry();
},
},
@@ -2523,7 +2556,7 @@ export class AgentSession {
extensionsResult.runtime,
this._cwd,
this.sessionManager,
this._modelRegistry,
new ModelRegistry(this._modelRuntime),
);
if (this._extensionRunnerRef) {
this._extensionRunnerRef.current = this._extensionRunner;
+50 -318
View File
@@ -1,19 +1,9 @@
/**
* Credential storage for API keys and OAuth tokens.
* Handles loading, saving, and refreshing credentials from auth.json.
*
* Uses file locking to prevent race conditions when multiple pi instances
* try to refresh tokens simultaneously.
* CredentialStore implementation backed by auth.json.
* Provider auth orchestration belongs to ModelRuntime and pi-ai Models.
*/
import {
findEnvKeys,
getEnvApiKey,
type OAuthCredentials,
type OAuthLoginCallbacks,
type OAuthProviderId,
} from "@earendil-works/pi-ai/compat";
import { getOAuthApiKey, getOAuthProvider, getOAuthProviders } from "@earendil-works/pi-ai/oauth";
import type { Credential, CredentialInfo, CredentialStore } from "@earendil-works/pi-ai";
import { chmodSync, existsSync, mkdirSync, readFileSync, writeFileSync } from "fs";
import { dirname, join } from "path";
import lockfile from "proper-lockfile";
@@ -21,29 +11,7 @@ import { getAgentDir } from "../config.ts";
import { normalizePath } from "../utils/paths.ts";
import { resolveConfigValue } from "./resolve-config-value.ts";
export type ApiKeyCredential = {
type: "api_key";
key: string;
env?: Record<string, string>;
};
export type OAuthCredential = {
type: "oauth";
} & OAuthCredentials;
export type AuthCredential = ApiKeyCredential | OAuthCredential;
export type AuthStorageData = Record<string, AuthCredential>;
export type AuthStatus = {
configured: boolean;
source?: "stored" | "runtime" | "environment" | "fallback" | "models_json_key" | "models_json_command";
label?: string;
};
export interface GetApiKeyOptions {
includeFallback?: boolean;
}
type AuthStorageData = Record<string, Credential>;
type LockResult<T> = {
result: T;
@@ -200,11 +168,8 @@ export class InMemoryAuthStorageBackend implements AuthStorageBackend {
/**
* Credential storage backed by a JSON file.
*/
export class AuthStorage {
export class AuthStorage implements CredentialStore {
private data: AuthStorageData = {};
private runtimeOverrides: Map<string, string> = new Map();
private loadError: Error | null = null;
private errors: Error[] = [];
private storage: AuthStorageBackend;
private constructor(storage: AuthStorageBackend) {
@@ -226,26 +191,6 @@ export class AuthStorage {
return AuthStorage.fromStorage(storage);
}
/**
* Set a runtime API key override (not persisted to disk).
* Used for CLI --api-key flag.
*/
setRuntimeApiKey(provider: string, apiKey: string): void {
this.runtimeOverrides.set(provider, apiKey);
}
/**
* Remove a runtime API key override.
*/
removeRuntimeApiKey(provider: string): void {
this.runtimeOverrides.delete(provider);
}
private recordError(error: unknown): void {
const normalizedError = error instanceof Error ? error : new Error(String(error));
this.errors.push(normalizedError);
}
private parseStorageData(content: string | undefined): AuthStorageData {
if (!content) {
return {};
@@ -264,276 +209,63 @@ export class AuthStorage {
return { result: undefined };
});
this.data = this.parseStorageData(content);
this.loadError = null;
} catch (error) {
this.loadError = error as Error;
this.recordError(error);
} catch {
// Preserve the last valid in-memory snapshot.
}
}
private persistProviderChange(provider: string, credential: AuthCredential | undefined): AuthStorageData {
if (this.loadError) {
this.reload();
}
if (this.loadError) {
const error = new Error(
`Cannot update auth storage because it could not be loaded: ${this.loadError.message}`,
);
this.recordError(error);
throw error;
}
try {
let persistedData: AuthStorageData = {};
this.storage.withLock((current) => {
const currentData = this.parseStorageData(current);
const merged: AuthStorageData = { ...currentData };
if (credential) {
merged[provider] = credential;
} else {
delete merged[provider];
}
persistedData = merged;
return { result: undefined, next: JSON.stringify(merged, null, 2) };
});
this.loadError = null;
return persistedData;
} catch (error) {
this.recordError(error);
throw error;
}
async read(provider: string): Promise<Credential | undefined> {
const credential = this.data[provider];
if (credential?.type !== "api_key") return credential;
if (credential.key === undefined) return credential;
return { ...credential, key: resolveConfigValue(credential.key, credential.env) };
}
/**
* Get credential for a provider.
*/
get(provider: string): AuthCredential | undefined {
return this.data[provider] ?? undefined;
}
/**
* Get provider-scoped environment values for an API key credential.
*/
getProviderEnv(provider: string): Record<string, string> | undefined {
const cred = this.data[provider];
return cred?.type === "api_key" && cred.env ? { ...cred.env } : undefined;
}
/**
* Set credential for a provider.
*/
set(provider: string, credential: AuthCredential): void {
this.data = this.persistProviderChange(provider, credential);
}
/**
* Remove credential for a provider.
*/
remove(provider: string): void {
this.data = this.persistProviderChange(provider, undefined);
}
/**
* List all providers with credentials.
*/
list(): string[] {
return Object.keys(this.data);
}
/**
* Check if credentials exist for a provider in auth.json.
*/
has(provider: string): boolean {
return provider in this.data;
}
/**
* Check if any form of auth is configured for a provider.
* Unlike getApiKey(), this doesn't refresh OAuth tokens.
*/
hasAuth(provider: string): boolean {
if (this.runtimeOverrides.has(provider)) return true;
if (this.data[provider]) return true;
if (getEnvApiKey(provider)) return true;
return false;
}
/**
* Return auth status without exposing credential values or refreshing tokens.
*/
getAuthStatus(provider: string): AuthStatus {
if (this.data[provider]) {
return { configured: true, source: "stored" };
}
if (this.runtimeOverrides.has(provider)) {
return { configured: false, source: "runtime", label: "--api-key" };
}
const envKeys = findEnvKeys(provider);
if (envKeys?.[0]) {
return { configured: false, source: "environment", label: envKeys[0] };
}
return { configured: false };
}
/**
* Get all credentials (for passing to getOAuthApiKey).
*/
getAll(): AuthStorageData {
return { ...this.data };
}
drainErrors(): Error[] {
const drained = [...this.errors];
this.errors = [];
return drained;
}
/**
* Login to an OAuth provider.
*/
async login(providerId: OAuthProviderId, callbacks: OAuthLoginCallbacks): Promise<void> {
const provider = getOAuthProvider(providerId);
if (!provider) {
throw new Error(`Unknown OAuth provider: ${providerId}`);
}
const credentials = await provider.login(callbacks);
this.set(providerId, { type: "oauth", ...credentials });
}
/**
* Logout from a provider.
*/
logout(provider: string): void {
this.remove(provider);
}
/**
* Refresh OAuth token with backend locking to prevent race conditions.
* Multiple pi instances may try to refresh simultaneously when tokens expire.
*/
private async refreshOAuthTokenWithLock(
providerId: OAuthProviderId,
): Promise<{ apiKey: string; newCredentials: OAuthCredentials } | null> {
const provider = getOAuthProvider(providerId);
if (!provider) {
return null;
}
const result = await this.storage.withLockAsync(async (current) => {
const currentData = this.parseStorageData(current);
this.data = currentData;
this.loadError = null;
const cred = currentData[providerId];
if (cred?.type !== "oauth") {
return { result: null };
async modify(
provider: string,
fn: (current: Credential | undefined) => Promise<Credential | undefined>,
): Promise<Credential | undefined> {
return this.storage.withLockAsync(async (content) => {
const currentData = this.parseStorageData(content);
const next = await fn(currentData[provider]);
if (next === undefined) {
this.data = currentData;
return { result: currentData[provider] };
}
if (Date.now() < cred.expires) {
return { result: { apiKey: provider.getApiKey(cred), newCredentials: cred } };
}
const oauthCreds: Record<string, OAuthCredentials> = {};
for (const [key, value] of Object.entries(currentData)) {
if (value.type === "oauth") {
oauthCreds[key] = value;
}
}
const refreshed = await getOAuthApiKey(providerId, oauthCreds);
if (!refreshed) {
return { result: null };
}
const merged: AuthStorageData = {
...currentData,
[providerId]: { type: "oauth", ...refreshed.newCredentials },
};
const merged: AuthStorageData = { ...currentData, [provider]: next };
this.data = merged;
this.loadError = null;
return { result: refreshed, next: JSON.stringify(merged, null, 2) };
return { result: next, next: JSON.stringify(merged, null, 2) };
});
return result;
}
/**
* Get API key for a provider.
* Priority:
* 1. Runtime override (CLI --api-key)
* 2. API key from auth.json
* 3. OAuth token from auth.json (auto-refreshed with locking)
* 4. Environment variable
*/
async getApiKey(providerId: string, options: GetApiKeyOptions = {}): Promise<string | undefined> {
// Runtime override takes highest priority
const runtimeKey = this.runtimeOverrides.get(providerId);
if (runtimeKey) {
return runtimeKey;
}
const cred = this.data[providerId];
if (cred?.type === "api_key") {
return resolveConfigValue(cred.key, cred.env);
}
if (cred?.type === "oauth") {
const provider = getOAuthProvider(providerId);
if (!provider) {
// Unknown OAuth provider, can't get API key
return undefined;
}
// Check if token needs refresh
const needsRefresh = Date.now() >= cred.expires;
if (needsRefresh) {
// Use locked refresh to prevent race conditions
try {
const result = await this.refreshOAuthTokenWithLock(providerId);
if (result) {
return result.apiKey;
}
} catch (error) {
this.recordError(error);
// Refresh failed - re-read file to check if another instance succeeded
this.reload();
const updatedCred = this.data[providerId];
if (updatedCred?.type === "oauth" && Date.now() < updatedCred.expires) {
// Another instance refreshed successfully, use those credentials
return provider.getApiKey(updatedCred);
}
// Refresh truly failed - return undefined so model discovery skips this provider
// User can /login to re-authenticate (credentials preserved for retry)
return undefined;
}
} else {
// Token not expired, use current access token
return provider.getApiKey(cred);
}
}
if (options.includeFallback === false) return undefined;
// Fall back to environment variable
const envKey = getEnvApiKey(providerId);
if (envKey) return envKey;
return undefined;
async delete(provider: string): Promise<void> {
await this.storage.withLockAsync(async (content) => {
const currentData = this.parseStorageData(content);
delete currentData[provider];
this.data = currentData;
return { result: undefined, next: JSON.stringify(currentData, null, 2) };
});
}
/**
* Get all registered OAuth providers
*/
getOAuthProviders() {
return getOAuthProviders();
/** List credential metadata without resolving configured key values. */
async list(): Promise<readonly CredentialInfo[]> {
return Object.entries(this.data).map(([providerId, credential]) => ({ providerId, type: credential.type }));
}
}
/**
* One-off synchronous read of a stored credential from an auth.json file,
* without instantiating a store or resolving configured key values.
*/
export function readStoredCredential(
providerId: string,
authPath: string = join(getAgentDir(), "auth.json"),
): Credential | undefined {
try {
const data = JSON.parse(readFileSync(normalizePath(authPath), "utf-8")) as AuthStorageData;
return data[providerId];
} catch {
return undefined;
}
}
@@ -29,9 +29,9 @@ export interface CacheWasteTotals {
missCount: number;
}
/** Minimal pricing lookup, satisfied by ModelRegistry. Cost is $/million tokens. */
/** Minimal pricing lookup, satisfied by ModelRuntime. Cost is $/million tokens. */
export interface ModelPriceSource {
find(provider: string, modelId: string): { cost: { cacheRead: number } } | undefined;
getModel(provider: string, modelId: string): { cost: { cacheRead: number } } | undefined;
}
/** The last request seen by the scan; everything in its prompt should be cached. */
@@ -79,7 +79,7 @@ function detectMiss(
const readPerToken =
usage.cacheRead > 0
? usage.cost.cacheRead / usage.cacheRead
: (models.find(message.provider, message.model)?.cost.cacheRead ?? 0) / 1_000_000;
: (models.getModel(message.provider, message.model)?.cost.cacheRead ?? 0) / 1_000_000;
return {
missedTokens,
@@ -10,6 +10,7 @@ import { fileURLToPath } from "node:url";
import * as _bundledPiAgentCore from "@earendil-works/pi-agent-core";
import * as _bundledPiAiCompat from "@earendil-works/pi-ai/compat";
import * as _bundledPiAiOauth from "@earendil-works/pi-ai/oauth";
import * as _bundledPiAiProviders from "@earendil-works/pi-ai/providers/all";
import type { KeyId } from "@earendil-works/pi-tui";
import * as _bundledPiTui from "@earendil-works/pi-tui";
import { createJiti } from "jiti/static";
@@ -58,12 +59,14 @@ const VIRTUAL_MODULES: Record<string, unknown> = {
"@earendil-works/pi-ai": _bundledPiAiCompat,
"@earendil-works/pi-ai/compat": _bundledPiAiCompat,
"@earendil-works/pi-ai/oauth": _bundledPiAiOauth,
"@earendil-works/pi-ai/providers/all": _bundledPiAiProviders,
"@earendil-works/pi-coding-agent": _bundledPiCodingAgent,
"@mariozechner/pi-agent-core": _bundledPiAgentCore,
"@mariozechner/pi-tui": _bundledPiTui,
"@mariozechner/pi-ai": _bundledPiAiCompat,
"@mariozechner/pi-ai/compat": _bundledPiAiCompat,
"@mariozechner/pi-ai/oauth": _bundledPiAiOauth,
"@mariozechner/pi-ai/providers/all": _bundledPiAiProviders,
"@mariozechner/pi-coding-agent": _bundledPiCodingAgent,
};
@@ -102,20 +105,26 @@ function getAliases(): Record<string, string> {
// global API keep working at runtime until compat is removed.
const piAiCompatEntry = resolveWorkspaceOrImport("ai/dist/compat.js", "@earendil-works/pi-ai/compat");
const piAiOauthEntry = resolveWorkspaceOrImport("ai/dist/oauth.js", "@earendil-works/pi-ai/oauth");
const piAiProvidersEntry = resolveWorkspaceOrImport(
"ai/dist/providers/all.js",
"@earendil-works/pi-ai/providers/all",
);
_aliases = {
"@earendil-works/pi-coding-agent": piCodingAgentEntry,
"@earendil-works/pi-agent-core": piAgentCoreEntry,
"@earendil-works/pi-tui": piTuiEntry,
"@earendil-works/pi-ai": piAiCompatEntry,
"@earendil-works/pi-ai/providers/all": piAiProvidersEntry,
"@earendil-works/pi-ai/compat": piAiCompatEntry,
"@earendil-works/pi-ai/oauth": piAiOauthEntry,
"@earendil-works/pi-ai": piAiCompatEntry,
"@mariozechner/pi-coding-agent": piCodingAgentEntry,
"@mariozechner/pi-agent-core": piAgentCoreEntry,
"@mariozechner/pi-tui": piTuiEntry,
"@mariozechner/pi-ai": piAiCompatEntry,
"@mariozechner/pi-ai/providers/all": piAiProvidersEntry,
"@mariozechner/pi-ai/compat": piAiCompatEntry,
"@mariozechner/pi-ai/oauth": piAiOauthEntry,
"@mariozechner/pi-ai": piAiCompatEntry,
typebox: typeboxEntry,
"typebox/compile": typeboxCompileEntry,
"typebox/value": typeboxValueEntry,
@@ -602,6 +602,10 @@ export class ExtensionRunner {
});
}
getModelRegistry(): ModelRegistry {
return this.modelRegistry;
}
getRegisteredCommands(): ResolvedCommand[] {
this.commandDiagnostics = [];
return this.resolveRegisteredCommands();
@@ -1424,14 +1424,14 @@ export interface ProviderConfig {
oauth?: {
/** Display name for the provider in login UI. */
name: string;
/** @deprecated Retained for source compatibility; canonical auth flows ignore it. */
usesCallbackServer?: boolean;
/** Run the login flow, return credentials to persist. */
login(callbacks: OAuthLoginCallbacks): Promise<OAuthCredentials>;
/** Refresh expired credentials, return updated credentials to persist. */
refreshToken(credentials: OAuthCredentials): Promise<OAuthCredentials>;
/** Convert credentials to API key string for the provider. */
getApiKey(credentials: OAuthCredentials): string;
/** Optional: modify models for this provider (e.g., update baseUrl based on credentials). */
modifyModels?(models: Model<Api>[], credentials: OAuthCredentials): Model<Api>[];
};
}
@@ -0,0 +1,277 @@
/** Immutable, credential-blind models.json snapshot. */
import { readFile } from "node:fs/promises";
import { type Static, Type } from "typebox";
import { Compile } from "typebox/compile";
import type { TLocalizedValidationError } from "typebox/error";
import { stripJsonComments } from "../utils/json.ts";
import { normalizePath } from "../utils/paths.ts";
const PercentileCutoffsSchema = Type.Object({
p50: Type.Optional(Type.Number()),
p75: Type.Optional(Type.Number()),
p90: Type.Optional(Type.Number()),
p99: Type.Optional(Type.Number()),
});
const OpenRouterRoutingSchema = Type.Object({
allow_fallbacks: Type.Optional(Type.Boolean()),
require_parameters: Type.Optional(Type.Boolean()),
data_collection: Type.Optional(Type.Union([Type.Literal("deny"), Type.Literal("allow")])),
zdr: Type.Optional(Type.Boolean()),
enforce_distillable_text: Type.Optional(Type.Boolean()),
order: Type.Optional(Type.Array(Type.String())),
only: Type.Optional(Type.Array(Type.String())),
ignore: Type.Optional(Type.Array(Type.String())),
quantizations: Type.Optional(Type.Array(Type.String())),
sort: Type.Optional(
Type.Union([
Type.String(),
Type.Object({
by: Type.Optional(Type.String()),
partition: Type.Optional(Type.Union([Type.String(), Type.Null()])),
}),
]),
),
max_price: Type.Optional(
Type.Object({
prompt: Type.Optional(Type.Union([Type.Number(), Type.String()])),
completion: Type.Optional(Type.Union([Type.Number(), Type.String()])),
image: Type.Optional(Type.Union([Type.Number(), Type.String()])),
audio: Type.Optional(Type.Union([Type.Number(), Type.String()])),
request: Type.Optional(Type.Union([Type.Number(), Type.String()])),
}),
),
preferred_min_throughput: Type.Optional(Type.Union([Type.Number(), PercentileCutoffsSchema])),
preferred_max_latency: Type.Optional(Type.Union([Type.Number(), PercentileCutoffsSchema])),
});
const VercelGatewayRoutingSchema = Type.Object({
only: Type.Optional(Type.Array(Type.String())),
order: Type.Optional(Type.Array(Type.String())),
});
const ThinkingLevelMapValueSchema = Type.Union([Type.String(), Type.Null()]);
const ThinkingLevelMapSchema = Type.Object({
off: Type.Optional(ThinkingLevelMapValueSchema),
minimal: Type.Optional(ThinkingLevelMapValueSchema),
low: Type.Optional(ThinkingLevelMapValueSchema),
medium: Type.Optional(ThinkingLevelMapValueSchema),
high: Type.Optional(ThinkingLevelMapValueSchema),
xhigh: Type.Optional(ThinkingLevelMapValueSchema),
max: Type.Optional(ThinkingLevelMapValueSchema),
});
const ChatTemplateKwargScalarSchema = Type.Union([Type.String(), Type.Number(), Type.Boolean(), Type.Null()]);
const ChatTemplateKwargVariableSchema = Type.Object({
$var: Type.Union([Type.Literal("thinking.enabled"), Type.Literal("thinking.effort")]),
omitWhenOff: Type.Optional(Type.Boolean()),
});
const ChatTemplateKwargSchema = Type.Union([ChatTemplateKwargScalarSchema, ChatTemplateKwargVariableSchema]);
const OpenAICompletionsCompatSchema = Type.Object({
supportsStore: Type.Optional(Type.Boolean()),
supportsDeveloperRole: Type.Optional(Type.Boolean()),
supportsReasoningEffort: Type.Optional(Type.Boolean()),
supportsUsageInStreaming: Type.Optional(Type.Boolean()),
maxTokensField: Type.Optional(Type.Union([Type.Literal("max_completion_tokens"), Type.Literal("max_tokens")])),
requiresToolResultName: Type.Optional(Type.Boolean()),
requiresAssistantAfterToolResult: Type.Optional(Type.Boolean()),
requiresThinkingAsText: Type.Optional(Type.Boolean()),
requiresReasoningContentOnAssistantMessages: Type.Optional(Type.Boolean()),
thinkingFormat: Type.Optional(
Type.Union([
Type.Literal("openai"),
Type.Literal("openrouter"),
Type.Literal("together"),
Type.Literal("deepseek"),
Type.Literal("zai"),
Type.Literal("qwen"),
Type.Literal("chat-template"),
Type.Literal("qwen-chat-template"),
Type.Literal("string-thinking"),
Type.Literal("ant-ling"),
]),
),
chatTemplateKwargs: Type.Optional(Type.Record(Type.String(), ChatTemplateKwargSchema)),
cacheControlFormat: Type.Optional(Type.Literal("anthropic")),
openRouterRouting: Type.Optional(OpenRouterRoutingSchema),
vercelGatewayRouting: Type.Optional(VercelGatewayRoutingSchema),
supportsStrictMode: Type.Optional(Type.Boolean()),
supportsLongCacheRetention: Type.Optional(Type.Boolean()),
});
const OpenAIResponsesCompatSchema = Type.Object({
supportsDeveloperRole: Type.Optional(Type.Boolean()),
sendSessionIdHeader: Type.Optional(Type.Boolean()),
supportsLongCacheRetention: Type.Optional(Type.Boolean()),
});
const AnthropicMessagesCompatSchema = Type.Object({
supportsEagerToolInputStreaming: Type.Optional(Type.Boolean()),
supportsLongCacheRetention: Type.Optional(Type.Boolean()),
sendSessionAffinityHeaders: Type.Optional(Type.Boolean()),
supportsCacheControlOnTools: Type.Optional(Type.Boolean()),
forceAdaptiveThinking: Type.Optional(Type.Boolean()),
});
const ProviderCompatSchema = Type.Union([
OpenAICompletionsCompatSchema,
OpenAIResponsesCompatSchema,
AnthropicMessagesCompatSchema,
]);
const ModelCostRatesSchema = {
input: Type.Number(),
output: Type.Number(),
cacheRead: Type.Number(),
cacheWrite: Type.Number(),
};
const ModelCostTierSchema = Type.Object({
inputTokensAbove: Type.Number(),
...ModelCostRatesSchema,
});
const ModelCostSchema = Type.Object({
...ModelCostRatesSchema,
tiers: Type.Optional(Type.Array(ModelCostTierSchema)),
});
const ModelDefinitionSchema = Type.Object({
id: Type.String({ minLength: 1 }),
name: Type.Optional(Type.String({ minLength: 1 })),
api: Type.Optional(Type.String({ minLength: 1 })),
baseUrl: Type.Optional(Type.String({ minLength: 1 })),
reasoning: Type.Optional(Type.Boolean()),
thinkingLevelMap: Type.Optional(ThinkingLevelMapSchema),
input: Type.Optional(Type.Array(Type.Union([Type.Literal("text"), Type.Literal("image")]))),
cost: Type.Optional(ModelCostSchema),
contextWindow: Type.Optional(Type.Number()),
maxTokens: Type.Optional(Type.Number()),
headers: Type.Optional(Type.Record(Type.String(), Type.String())),
compat: Type.Optional(ProviderCompatSchema),
});
const ModelOverrideSchema = Type.Object({
name: Type.Optional(Type.String({ minLength: 1 })),
reasoning: Type.Optional(Type.Boolean()),
thinkingLevelMap: Type.Optional(ThinkingLevelMapSchema),
input: Type.Optional(Type.Array(Type.Union([Type.Literal("text"), Type.Literal("image")]))),
cost: Type.Optional(
Type.Object({
input: Type.Optional(Type.Number()),
output: Type.Optional(Type.Number()),
cacheRead: Type.Optional(Type.Number()),
cacheWrite: Type.Optional(Type.Number()),
tiers: Type.Optional(Type.Array(ModelCostTierSchema)),
}),
),
contextWindow: Type.Optional(Type.Number()),
maxTokens: Type.Optional(Type.Number()),
headers: Type.Optional(Type.Record(Type.String(), Type.String())),
compat: Type.Optional(ProviderCompatSchema),
});
const ProviderConfigSchema = Type.Object({
name: Type.Optional(Type.String({ minLength: 1 })),
baseUrl: Type.Optional(Type.String({ minLength: 1 })),
apiKey: Type.Optional(Type.String({ minLength: 1 })),
api: Type.Optional(Type.String({ minLength: 1 })),
headers: Type.Optional(Type.Record(Type.String(), Type.String())),
compat: Type.Optional(ProviderCompatSchema),
authHeader: Type.Optional(Type.Boolean()),
models: Type.Optional(Type.Array(ModelDefinitionSchema)),
modelOverrides: Type.Optional(Type.Record(Type.String(), ModelOverrideSchema)),
});
const ModelsConfigSchema = Type.Object({
providers: Type.Record(Type.String(), ProviderConfigSchema),
});
const validateModelsConfig = Compile(ModelsConfigSchema);
export type ModelsJsonModel = Static<typeof ModelDefinitionSchema>;
export type ModelsJsonModelOverride = Static<typeof ModelOverrideSchema>;
export type ModelsJsonProvider = Static<typeof ProviderConfigSchema>;
type ModelsJson = Static<typeof ModelsConfigSchema>;
function formatValidationPath(error: TLocalizedValidationError): string {
if (error.keyword === "required") {
const requiredProperties = (error.params as { requiredProperties?: string[] }).requiredProperties;
const requiredProperty = requiredProperties?.[0];
if (requiredProperty) {
const basePath = error.instancePath.replace(/^\//, "").replace(/\//g, ".");
return basePath ? `${basePath}.${requiredProperty}` : requiredProperty;
}
}
const path = error.instancePath.replace(/^\//, "").replace(/\//g, ".");
return path || "root";
}
function deepFreeze<T>(value: T): T {
if (typeof value !== "object" || value === null || Object.isFrozen(value)) return value;
for (const child of Object.values(value)) deepFreeze(child);
return Object.freeze(value);
}
/** One immutable load of models.json. */
export class ModelConfig {
private readonly providers: ReadonlyMap<string, ModelsJsonProvider>;
private readonly error: string | undefined;
private constructor(providers: ReadonlyMap<string, ModelsJsonProvider>, error?: string) {
this.providers = providers;
this.error = error;
}
static async load(modelsJsonPath: string | undefined): Promise<ModelConfig> {
if (!modelsJsonPath) return new ModelConfig(new Map());
const path = normalizePath(modelsJsonPath);
let content: string;
try {
content = await readFile(path, "utf-8");
} catch (error) {
if ((error as NodeJS.ErrnoException).code === "ENOENT") return new ModelConfig(new Map());
return new ModelConfig(
new Map(),
`Failed to load models.json: ${error instanceof Error ? error.message : error}\n\nFile: ${path}`,
);
}
let parsed: unknown;
try {
parsed = JSON.parse(stripJsonComments(content));
} catch (error) {
return new ModelConfig(
new Map(),
`Failed to parse models.json: ${error instanceof Error ? error.message : error}\n\nFile: ${path}`,
);
}
if (!validateModelsConfig.Check(parsed)) {
const errors =
validateModelsConfig
.Errors(parsed)
.map((error) => ` - ${formatValidationPath(error)}: ${error.message}`)
.join("\n") || "Unknown schema error";
return new ModelConfig(new Map(), `Invalid models.json schema:\n${errors}\n\nFile: ${path}`);
}
const config = parsed as ModelsJson;
const providers = new Map<string, ModelsJsonProvider>();
for (const [providerId, provider] of Object.entries(config.providers)) {
providers.set(providerId, deepFreeze(structuredClone(provider)));
}
return new ModelConfig(providers);
}
getProvider(providerId: string): ModelsJsonProvider | undefined {
return this.providers.get(providerId);
}
getProviderIds(): readonly string[] {
return [...this.providers.keys()];
}
getError(): string | undefined {
return this.error;
}
}
File diff suppressed because it is too large Load Diff
@@ -8,7 +8,7 @@ import chalk from "chalk";
import { minimatch } from "minimatch";
import { isValidThinkingLevel } from "../cli/args.ts";
import { DEFAULT_THINKING_LEVEL } from "./defaults.ts";
import type { ModelRegistry } from "./model-registry.ts";
import type { ModelRuntime } from "./model-runtime.ts";
/** Default model IDs for each known provider */
export const defaultModelPerProvider: Record<KnownProvider, string> = {
@@ -268,9 +268,9 @@ export interface ResolveModelScopeResult {
export async function resolveModelScopeWithDiagnostics(
patterns: string[],
modelRegistry: ModelRegistry,
modelRuntime: ModelRuntime,
): Promise<ResolveModelScopeResult> {
const availableModels = await modelRegistry.getAvailable();
const availableModels = [...(await modelRuntime.getAvailable())];
const scopedModels: ScopedModel[] = [];
const diagnostics: ModelScopeDiagnostic[] = [];
@@ -330,8 +330,8 @@ export async function resolveModelScopeWithDiagnostics(
return { scopedModels, diagnostics };
}
export async function resolveModelScope(patterns: string[], modelRegistry: ModelRegistry): Promise<ScopedModel[]> {
const { scopedModels, diagnostics } = await resolveModelScopeWithDiagnostics(patterns, modelRegistry);
export async function resolveModelScope(patterns: string[], modelRuntime: ModelRuntime): Promise<ScopedModel[]> {
const { scopedModels, diagnostics } = await resolveModelScopeWithDiagnostics(patterns, modelRuntime);
for (const diagnostic of diagnostics) {
console.warn(chalk.yellow(`Warning: ${diagnostic.message}`));
}
@@ -364,9 +364,9 @@ export function resolveCliModel(options: {
cliProvider?: string;
cliModel?: string;
cliThinking?: ThinkingLevel;
modelRegistry: ModelRegistry;
modelRuntime: ModelRuntime;
}): ResolveCliModelResult {
const { cliProvider, cliModel, cliThinking, modelRegistry } = options;
const { cliProvider, cliModel, cliThinking, modelRuntime } = options;
if (!cliModel) {
return { model: undefined, warning: undefined, error: undefined };
@@ -374,7 +374,7 @@ export function resolveCliModel(options: {
// Important: use *all* models here, not just models with pre-configured auth.
// This allows "--api-key" to be used for first-time setup.
const availableModels = modelRegistry.getAll();
const availableModels = [...modelRuntime.getModels()];
if (availableModels.length === 0) {
return {
model: undefined,
@@ -454,8 +454,8 @@ export function resolveCliModel(options: {
const rawExactMatches = availableModels.filter(
(m) => m.id.toLowerCase() === cliModel.toLowerCase() && !modelsAreEqual(m, model),
);
if (rawExactMatches.length > 0 && !modelRegistry.hasConfiguredAuth(model)) {
const authenticatedRawMatches = rawExactMatches.filter((m) => modelRegistry.hasConfiguredAuth(m));
if (rawExactMatches.length > 0 && !modelRuntime.hasConfiguredAuth(model.provider)) {
const authenticatedRawMatches = rawExactMatches.filter((m) => modelRuntime.hasConfiguredAuth(m.provider));
if (authenticatedRawMatches.length === 1) {
return {
model: authenticatedRawMatches[0],
@@ -555,7 +555,7 @@ export async function findInitialModel(options: {
defaultProvider?: string;
defaultModelId?: string;
defaultThinkingLevel?: ThinkingLevel;
modelRegistry: ModelRegistry;
modelRuntime: ModelRuntime;
}): Promise<InitialModelResult> {
const {
cliProvider,
@@ -565,7 +565,7 @@ export async function findInitialModel(options: {
defaultProvider,
defaultModelId,
defaultThinkingLevel,
modelRegistry,
modelRuntime,
} = options;
let model: Model<Api> | undefined;
@@ -576,7 +576,7 @@ export async function findInitialModel(options: {
const resolved = resolveCliModel({
cliProvider,
cliModel,
modelRegistry,
modelRuntime,
});
if (resolved.error) {
console.error(chalk.red(resolved.error));
@@ -598,8 +598,8 @@ export async function findInitialModel(options: {
// 3. Try saved default from settings if auth is configured.
if (defaultProvider && defaultModelId) {
const found = modelRegistry.find(defaultProvider, defaultModelId);
if (found && modelRegistry.hasConfiguredAuth(found)) {
const found = modelRuntime.getModel(defaultProvider, defaultModelId);
if (found && modelRuntime.hasConfiguredAuth(found.provider)) {
model = found;
if (defaultThinkingLevel) {
thinkingLevel = defaultThinkingLevel;
@@ -609,7 +609,7 @@ export async function findInitialModel(options: {
}
// 4. Try first available model with valid API key
const availableModels = await modelRegistry.getAvailable();
const availableModels = [...(await modelRuntime.getAvailable())];
if (availableModels.length > 0) {
// Try to find a default model from known providers
@@ -637,12 +637,12 @@ export async function restoreModelFromSession(
savedModelId: string,
currentModel: Model<Api> | undefined,
shouldPrintMessages: boolean,
modelRegistry: ModelRegistry,
modelRuntime: ModelRuntime,
): Promise<{ model: Model<Api> | undefined; fallbackMessage: string | undefined }> {
const restoredModel = modelRegistry.find(savedProvider, savedModelId);
const restoredModel = modelRuntime.getModel(savedProvider, savedModelId);
// Check if restored model exists and still has auth configured
const hasConfiguredAuth = restoredModel ? modelRegistry.hasConfiguredAuth(restoredModel) : false;
const hasConfiguredAuth = restoredModel ? modelRuntime.hasConfiguredAuth(restoredModel.provider) : false;
if (restoredModel && hasConfiguredAuth) {
if (shouldPrintMessages) {
@@ -670,7 +670,7 @@ export async function restoreModelFromSession(
}
// Try to find any available model
const availableModels = await modelRegistry.getAvailable();
const availableModels = [...(await modelRuntime.getAvailable())];
if (availableModels.length > 0) {
// Try to find a default model from known providers
@@ -0,0 +1,489 @@
import { join } from "node:path";
import {
type Api,
type ApiStreamOptions,
type AssistantMessage,
type AssistantMessageEventStream,
type AuthCheck,
type AuthInteraction,
type AuthResult,
type AuthType,
type Context,
type Credential,
type CredentialInfo,
type CredentialStore,
createModels,
lazyStream,
type Model,
type Models,
type ModelsApiStreamOptions,
ModelsError,
type ModelsSimpleStreamOptions,
type ModelsStreamTransforms,
type MutableModels,
type Provider,
type ProviderHeaders,
type SimpleStreamOptions,
type StreamOptions,
} from "@earendil-works/pi-ai";
import { builtinProviders } from "@earendil-works/pi-ai/providers/all";
import { getAgentDir } from "../config.ts";
import { AuthStorage as DefaultAuthStorage } from "./auth-storage.ts";
import { ModelConfig } from "./model-config.ts";
import {
type AuthStatus,
type CompatibilityRequestConfig,
composeModelProvider,
configuredRequestAuthStatus,
type ProviderConfigInput,
resolveCompatibilityRequestConfig,
resolveConfiguredModelHeaders,
validateExtensionProvider,
} from "./provider-composer.ts";
import { RuntimeCredentials } from "./runtime-credentials.ts";
interface ModelRuntimeSnapshot {
all: readonly Model<Api>[];
available: readonly Model<Api>[];
configuredProviders: ReadonlySet<string>;
storedProviders: ReadonlySet<string>;
auth: ReadonlyMap<string, AuthCheck | undefined>;
}
export interface CreateModelRuntimeOptions {
/** Credential storage. Defaults to the file at authPath. */
credentials?: CredentialStore;
authPath?: string;
modelsPath?: string | null;
}
export interface ModelRuntimeAuthOverrides {
apiKey?: string;
env?: Record<string, string>;
}
function mergeHeaders(
base: ProviderHeaders | undefined,
override: ProviderHeaders | undefined,
): ProviderHeaders | undefined {
if (!base && !override) return undefined;
const merged = { ...base };
for (const [name, value] of Object.entries(override ?? {})) {
const lowerName = name.toLowerCase();
for (const existingName of Object.keys(merged)) {
if (existingName.toLowerCase() === lowerName) delete merged[existingName];
}
merged[name] = value;
}
return merged;
}
/** Configured pi-ai Models collection used by coding-agent and SDK consumers. */
export class ModelRuntime implements Models {
private readonly models: MutableModels;
private readonly credentials: RuntimeCredentials;
private readonly builtins: ReadonlyMap<string, Provider>;
private readonly extensionProviders = new Map<string, ProviderConfigInput>();
private readonly compositionErrors = new Map<string, string>();
private readonly modelsPath: string | undefined;
private config: ModelConfig;
private snapshot: ModelRuntimeSnapshot = {
all: [],
available: [],
configuredProviders: new Set(),
storedProviders: new Set(),
auth: new Map(),
};
private availabilityRefresh: Promise<void> | undefined;
private availabilityError: string | undefined;
private constructor(
credentials: RuntimeCredentials,
config: ModelConfig,
modelsPath: string | undefined,
providers: readonly Provider[],
) {
this.credentials = credentials;
this.config = config;
this.modelsPath = modelsPath;
this.builtins = new Map(providers.map((provider) => [provider.id, provider]));
this.models = createModels({ credentials });
this.rebuildProviders();
}
static async create(options: CreateModelRuntimeOptions = {}): Promise<ModelRuntime> {
const credentials = new RuntimeCredentials(options.credentials ?? DefaultAuthStorage.create(options.authPath));
const modelsPath =
options.modelsPath === null ? undefined : (options.modelsPath ?? join(getAgentDir(), "models.json"));
const config = await ModelConfig.load(modelsPath);
const runtime = new ModelRuntime(credentials, config, modelsPath, builtinProviders());
await runtime.refreshAvailability();
return runtime;
}
private providerIds(): Set<string> {
return new Set([...this.builtins.keys(), ...this.config.getProviderIds(), ...this.extensionProviders.keys()]);
}
private recomposeProvider(providerId: string): void {
const base = this.builtins.get(providerId);
const extension = this.extensionProviders.get(providerId);
if (!base && !this.config.getProvider(providerId) && !extension) {
this.models.deleteProvider(providerId);
this.compositionErrors.delete(providerId);
return;
}
if (base && !this.config.getProvider(providerId) && !extension) {
// No overlays: use the builtin untouched so its auth/login/stream behavior is exact.
this.models.setProvider(base);
this.compositionErrors.delete(providerId);
return;
}
try {
this.models.setProvider(composeModelProvider(providerId, base, this.config, extension));
this.compositionErrors.delete(providerId);
} catch (error) {
this.compositionErrors.set(providerId, error instanceof Error ? error.message : String(error));
if (base) this.models.setProvider(base);
else this.models.deleteProvider(providerId);
}
}
private rebuildProviders(): void {
this.models.clearProviders();
this.compositionErrors.clear();
for (const providerId of this.providerIds()) this.recomposeProvider(providerId);
this.updateModelSnapshot();
}
private updateModelSnapshot(): void {
const all = [...this.models.getModels()];
this.snapshot = {
...this.snapshot,
all,
available: all.filter((model) => this.snapshot.configuredProviders.has(model.provider)),
};
}
private async runAvailabilityRefresh(): Promise<void> {
const providers = this.models.getProviders();
const [available, checks, credentials] = await Promise.all([
this.models.getAvailable(),
Promise.all(
providers.map(
async (provider): Promise<[string, AuthCheck | undefined]> => [
provider.id,
await this.models.checkAuth(provider.id),
],
),
),
this.credentials.list(),
]);
const auth = new Map(checks);
const configuredProviders = new Set(
checks
.filter((entry): entry is [string, AuthCheck] => entry[1] !== undefined)
.map(([providerId]) => providerId),
);
this.snapshot = {
all: [...this.models.getModels()],
available: [...available],
configuredProviders,
storedProviders: new Set(credentials.map((entry) => entry.providerId)),
auth,
};
this.availabilityError = undefined;
}
private queueAvailabilityRefresh(after: Promise<void> | undefined): Promise<void> {
const refresh = (after ?? Promise.resolve()).catch(() => {}).then(() => this.runAvailabilityRefresh());
const recorded = refresh.catch((error) => {
this.availabilityError = error instanceof Error ? error.message : String(error);
throw error;
});
const tracked = recorded.finally(() => {
if (this.availabilityRefresh === tracked) this.availabilityRefresh = undefined;
});
this.availabilityRefresh = tracked;
return tracked;
}
/** Coalesce concurrent readers onto the pending refresh. */
private refreshAvailability(): Promise<void> {
return this.availabilityRefresh ?? this.queueAvailabilityRefresh(undefined);
}
/** Mutations must not observe an in-flight refresh started before them. */
private forceRefreshAvailability(): Promise<void> {
return this.queueAvailabilityRefresh(this.availabilityRefresh);
}
getProviders(): readonly Provider[] {
return this.models.getProviders();
}
getProvider(providerId: string): Provider | undefined {
return this.models.getProvider(providerId);
}
getModels(providerId?: string): readonly Model<Api>[] {
return this.models.getModels(providerId);
}
getModel(providerId: string, modelId: string): Model<Api> | undefined {
return this.models.getModel(providerId, modelId);
}
async checkAuth(providerId: string): Promise<AuthCheck | undefined> {
return this.models.checkAuth(providerId);
}
async getAvailable(providerId?: string): Promise<readonly Model<Api>[]> {
if (providerId) {
if (this.availabilityRefresh) {
await this.availabilityRefresh;
return this.snapshot.available.filter((model) => model.provider === providerId);
}
try {
return await this.models.getAvailable(providerId);
} catch (error) {
this.availabilityError = error instanceof Error ? error.message : String(error);
throw error;
}
}
await this.refreshAvailability();
return this.snapshot.available;
}
getAvailableSnapshot(): readonly Model<Api>[] {
return this.snapshot.available;
}
getError(): string | undefined {
const errors: string[] = [];
const configError = this.config.getError();
if (configError) errors.push(configError);
for (const [providerId, error] of this.compositionErrors) {
errors.push(`Provider "${providerId}": ${error}`);
}
if (this.availabilityError) errors.push(`Availability refresh: ${this.availabilityError}`);
return errors.length > 0 ? errors.join("\n\n") : undefined;
}
getRegisteredProviderConfig(providerId: string): ProviderConfigInput | undefined {
return this.extensionProviders.get(providerId);
}
getRegisteredProviderIds(): readonly string[] {
return [...this.extensionProviders.keys()];
}
/** @internal Compatibility fallback for ModelRegistry when provider auth is unconfigured. */
getCompatibilityRequestConfig(model: Model<Api>): CompatibilityRequestConfig {
return resolveCompatibilityRequestConfig(
model,
this.config.getProvider(model.provider),
this.extensionProviders.get(model.provider),
);
}
isUsingOAuth(providerId: string): boolean {
return this.snapshot.auth.get(providerId)?.type === "oauth";
}
hasConfiguredAuth(providerId: string): boolean {
return this.snapshot.configuredProviders.has(providerId);
}
getAuth(providerId: string, overrides?: ModelRuntimeAuthOverrides): Promise<AuthResult | undefined>;
getAuth(model: Model<Api>, overrides?: ModelRuntimeAuthOverrides): Promise<AuthResult | undefined>;
async getAuth(
providerOrModel: string | Model<Api>,
overrides: ModelRuntimeAuthOverrides = {},
): Promise<AuthResult | undefined> {
if (typeof providerOrModel === "string") return this.models.getAuth(providerOrModel, overrides);
const resolution = await this.models.getAuth(providerOrModel, overrides);
if (!resolution) return undefined;
const configuredHeaders = resolveConfiguredModelHeaders(
providerOrModel,
this.config.getProvider(providerOrModel.provider),
this.extensionProviders.get(providerOrModel.provider),
{ ...(resolution.env ?? {}), ...(overrides.env ?? {}) },
);
return {
...resolution,
auth: {
...resolution.auth,
headers: mergeHeaders(resolution.auth.headers, configuredHeaders),
},
};
}
setRuntimeApiKey(providerId: string, apiKey: string): void {
this.credentials.setRuntimeApiKey(providerId, apiKey);
const auth = new Map(this.snapshot.auth).set(providerId, { type: "api_key", source: "runtime API key" });
const configuredProviders = new Set(this.snapshot.configuredProviders).add(providerId);
const storedProviders = new Set(this.snapshot.storedProviders).add(providerId);
this.snapshot = {
...this.snapshot,
auth,
configuredProviders,
storedProviders,
available: this.snapshot.all.filter((model) => configuredProviders.has(model.provider)),
};
void this.forceRefreshAvailability().catch(() => {});
}
removeRuntimeApiKey(providerId: string): void {
this.credentials.removeRuntimeApiKey(providerId);
void this.forceRefreshAvailability().catch(() => {});
}
listCredentials(): Promise<readonly CredentialInfo[]> {
return this.credentials.list();
}
getProviderAuthStatus(providerId: string): AuthStatus {
if (this.credentials.hasRuntimeApiKey(providerId)) return { configured: true, source: "runtime" };
if (this.snapshot.storedProviders.has(providerId)) return { configured: true, source: "stored" };
const configured = configuredRequestAuthStatus(
this.config.getProvider(providerId),
this.extensionProviders.get(providerId),
);
if (configured) return configured;
const check = this.snapshot.auth.get(providerId);
return check ? { configured: true, source: "environment", label: check.source } : { configured: false };
}
private async prepareRequest(
model: Model<Api>,
options: (StreamOptions & ModelsStreamTransforms) | undefined,
): Promise<{ provider: Provider; model: Model<Api>; options: StreamOptions }> {
const provider = this.models.getProvider(model.provider);
if (!provider) throw new ModelsError("provider", `Unknown provider: ${model.provider}`);
const resolution = await this.getAuth(model, { apiKey: options?.apiKey, env: options?.env });
if (!resolution) throw new ModelsError("auth", `Provider is not configured: ${model.provider}`);
const { transformHeaders, ...providerOptions } = options ?? {};
let headers = mergeHeaders(resolution.auth.headers, providerOptions.headers);
if (transformHeaders) headers = await transformHeaders(headers ?? {});
const env =
resolution.env || providerOptions.env
? { ...(resolution.env ?? {}), ...(providerOptions.env ?? {}) }
: undefined;
return {
provider,
model: resolution.auth.baseUrl ? { ...model, baseUrl: resolution.auth.baseUrl } : model,
options: {
...providerOptions,
apiKey: providerOptions.apiKey ?? resolution.auth.apiKey,
headers,
env,
},
};
}
stream<TApi extends Api>(
model: Model<TApi>,
context: Context,
options?: ModelsApiStreamOptions<TApi>,
): AssistantMessageEventStream {
return lazyStream(model, async () => {
const prepared = await this.prepareRequest(
model,
options as (StreamOptions & ModelsStreamTransforms) | undefined,
);
return prepared.provider.stream(
prepared.model as Model<TApi>,
context,
prepared.options as ApiStreamOptions<TApi>,
);
});
}
complete<TApi extends Api>(
model: Model<TApi>,
context: Context,
options?: ModelsApiStreamOptions<TApi>,
): Promise<AssistantMessage> {
return this.stream(model, context, options).result();
}
streamSimple(model: Model<Api>, context: Context, options?: ModelsSimpleStreamOptions): AssistantMessageEventStream {
return lazyStream(model, async () => {
const prepared = await this.prepareRequest(model, options);
return prepared.provider.streamSimple(prepared.model, context, prepared.options as SimpleStreamOptions);
});
}
completeSimple(model: Model<Api>, context: Context, options?: ModelsSimpleStreamOptions): Promise<AssistantMessage> {
return this.streamSimple(model, context, options).result();
}
async login(providerId: string, type: AuthType, interaction: AuthInteraction): Promise<Credential> {
const credential = await this.models.login(providerId, type, interaction);
await this.forceRefreshAvailability();
return credential;
}
async logout(providerId: string): Promise<void> {
await this.models.logout(providerId);
await this.forceRefreshAvailability();
}
async reloadConfig(): Promise<void> {
this.config = await ModelConfig.load(this.modelsPath);
this.rebuildProviders();
await this.forceRefreshAvailability();
}
async refresh(providerId?: string): Promise<void> {
await this.models.refresh(providerId);
this.updateModelSnapshot();
await this.forceRefreshAvailability();
}
registerProvider(providerId: string, config: ProviderConfigInput): void {
// Validate the incoming registration on its own, like the legacy registry:
// a broken re-registration must throw without touching the stored config.
validateExtensionProvider(providerId, this.builtins.get(providerId), this.config.getProvider(providerId), config);
// Re-registration merges defined values over the previous registration and
// preserves undefined ones, matching the legacy ModelRegistry contract.
const previous = this.extensionProviders.get(providerId);
const effective: ProviderConfigInput = { ...previous };
for (const [key, value] of Object.entries(config)) {
if (value !== undefined) (effective as Record<string, unknown>)[key] = value;
}
this.extensionProviders.set(providerId, effective);
this.recomposeProvider(providerId);
this.updateModelSnapshot();
if (
this.snapshot.storedProviders.has(providerId) ||
configuredRequestAuthStatus(this.config.getProvider(providerId), effective)?.configured
) {
const configuredProviders = new Set(this.snapshot.configuredProviders).add(providerId);
const auth = new Map(this.snapshot.auth);
// Provisional entry until the async refresh lands; never clobber a real check result.
if (!auth.get(providerId)) {
auth.set(providerId, {
type: effective.oauth && !effective.apiKey ? "oauth" : "api_key",
source: "configured provider",
});
}
this.snapshot = {
...this.snapshot,
auth,
configuredProviders,
available: this.snapshot.all.filter((model) => configuredProviders.has(model.provider)),
};
}
void this.forceRefreshAvailability().catch(() => {});
}
unregisterProvider(providerId: string): void {
this.extensionProviders.delete(providerId);
this.recomposeProvider(providerId);
this.updateModelSnapshot();
void this.forceRefreshAvailability().catch(() => {});
}
}
@@ -0,0 +1,513 @@
import {
type Api,
type ApiKeyAuth,
type AssistantMessageEventStream,
type AuthContext,
type AuthInteraction,
type AuthResult,
type Context,
type Credential,
lazyStream,
type Model,
type ModelAuth,
type OAuthAuth,
type OAuthCredentials,
type OAuthLoginCallbacks,
type Provider,
type ProviderHeaders,
type SimpleStreamOptions,
type StreamOptions,
} from "@earendil-works/pi-ai";
import { getApiProvider } from "@earendil-works/pi-ai/compat";
import type { ModelConfig, ModelsJsonModel, ModelsJsonModelOverride, ModelsJsonProvider } from "./model-config.ts";
import {
clearConfigValueCache,
getConfigValueEnvVarNames,
isCommandConfigValue,
isConfigValueConfigured,
resolveConfigValueOrThrow,
resolveHeadersOrThrow,
} from "./resolve-config-value.ts";
export interface ExtensionOAuthConfig {
name: string;
/** @deprecated Retained for extension source compatibility; ignored by canonical auth flows. */
usesCallbackServer?: boolean;
login(callbacks: OAuthLoginCallbacks): Promise<OAuthCredentials>;
refreshToken(credentials: OAuthCredentials): Promise<OAuthCredentials>;
getApiKey(credentials: OAuthCredentials): string;
}
/** Input type for the extension registerProvider API. */
export interface ProviderConfigInput {
name?: string;
baseUrl?: string;
apiKey?: string;
api?: Api;
streamSimple?: (model: Model<Api>, context: Context, options?: SimpleStreamOptions) => AssistantMessageEventStream;
headers?: Record<string, string>;
authHeader?: boolean;
oauth?: ExtensionOAuthConfig;
models?: Array<{
id: string;
name: string;
api?: Api;
baseUrl?: string;
reasoning: boolean;
thinkingLevelMap?: Model<Api>["thinkingLevelMap"];
input: ("text" | "image")[];
cost: Model<Api>["cost"];
contextWindow: number;
maxTokens: number;
headers?: Record<string, string>;
compat?: Model<Api>["compat"];
}>;
}
export type AuthStatus = {
configured: boolean;
source?: "stored" | "runtime" | "environment" | "fallback" | "models_json_key" | "models_json_command";
label?: string;
};
export const clearApiKeyCache = clearConfigValueCache;
function mergeCompat(
base: Model<Api>["compat"],
override: Model<Api>["compat"] | ModelsJsonModelOverride["compat"],
): Model<Api>["compat"] {
if (!override) return base;
const merged = { ...base, ...override } as NonNullable<Model<Api>["compat"]>;
const baseNested = base as Record<string, unknown> | undefined;
const overrideNested = override as Record<string, unknown>;
const mergedNested = merged as Record<string, unknown>;
for (const key of ["openRouterRouting", "vercelGatewayRouting", "chatTemplateKwargs"] as const) {
const baseValue = baseNested?.[key];
const overrideValue = overrideNested[key];
if (
(typeof baseValue === "object" && baseValue !== null) ||
(typeof overrideValue === "object" && overrideValue !== null)
) {
mergedNested[key] = { ...(baseValue as object | undefined), ...(overrideValue as object | undefined) };
}
}
return merged;
}
function applyModelOverride(model: Model<Api>, override: ModelsJsonModelOverride): Model<Api> {
return {
...model,
name: override.name ?? model.name,
reasoning: override.reasoning ?? model.reasoning,
thinkingLevelMap: override.thinkingLevelMap
? { ...model.thinkingLevelMap, ...override.thinkingLevelMap }
: model.thinkingLevelMap,
input: (override.input as ("text" | "image")[] | undefined) ?? model.input,
cost: override.cost
? {
input: override.cost.input ?? model.cost.input,
output: override.cost.output ?? model.cost.output,
cacheRead: override.cost.cacheRead ?? model.cost.cacheRead,
cacheWrite: override.cost.cacheWrite ?? model.cost.cacheWrite,
tiers: override.cost.tiers ?? model.cost.tiers,
}
: model.cost,
contextWindow: override.contextWindow ?? model.contextWindow,
maxTokens: override.maxTokens ?? model.maxTokens,
compat: mergeCompat(model.compat, override.compat),
};
}
function modelFromJson(
providerId: string,
definition: ModelsJsonModel,
providerConfig: ModelsJsonProvider,
defaults: Model<Api> | undefined,
): Model<Api> {
const api = definition.api ?? providerConfig.api ?? defaults?.api;
if (!api) {
throw new Error(
`Provider ${providerId}, model ${definition.id}: no "api" specified. Set at provider or model level.`,
);
}
const baseUrl = definition.baseUrl ?? providerConfig.baseUrl ?? defaults?.baseUrl;
if (!baseUrl) throw new Error(`Provider ${providerId}: "baseUrl" is required when defining custom models.`);
if (definition.contextWindow !== undefined && definition.contextWindow <= 0) {
throw new Error(`Provider ${providerId}, model ${definition.id}: invalid contextWindow`);
}
if (definition.maxTokens !== undefined && definition.maxTokens <= 0) {
throw new Error(`Provider ${providerId}, model ${definition.id}: invalid maxTokens`);
}
return {
id: definition.id,
name: definition.name ?? definition.id,
api: api as Api,
provider: providerId,
baseUrl,
reasoning: definition.reasoning ?? false,
thinkingLevelMap: definition.thinkingLevelMap,
input: (definition.input ?? ["text"]) as ("text" | "image")[],
cost: definition.cost ?? { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
contextWindow: definition.contextWindow ?? 128000,
maxTokens: definition.maxTokens ?? 16384,
headers: undefined,
compat: mergeCompat(providerConfig.compat, definition.compat),
};
}
function applyModelsJson(
providerId: string,
baseModels: readonly Model<Api>[],
config: ModelsJsonProvider | undefined,
): Model<Api>[] {
if (!config) return [...baseModels];
const hasOverrides = config.modelOverrides && Object.keys(config.modelOverrides).length > 0;
if (
!config.models?.length &&
!config.baseUrl &&
!config.headers &&
!config.compat &&
!hasOverrides &&
!config.apiKey &&
config.authHeader === undefined
) {
throw new Error(
`Provider ${providerId}: must specify "baseUrl", "headers", "compat", "modelOverrides", or "models".`,
);
}
const models: Model<Api>[] = baseModels.map((model) => ({
...model,
baseUrl: config.baseUrl ?? model.baseUrl,
compat: mergeCompat(model.compat, config.compat),
}));
for (const definition of config.models ?? []) {
const existingIndex = models.findIndex((model) => model.id === definition.id);
const defaults = existingIndex >= 0 ? models[existingIndex] : models[0];
const model = modelFromJson(providerId, definition, config, defaults);
if (existingIndex >= 0) models[existingIndex] = model;
else models.push(model);
}
return models;
}
function applyExtension(
providerId: string,
models: readonly Model<Api>[],
config: ProviderConfigInput | undefined,
): Model<Api>[] {
if (!config) return [...models];
if (!config.models) {
return config.baseUrl ? models.map((model) => ({ ...model, baseUrl: config.baseUrl! })) : [...models];
}
return config.models.map((definition) => {
const defaults = models.find((model) => model.id === definition.id) ?? models[0];
const api = definition.api ?? config.api ?? defaults?.api;
if (!api) {
throw new Error(
`Provider ${providerId}, model ${definition.id}: no "api" specified. Set at provider or model level.`,
);
}
const baseUrl = definition.baseUrl ?? config.baseUrl ?? defaults?.baseUrl;
if (!baseUrl) throw new Error(`Provider ${providerId}: "baseUrl" is required when defining custom models.`);
return {
...definition,
api,
provider: providerId,
baseUrl,
headers: undefined,
};
});
}
function adaptOAuth(config: ExtensionOAuthConfig): OAuthAuth {
return {
name: config.name,
login: async (callbacks) => {
const credential = await config.login({
onAuth: (info) => callbacks.notify({ type: "auth_url", ...info }),
onDeviceCode: (info) => callbacks.notify({ type: "device_code", ...info }),
onPrompt: (prompt) => callbacks.prompt({ type: "text", ...prompt }),
onProgress: (message) => callbacks.notify({ type: "progress", message }),
onManualCodeInput: () => callbacks.prompt({ type: "manual_code", message: "Paste the authorization code" }),
onSelect: (prompt) => callbacks.prompt({ type: "select", ...prompt }),
signal: callbacks.signal,
});
return { ...credential, type: "oauth" };
},
refresh: async (credential) => ({ ...(await config.refreshToken(credential)), type: "oauth" }),
toAuth: async (credential) => ({ apiKey: config.getApiKey(credential) }),
};
}
function withConfiguredAuth(
auth: ModelAuth,
headers: Record<string, string> | undefined,
authHeader: boolean,
): ModelAuth {
let mergedHeaders: ProviderHeaders | undefined =
auth.headers || headers ? { ...auth.headers, ...headers } : undefined;
if (authHeader) {
if (!auth.apiKey) throw new Error("authHeader requires a resolved API key");
mergedHeaders = { ...mergedHeaders, Authorization: `Bearer ${auth.apiKey}` };
}
return { ...auth, headers: mergedHeaders };
}
function configuredApiKey(
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): string | undefined {
return extension?.apiKey ?? config?.apiKey;
}
function configuredHeaders(
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): Record<string, string> | undefined {
if (!config?.headers && !extension?.headers) return undefined;
return { ...config?.headers, ...extension?.headers };
}
async function configContextEnv(
values: readonly string[],
ctx: AuthContext,
explicit?: Record<string, string>,
): Promise<Record<string, string> | undefined> {
const env = { ...explicit };
for (const name of new Set(values.flatMap(getConfigValueEnvVarNames))) {
if (env[name] !== undefined) continue;
const value = await ctx.env(name);
if (value !== undefined) env[name] = value;
}
return Object.keys(env).length > 0 ? env : undefined;
}
function composeApiKeyAuth(
providerId: string,
base: Provider | undefined,
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): ApiKeyAuth | undefined {
const inherited = base?.auth.apiKey;
const rawKey = configuredApiKey(config, extension);
const oauth = extension?.oauth ?? base?.auth.oauth;
// OAuth-only providers get no fabricated API-key login method.
if (!inherited && rawKey === undefined && oauth) return undefined;
const rawHeaders = configuredHeaders(config, extension);
const authHeader = extension?.authHeader ?? config?.authHeader ?? false;
return {
name: inherited?.name ?? "API key",
login:
inherited?.login ??
(async (interaction: AuthInteraction) => ({
type: "api_key",
key: await interaction.prompt({ type: "secret", message: "Enter API key" }),
})),
check: async (input) => {
if (input.credential) {
if (inherited?.check) return inherited.check(input);
if (input.credential.key) return { type: "api_key", source: "stored credential" };
const resolved = await inherited?.resolve(input);
return resolved ? { type: "api_key", source: resolved.source } : undefined;
}
if (rawKey !== undefined) {
if (isCommandConfigValue(rawKey)) return { type: "api_key", source: "configured API key" };
const envNames = getConfigValueEnvVarNames(rawKey);
for (const name of envNames) {
if ((await input.ctx.env(name)) === undefined) return undefined;
}
return { type: "api_key", source: "configured API key" };
}
if (inherited?.check) return inherited.check(input);
const resolved = await inherited?.resolve(input);
return resolved ? { type: "api_key", source: resolved.source } : undefined;
},
resolve: async (input) => {
let result: AuthResult | undefined;
if (input.credential) {
result = inherited
? await inherited.resolve(input)
: input.credential.key
? { auth: { apiKey: input.credential.key }, env: input.credential.env, source: "stored credential" }
: undefined;
} else if (rawKey !== undefined) {
const env = await configContextEnv([rawKey], input.ctx);
const key = resolveConfigValueOrThrow(rawKey, `API key for provider "${providerId}"`, env);
result = inherited
? await inherited.resolve({ ...input, credential: { type: "api_key", key } })
: { auth: { apiKey: key }, source: "configured API key" };
} else {
result = await inherited?.resolve(input);
}
if (!result) return undefined;
const explicitEnv = { ...(input.credential?.env ?? {}), ...(result.env ?? {}) };
const headerEnv = await configContextEnv(Object.values(rawHeaders ?? {}), input.ctx, explicitEnv);
const headers = resolveHeadersOrThrow(rawHeaders, `provider "${providerId}"`, headerEnv);
return { ...result, auth: withConfiguredAuth(result.auth, headers, authHeader) };
},
};
}
function composeOAuthAuth(
providerId: string,
base: Provider | undefined,
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): OAuthAuth | undefined {
const oauth = extension?.oauth ? adaptOAuth(extension.oauth) : base?.auth.oauth;
if (!oauth) return undefined;
const rawHeaders = configuredHeaders(config, extension);
const authHeader = extension?.authHeader ?? config?.authHeader ?? false;
return {
...oauth,
toAuth: async (credential) => {
const auth = await oauth.toAuth(credential);
const env = credential.env;
const headers = resolveHeadersOrThrow(
rawHeaders,
`provider "${providerId}"`,
typeof env === "object" && env !== null ? (env as Record<string, string>) : undefined,
);
return withConfiguredAuth(auth, headers, authHeader);
},
};
}
function rawModelHeaders(
model: Model<Api>,
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): Record<string, string> | undefined {
const definition = config?.models?.find((entry) => entry.id === model.id);
const extensionModel = extension?.models?.find((entry) => entry.id === model.id);
const headers = {
...config?.modelOverrides?.[model.id]?.headers,
...definition?.headers,
...extensionModel?.headers,
};
return Object.keys(headers).length > 0 ? headers : undefined;
}
export function validateExtensionProvider(
providerId: string,
base: Provider | undefined,
modelsConfig: ModelsJsonProvider | undefined,
extension: ProviderConfigInput,
): void {
if (extension.streamSimple && !extension.api) {
throw new Error(`Provider ${providerId}: "api" is required when registering streamSimple.`);
}
applyExtension(providerId, applyModelsJson(providerId, base?.getModels() ?? [], modelsConfig), extension);
}
/** Compose built-in, models.json, and extension layers without reading credentials. */
export function composeModelProvider(
providerId: string,
base: Provider | undefined,
modelConfig: ModelConfig,
extension: ProviderConfigInput | undefined,
): Provider {
const config = modelConfig.getProvider(providerId);
// models.json modelOverrides are the topmost user-config layer: they apply once,
// after custom-model upserts and extension model replacement.
const getModels = () =>
applyExtension(providerId, applyModelsJson(providerId, base?.getModels() ?? [], config), extension).map(
(model) => {
const override = config?.modelOverrides?.[model.id];
return override ? applyModelOverride(model, override) : model;
},
);
// Validate eagerly so registration/reload reports structural errors immediately.
getModels();
const apiKey = composeApiKeyAuth(providerId, base, config, extension);
const oauth = composeOAuthAuth(providerId, base, config, extension);
if (!apiKey && !oauth) throw new Error(`Provider ${providerId}: no authentication method configured.`);
const supportsBaseApi = (model: Model<Api>) => base?.getModels().some((entry) => entry.api === model.api) ?? false;
const streamWith = (
model: Model<Api>,
context: Context,
options: StreamOptions | undefined,
simple: boolean,
): AssistantMessageEventStream =>
lazyStream(model, async () => {
if (extension?.streamSimple && model.api === extension.api) {
return extension.streamSimple(model, context, options as SimpleStreamOptions);
}
if (base && supportsBaseApi(model)) {
return simple
? base.streamSimple(model, context, options as SimpleStreamOptions)
: base.stream(model, context, options);
}
const api = getApiProvider(model.api);
if (!api) throw new Error(`No API provider registered for api: ${model.api}`);
return simple
? api.streamSimple(model, context, options as SimpleStreamOptions)
: api.stream(model, context, options);
});
return {
id: providerId,
name: extension?.name ?? config?.name ?? base?.name ?? extension?.oauth?.name ?? providerId,
baseUrl: extension?.baseUrl ?? config?.baseUrl ?? base?.baseUrl,
headers: base?.headers,
auth: { ...(apiKey ? { apiKey } : {}), ...(oauth ? { oauth } : {}) },
getModels,
refreshModels: base?.refreshModels ? () => base.refreshModels!() : undefined,
filterModels: base?.filterModels
? (models, credential: Credential | undefined) => base.filterModels!(models, credential)
: undefined,
stream: (model, context, options) => streamWith(model, context, options, false),
streamSimple: (model, context, options) => streamWith(model, context, options, true),
};
}
export function resolveConfiguredModelHeaders(
model: Model<Api>,
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
env?: Record<string, string>,
): Record<string, string> | undefined {
return resolveHeadersOrThrow(
rawModelHeaders(model, config, extension),
`model "${model.provider}/${model.id}"`,
env,
);
}
export interface CompatibilityRequestConfig {
headers?: ProviderHeaders;
authHeader: boolean;
}
export function resolveCompatibilityRequestConfig(
model: Model<Api>,
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): CompatibilityRequestConfig {
const configured = resolveHeadersOrThrow(
{ ...configuredHeaders(config, extension), ...rawModelHeaders(model, config, extension) },
`model "${model.provider}/${model.id}"`,
);
return {
headers: model.headers || configured ? { ...model.headers, ...configured } : undefined,
authHeader: extension?.authHeader ?? config?.authHeader ?? false,
};
}
export function configuredRequestAuthStatus(
config: ModelsJsonProvider | undefined,
extension: ProviderConfigInput | undefined,
): AuthStatus | undefined {
const value = configuredApiKey(config, extension);
if (value === undefined) return undefined;
if (isCommandConfigValue(value)) return { configured: true, source: "models_json_command" };
const names = getConfigValueEnvVarNames(value);
if (names.length > 0) {
return isConfigValueConfigured(value)
? { configured: true, source: "environment", label: names.join(", ") }
: { configured: false };
}
return { configured: true, source: extension?.apiKey !== undefined ? "fallback" : "models_json_key" };
}
@@ -1,35 +0,0 @@
export const BUILT_IN_PROVIDER_DISPLAY_NAMES: Record<string, string> = {
anthropic: "Anthropic",
"amazon-bedrock": "Amazon Bedrock",
"ant-ling": "Ant Ling",
"azure-openai-responses": "Azure OpenAI Responses",
cerebras: "Cerebras",
"cloudflare-ai-gateway": "Cloudflare AI Gateway",
"cloudflare-workers-ai": "Cloudflare Workers AI",
deepseek: "DeepSeek",
fireworks: "Fireworks",
google: "Google Gemini",
"google-vertex": "Google Vertex AI",
groq: "Groq",
huggingface: "Hugging Face",
"kimi-coding": "Kimi For Coding",
mistral: "Mistral",
minimax: "MiniMax",
"minimax-cn": "MiniMax (China)",
moonshotai: "Moonshot AI",
"moonshotai-cn": "Moonshot AI (China)",
nvidia: "NVIDIA NIM",
opencode: "OpenCode Zen",
"opencode-go": "OpenCode Go",
openai: "OpenAI",
openrouter: "OpenRouter",
together: "Together AI",
"vercel-ai-gateway": "Vercel AI Gateway",
xai: "xAI",
zai: "ZAI Coding Plan (Global)",
"zai-coding-cn": "ZAI Coding Plan (China)",
xiaomi: "Xiaomi MiMo",
"xiaomi-token-plan-cn": "Xiaomi MiMo Token Plan (China)",
"xiaomi-token-plan-ams": "Xiaomi MiMo Token Plan (Amsterdam)",
"xiaomi-token-plan-sgp": "Xiaomi MiMo Token Plan (Singapore)",
};
@@ -0,0 +1,48 @@
import type { Credential, CredentialInfo, CredentialStore } from "@earendil-works/pi-ai";
/** Async credential store overlay for non-persistent runtime API keys. */
export class RuntimeCredentials implements CredentialStore {
private readonly store: CredentialStore;
private readonly overrides = new Map<string, string>();
constructor(store: CredentialStore) {
this.store = store;
}
setRuntimeApiKey(providerId: string, apiKey: string): void {
this.overrides.set(providerId, apiKey);
}
removeRuntimeApiKey(providerId: string): void {
this.overrides.delete(providerId);
}
hasRuntimeApiKey(providerId: string): boolean {
return this.overrides.has(providerId);
}
async read(providerId: string): Promise<Credential | undefined> {
const override = this.overrides.get(providerId);
return override ? { type: "api_key", key: override } : this.store.read(providerId);
}
async list(): Promise<readonly CredentialInfo[]> {
const entries = new Map((await this.store.list()).map((entry) => [entry.providerId, entry]));
for (const providerId of this.overrides.keys()) {
entries.set(providerId, { providerId, type: "api_key" });
}
return [...entries.values()];
}
modify(
providerId: string,
fn: (current: Credential | undefined) => Promise<Credential | undefined>,
): Promise<Credential | undefined> {
return this.store.modify(providerId, fn);
}
async delete(providerId: string): Promise<void> {
this.overrides.delete(providerId);
await this.store.delete(providerId);
}
}
+21 -35
View File
@@ -1,16 +1,15 @@
import { join } from "node:path";
import { Agent, type AgentMessage, type ThinkingLevel } from "@earendil-works/pi-agent-core";
import { clampThinkingLevel, type Message, type Model, streamSimple } from "@earendil-works/pi-ai/compat";
import { clampThinkingLevel, type Message, type Model } from "@earendil-works/pi-ai/compat";
import { getAgentDir } from "../config.ts";
import { resolvePath } from "../utils/paths.ts";
import { AgentSession } from "./agent-session.ts";
import { formatNoModelsAvailableMessage } from "./auth-guidance.ts";
import { AuthStorage } from "./auth-storage.ts";
import { DEFAULT_THINKING_LEVEL } from "./defaults.ts";
import type { ExtensionRunner, LoadExtensionsResult, SessionStartEvent, ToolDefinition } from "./extensions/index.ts";
import { convertToLlm } from "./messages.ts";
import { ModelRegistry } from "./model-registry.ts";
import { findInitialModel } from "./model-resolver.ts";
import { ModelRuntime } from "./model-runtime.ts";
import { mergeProviderAttributionHeaders } from "./provider-attribution.ts";
import type { ResourceLoader } from "./resource-loader.ts";
import { DefaultResourceLoader } from "./resource-loader.ts";
@@ -37,10 +36,8 @@ export interface CreateAgentSessionOptions {
/** Global config directory. Default: ~/.pi/agent */
agentDir?: string;
/** Auth storage for credentials. Default: AuthStorage.create(agentDir/auth.json) */
authStorage?: AuthStorage;
/** Model registry. Default: ModelRegistry.create(authStorage, agentDir/models.json) */
modelRegistry?: ModelRegistry;
/** Canonical model/auth runtime. Defaults to a runtime using agentDir/auth.json and models.json. */
modelRuntime?: ModelRuntime;
/** Model to use. Default: from settings, else first available */
model?: Model<any>;
@@ -169,11 +166,9 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
const agentDir = options.agentDir ? resolvePath(options.agentDir) : getDefaultAgentDir();
let resourceLoader = options.resourceLoader;
// Use provided or create AuthStorage and ModelRegistry
const authPath = options.agentDir ? join(agentDir, "auth.json") : undefined;
const modelsPath = options.agentDir ? join(agentDir, "models.json") : undefined;
const authStorage = options.authStorage ?? AuthStorage.create(authPath);
const modelRegistry = options.modelRegistry ?? ModelRegistry.create(authStorage, modelsPath);
const modelRuntime = options.modelRuntime ?? (await ModelRuntime.create({ authPath, modelsPath }));
const settingsManager = options.settingsManager ?? SettingsManager.create(cwd, agentDir);
const sessionManager = options.sessionManager ?? SessionManager.create(cwd, getDefaultSessionDir(cwd, agentDir));
@@ -194,8 +189,8 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
// If session has data, try to restore model from it
if (!model && hasExistingSession && existingSession.model) {
const restoredModel = modelRegistry.find(existingSession.model.provider, existingSession.model.modelId);
if (restoredModel && modelRegistry.hasConfiguredAuth(restoredModel)) {
const restoredModel = modelRuntime.getModel(existingSession.model.provider, existingSession.model.modelId);
if (restoredModel && modelRuntime.hasConfiguredAuth(restoredModel.provider)) {
model = restoredModel;
}
if (!model) {
@@ -211,7 +206,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
defaultProvider: settingsManager.getDefaultProvider(),
defaultModelId: settingsManager.getDefaultModel(),
defaultThinkingLevel: settingsManager.getDefaultThinkingLevel(),
modelRegistry,
modelRuntime,
});
model = result.model;
if (!model) {
@@ -300,11 +295,6 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
},
convertToLlm: convertToLlmWithBlockImages,
streamFn: async (model, context, options) => {
const auth = await modelRegistry.getApiKeyAndHeaders(model);
if (!auth.ok) {
throw new Error(auth.error);
}
const env = auth.env || options?.env ? { ...(auth.env ?? {}), ...(options?.env ?? {}) } : undefined;
const providerRetrySettings = settingsManager.getProviderRetrySettings();
const httpIdleTimeoutMs = settingsManager.getHttpIdleTimeoutMs();
// SDKs treat timeout=0 as 0ms (immediate timeout), not "no timeout".
@@ -313,28 +303,24 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
const timeoutMs = options?.timeoutMs ?? providerRetrySettings.timeoutMs ?? effectiveTimeoutMs;
const websocketConnectTimeoutMs =
options?.websocketConnectTimeoutMs ?? settingsManager.getWebSocketConnectTimeoutMs();
let headers = mergeProviderAttributionHeaders(
model,
settingsManager,
options?.sessionId,
auth.headers,
options?.headers,
);
// Let extensions inject/adjust per-request headers (e.g. tracing, session correlation)
// after static assembly, before the provider HTTP call.
const headerRunner = extensionRunnerRef.current;
if (headerRunner?.hasHandlers("before_provider_headers")) {
headers = await headerRunner.emitBeforeProviderHeaders(headers ?? {});
}
return streamSimple(model, context, {
return modelRuntime.streamSimple(model, context, {
...options,
apiKey: auth.apiKey,
env,
timeoutMs,
websocketConnectTimeoutMs,
maxRetries: options?.maxRetries ?? providerRetrySettings.maxRetries,
maxRetryDelayMs: options?.maxRetryDelayMs ?? providerRetrySettings.maxRetryDelayMs,
headers,
transformHeaders: async (requestHeaders) => {
const headers = mergeProviderAttributionHeaders(
model,
settingsManager,
options?.sessionId,
requestHeaders,
);
return headerRunner?.hasHandlers("before_provider_headers")
? headerRunner.emitBeforeProviderHeaders(headers ?? {})
: (headers ?? {});
},
});
},
onPayload: async (payload, _model) => {
@@ -390,7 +376,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
scopedModels: options.scopedModels,
resourceLoader,
customTools: options.customTools,
modelRegistry,
modelRuntime,
initialActiveToolNames,
allowedToolNames,
excludedToolNames,