import { readFile } from "node:fs/promises"; import { join } from "node:path"; import { getBuiltinProviders } from "@earendil-works/pi-ai/providers/all"; import { type ExtensionAPI, type ExtensionContext, getAgentDir, ModelRuntime, } from "@earendil-works/pi-coding-agent"; import { type AliasConfig, loadAliasConfig, saveAliasConfig, validateAliasName } from "./config.ts"; import { appendCleanupFailure, describeFailure } from "./error-diagnostics.ts"; import { createProviderAlias, type SourceHeadersResolver } from "./provider-alias.ts"; const CONFIG_FILE = "login-aliases.json"; type ProviderOption = { id: string; label: string; }; type RegistryProvider = NonNullable>; type ExtensionDependencies = { createStartupProviderCatalog: typeof createStartupProviderCatalog; saveAliasConfig: typeof saveAliasConfig; }; const defaultDependencies: ExtensionDependencies = { createStartupProviderCatalog, saveAliasConfig, }; function getAliasPath(): string { return join(getAgentDir(), CONFIG_FILE); } function isRecord(value: unknown): value is Record { return Boolean(value) && typeof value === "object" && !Array.isArray(value); } async function getStartupSourceIds(aliases: AliasConfig): Promise> { if (Object.keys(aliases).length === 0) { return new Set(); } const knownSourceIds = new Set([...getBuiltinProviders(), "radius"]); try { const contents = await readFile(join(getAgentDir(), "models.json"), "utf8"); const value: unknown = JSON.parse(contents); const providers = isRecord(value) ? value.providers : undefined; if (isRecord(providers)) { for (const providerId of Object.keys(providers)) { knownSourceIds.add(providerId); } } } catch (error) { if (!(isRecord(error) && error.code === "ENOENT")) { console.error( `[login-alias] Could not inspect models.json: ${error instanceof Error ? error.message : String(error)}`, ); } } return new Set( Object.values(aliases) .map(({ source }) => source) .filter((sourceId) => knownSourceIds.has(sourceId)), ); } export async function createStartupProviderCatalog( createRuntime: () => Promise> = () => ModelRuntime.create({ allowModelNetwork: false }), ): Promise> { const runtime = await createRuntime(); return new Map(runtime.getProviders().map((provider) => [provider.id, provider])); } function getProviderOptions(ctx: ExtensionContext, aliases: AliasConfig): ProviderOption[] { const aliasNames = new Set(Object.keys(aliases)); const providerIds = new Set([ ...ctx.modelRegistry.getAll().map((model) => model.provider), ...ctx.modelRegistry.getRegisteredProviderIds(), ]); return [...providerIds] .filter((providerId) => !aliasNames.has(providerId)) .map((providerId) => { const provider = ctx.modelRegistry.getProvider(providerId); return { id: providerId, label: provider ? `${providerId} — ${provider.name}` : providerId, }; }) .sort((a, b) => a.label.localeCompare(b.label)); } function registerAlias( pi: ExtensionAPI, alias: string, resolveSource: () => ReturnType, resolveSourceHeaders?: SourceHeadersResolver, ): void { pi.registerProvider(createProviderAlias(resolveSource, alias, undefined, resolveSourceHeaders)); } function reconcileConfiguredAliases( pi: ExtensionAPI, aliases: AliasConfig, startupProviders: ReadonlyMap, registeredAliases: Set, createSourceResolver: (sourceId: string) => () => RegistryProvider | undefined, resolveSourceHeaders: SourceHeadersResolver, existingProviderIds: ReadonlySet, reportUnavailable: boolean, ): string[] { const errors: string[] = []; for (const [alias, definition] of Object.entries(aliases)) { if (registeredAliases.has(alias)) { continue; } if (startupProviders.has(alias) || existingProviderIds.has(alias)) { errors.push(`Alias "${alias}" conflicts with an existing provider.`); continue; } try { registerAlias(pi, alias, createSourceResolver(definition.source), resolveSourceHeaders); registeredAliases.add(alias); } catch (error) { const unavailableMessage = `Provider "${alias}" source is unavailable.`; if (error instanceof Error && error.message === unavailableMessage) { if (reportUnavailable) { errors.push(`Alias "${alias}" refers to unavailable provider "${definition.source}".`); } continue; } errors.push( `Alias "${alias}" could not be registered: ${error instanceof Error ? error.message : String(error)}`, ); } } return errors; } export default async function loginAliasExtension( pi: ExtensionAPI, dependencies: ExtensionDependencies = defaultDependencies, ): Promise { const aliasPath = getAliasPath(); let aliases: AliasConfig = {}; let configError: string | undefined; let configErrorReported = false; try { aliases = await loadAliasConfig(aliasPath); } catch (error) { configError = error instanceof Error ? error.message : String(error); console.error(`[login-alias] ${configError}`); } // Pi has no factory-time registry lookup. Build the shadow catalog only when // an alias points at a built-in or models.json provider; extension-provided and // unknown sources can wait for the effective registry at session_start. const startupSourceIds = await getStartupSourceIds(aliases); const startupProviders = startupSourceIds.size === 0 ? new Map() : await dependencies.createStartupProviderCatalog(); const registeredAliases = new Set(); let effectiveRegistry: ExtensionContext["modelRegistry"] | undefined; const createSourceResolver = (sourceId: string) => () => effectiveRegistry?.getProvider(sourceId) ?? startupProviders.get(sourceId); const resolveSourceHeaders: SourceHeadersResolver = async (model) => { const registry = effectiveRegistry; if (!registry) { return undefined; } const resolved = await registry.getApiKeyAndHeaders(model); return resolved.ok ? resolved.headers : undefined; }; for (const error of reconcileConfiguredAliases( pi, aliases, startupProviders, registeredAliases, createSourceResolver, resolveSourceHeaders, new Set(startupProviders.keys()), false, )) { console.error(`[login-alias] ${error}`); } pi.registerCommand("login-alias", { description: "Create a provider alias for a separate login", handler: async (args, ctx) => { if (!ctx.hasUI) { ctx.ui.notify("/login-alias requires interactive input.", "error"); return; } if (configError) { ctx.ui.notify(`Could not load ${aliasPath}: ${configError}`, "error"); return; } const alias = ( args.trim() || (await ctx.ui.input("Alias name:", "my-custom-provider")) || "" ).trim(); try { validateAliasName(alias); } catch (error) { ctx.ui.notify(error instanceof Error ? error.message : String(error), "error"); return; } const existingProvider = ctx.modelRegistry.getProvider(alias); if (existingProvider || aliases[alias]) { ctx.ui.notify(`Provider or alias "${alias}" already exists.`, "error"); return; } const options = getProviderOptions(ctx, aliases); if (options.length === 0) { ctx.ui.notify("No providers are available to alias.", "error"); return; } const labels = options.map((option) => option.label); const selected = await ctx.ui.select("Select provider to alias:", labels); if (!selected) { ctx.ui.notify("Alias creation cancelled.", "info"); return; } const sourceId = options.find((option) => option.label === selected)?.id; const source = sourceId ? ctx.modelRegistry.getProvider(sourceId) : undefined; if (!sourceId || !source) { ctx.ui.notify("The selected provider is no longer available.", "error"); return; } const nextAliases: AliasConfig = { ...aliases, [alias]: { source: sourceId }, }; try { registerAlias( pi, alias, () => effectiveRegistry?.getProvider(sourceId) ?? source, resolveSourceHeaders, ); } catch (error) { ctx.ui.notify(`Could not create alias "${alias}": ${describeFailure(error)}`, "error"); return; } try { await dependencies.saveAliasConfig(aliasPath, nextAliases); } catch (persistenceError) { let failure = persistenceError; try { pi.unregisterProvider(alias); } catch (cleanupError) { failure = appendCleanupFailure( persistenceError, `unregister transient provider "${alias}"`, cleanupError, ); } ctx.ui.notify(`Could not create alias "${alias}": ${describeFailure(failure)}`, "error"); return; } aliases = nextAliases; registeredAliases.add(alias); ctx.ui.notify(`Alias "${alias}" created. Run /login ${alias} to authenticate it.`, "info"); }, }); pi.on("session_start", (_event, ctx) => { effectiveRegistry = ctx.modelRegistry; const existingProviderIds = new Set([ ...ctx.modelRegistry.getAll().map((model) => model.provider), ...ctx.modelRegistry.getRegisteredProviderIds(), ]); const errors = reconcileConfiguredAliases( pi, aliases, startupProviders, registeredAliases, createSourceResolver, resolveSourceHeaders, existingProviderIds, true, ); if (configError && !configErrorReported) { ctx.ui.notify(`Could not load ${aliasPath}: ${configError}`, "error"); configErrorReported = true; } for (const error of errors) { ctx.ui.notify(error, "warning"); } }); }