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:
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user