feat(coding-agent): accept native extension providers

Allow extensions to register complete pi-ai Provider objects while preserving models.json composition and resolved provider auth access.

Refs #4823 and #4824.
This commit is contained in:
Mario Zechner
2026-07-17 10:58:24 +02:00
parent c8560b8d70
commit 019e4ad687
12 changed files with 349 additions and 13 deletions
@@ -165,6 +165,18 @@ export async function createAgentSessionServices(
}
}
extensionsResult.runtime.pendingProviderRegistrations = [];
for (const { provider, extensionPath } of extensionsResult.runtime.pendingNativeProviderRegistrations) {
try {
modelRuntime.registerNativeProvider(provider);
} catch (error) {
const message = error instanceof Error ? error.message : String(error);
diagnostics.push({
type: "error",
message: `Extension "${extensionPath}" error: ${message}`,
});
}
}
extensionsResult.runtime.pendingNativeProviderRegistrations = [];
await modelRuntime.refresh({ allowNetwork: false });
diagnostics.push(...applyExtensionFlagValues(resourceLoader, options.extensionFlagValues));
@@ -2419,6 +2419,10 @@ export class AgentSession {
this._modelRuntime.registerProvider(name, config);
this._refreshCurrentModelFromRegistry();
},
registerNativeProvider: (provider) => {
this._modelRuntime.registerNativeProvider(provider);
this._refreshCurrentModelFromRegistry();
},
unregisterProvider: (name) => {
this._modelRuntime.unregisterProvider(name);
this._refreshCurrentModelFromRegistry();
@@ -8,6 +8,7 @@ import { createRequire } from "node:module";
import * as path from "node:path";
import { fileURLToPath } from "node:url";
import * as _bundledPiAgentCore from "@earendil-works/pi-agent-core";
import type { Provider } from "@earendil-works/pi-ai";
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";
@@ -195,6 +196,7 @@ export function createExtensionRuntime(): ExtensionRuntime {
setThinkingLevel: notInitialized,
flagValues: new Map(),
pendingProviderRegistrations: [],
pendingNativeProviderRegistrations: [],
assertActive,
invalidate: (message) => {
state.staleMessage ??=
@@ -206,8 +208,14 @@ export function createExtensionRuntime(): ExtensionRuntime {
registerProvider: (name, config, extensionPath = "<unknown>") => {
runtime.pendingProviderRegistrations.push({ name, config, extensionPath });
},
registerNativeProvider: (provider, extensionPath = "<unknown>") => {
runtime.pendingNativeProviderRegistrations.push({ provider, extensionPath });
},
unregisterProvider: (name) => {
runtime.pendingProviderRegistrations = runtime.pendingProviderRegistrations.filter((r) => r.name !== name);
runtime.pendingNativeProviderRegistrations = runtime.pendingNativeProviderRegistrations.filter(
(r) => r.provider.id !== name,
);
},
};
@@ -363,9 +371,14 @@ function createExtensionAPI(
runtime.setThinkingLevel(level);
},
registerProvider(name: string, config: ProviderConfig) {
registerProvider(providerOrName: Provider | string, config?: ProviderConfig) {
runtime.assertActive();
runtime.registerProvider(name, config, extension.path);
if (typeof providerOrName === "string") {
if (!config) throw new Error("Provider config is required when registering by name");
runtime.registerProvider(providerOrName, config, extension.path);
return;
}
runtime.registerNativeProvider(providerOrName, extension.path);
},
unregisterProvider(name: string) {
@@ -3,7 +3,7 @@
*/
import type { AgentMessage } from "@earendil-works/pi-agent-core";
import type { ImageContent, Model, ProviderHeaders } from "@earendil-works/pi-ai";
import type { ImageContent, Model, Provider, ProviderHeaders } from "@earendil-works/pi-ai";
import type { KeyId } from "@earendil-works/pi-tui";
import { type Theme, theme } from "../../modes/interactive/theme/theme.ts";
import type { ResourceDiagnostic } from "../diagnostics.ts";
@@ -313,6 +313,7 @@ export class ExtensionRunner {
contextActions: ExtensionContextActions,
providerActions?: {
registerProvider?: (name: string, config: ProviderConfig) => void;
registerNativeProvider?: (provider: Provider) => void;
unregisterProvider?: (name: string) => void;
},
): void {
@@ -363,6 +364,23 @@ export class ExtensionRunner {
}
}
this.runtime.pendingProviderRegistrations = [];
for (const { provider, extensionPath } of this.runtime.pendingNativeProviderRegistrations) {
try {
if (providerActions?.registerNativeProvider) {
providerActions.registerNativeProvider(provider);
} else {
this.modelRegistry.registerProvider(provider);
}
} catch (err) {
this.emitError({
extensionPath,
event: "register_provider",
error: err instanceof Error ? err.message : String(err),
stack: err instanceof Error ? err.stack : undefined,
});
}
}
this.runtime.pendingNativeProviderRegistrations = [];
// From this point on, provider registration/unregistration takes effect immediately
// without requiring a /reload.
@@ -373,6 +391,13 @@ export class ExtensionRunner {
}
this.modelRegistry.registerProvider(name, config);
};
this.runtime.registerNativeProvider = (provider) => {
if (providerActions?.registerNativeProvider) {
providerActions.registerNativeProvider(provider);
return;
}
this.modelRegistry.registerProvider(provider);
};
this.runtime.unregisterProvider = (name) => {
if (providerActions?.unregisterProvider) {
providerActions.unregisterProvider(name);
@@ -24,6 +24,7 @@ import type {
Model,
OAuthCredentials,
OAuthLoginCallbacks,
Provider,
ProviderHeaders,
RefreshModelsContext,
SimpleStreamOptions,
@@ -1378,6 +1379,7 @@ export interface ExtensionAPI {
* }
* });
*/
registerProvider(provider: Provider): void;
registerProvider(name: string, config: ProviderConfig): void;
/**
@@ -1553,8 +1555,10 @@ export type SetLabelHandler = (entryId: string, label: string | undefined) => vo
*/
export interface ExtensionRuntimeState {
flagValues: Map<string, boolean | string>;
/** Provider registrations queued during extension loading, processed when runner binds */
/** Legacy provider-config registrations queued during extension loading, processed when runner binds. */
pendingProviderRegistrations: Array<{ name: string; config: ProviderConfig; extensionPath: string }>;
/** Native pi-ai provider registrations queued during extension loading, processed when runner binds. */
pendingNativeProviderRegistrations: Array<{ provider: Provider; extensionPath: string }>;
/** Throws when this extension instance is stale after runtime replacement. */
assertActive: () => void;
/** Marks this extension instance as stale after runtime replacement or reload. */
@@ -1566,6 +1570,7 @@ export interface ExtensionRuntimeState {
* After bindCore(): calls ModelRegistry directly for immediate effect.
*/
registerProvider: (name: string, config: ProviderConfig, extensionPath?: string) => void;
registerNativeProvider: (provider: Provider, extensionPath?: string) => void;
unregisterProvider: (name: string, extensionPath?: string) => void;
}
@@ -1,4 +1,4 @@
import type { Api, Model } from "@earendil-works/pi-ai";
import type { Api, AuthResult, Model, Provider } from "@earendil-works/pi-ai";
import type { ModelRuntime } from "./model-runtime.ts";
import type { AuthStatus, ProviderConfigInput } from "./provider-composer.ts";
@@ -92,10 +92,18 @@ export class ModelRegistry {
return this.runtime.getProviderAuthStatus(provider);
}
getProvider(provider: string): Provider | undefined {
return this.runtime.getProvider(provider);
}
getProviderDisplayName(provider: string): string {
return this.runtime.getProvider(provider)?.name ?? provider;
}
getProviderAuth(provider: string): Promise<AuthResult | undefined> {
return this.runtime.getAuth(provider);
}
async getApiKeyForProvider(provider: string): Promise<string | undefined> {
try {
return (await this.runtime.getAuth(provider))?.auth.apiKey;
@@ -108,8 +116,15 @@ export class ModelRegistry {
return this.runtime.isUsingOAuth(model.provider);
}
registerProvider(providerName: string, config: ProviderConfigInput): void {
this.runtime.registerProvider(providerName, config);
registerProvider(provider: Provider): void;
registerProvider(providerName: string, config: ProviderConfigInput): void;
registerProvider(providerOrName: Provider | string, config?: ProviderConfigInput): void {
if (typeof providerOrName === "string") {
if (!config) throw new Error("Provider config is required when registering by name");
this.runtime.registerProvider(providerOrName, config);
return;
}
this.runtime.registerNativeProvider(providerOrName);
}
unregisterProvider(providerName: string): void {
@@ -120,6 +135,10 @@ export class ModelRegistry {
return this.runtime.getRegisteredProviderConfig(providerName);
}
getRegisteredNativeProvider(providerName: string): Provider | undefined {
return this.runtime.getRegisteredNativeProvider(providerName);
}
getRegisteredProviderIds(): readonly string[] {
return this.runtime.getRegisteredProviderIds();
}
@@ -94,6 +94,7 @@ export class ModelRuntime implements Models {
private readonly credentials: RuntimeCredentials;
private readonly defaultBuiltins: ReadonlyMap<string, Provider>;
private readonly builtins = new Map<string, Provider>();
private readonly nativeExtensionProviders = new Map<string, Provider>();
private readonly extensionProviders = new Map<string, ProviderConfigInput>();
private readonly compositionErrors = new Map<string, string>();
private readonly modelsPath: string | undefined;
@@ -182,11 +183,16 @@ export class ModelRuntime implements Models {
}
private providerIds(): Set<string> {
return new Set([...this.builtins.keys(), ...this.config.getProviderIds(), ...this.extensionProviders.keys()]);
return new Set([
...this.builtins.keys(),
...this.nativeExtensionProviders.keys(),
...this.config.getProviderIds(),
...this.extensionProviders.keys(),
]);
}
private recomposeProvider(providerId: string): void {
const base = this.builtins.get(providerId);
const base = this.nativeExtensionProviders.get(providerId) ?? this.builtins.get(providerId);
const extension = this.extensionProviders.get(providerId);
if (!base && !this.config.getProvider(providerId) && !extension) {
this.models.deleteProvider(providerId);
@@ -335,7 +341,11 @@ export class ModelRuntime implements Models {
}
getRegisteredProviderIds(): readonly string[] {
return [...this.extensionProviders.keys()];
return [...new Set([...this.extensionProviders.keys(), ...this.nativeExtensionProviders.keys()])];
}
getRegisteredNativeProvider(providerId: string): Provider | undefined {
return this.nativeExtensionProviders.get(providerId);
}
/** @internal Compatibility fallback for ModelRegistry when provider auth is unconfigured. */
@@ -520,10 +530,20 @@ export class ModelRuntime implements Models {
return result;
}
registerNativeProvider(provider: Provider): void {
if (!provider.id.trim()) throw new Error("Provider id must not be empty.");
this.extensionProviders.delete(provider.id);
this.nativeExtensionProviders.set(provider.id, provider);
this.recomposeProvider(provider.id);
this.updateModelSnapshot();
void this.refresh({ allowNetwork: false });
}
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);
this.nativeExtensionProviders.delete(providerId);
// Re-registration merges defined values over the previous registration and
// preserves undefined ones, matching the legacy ModelRegistry contract.
const previous = this.extensionProviders.get(providerId);
@@ -559,6 +579,7 @@ export class ModelRuntime implements Models {
unregisterProvider(providerId: string): void {
this.extensionProviders.delete(providerId);
this.nativeExtensionProviders.delete(providerId);
this.recomposeProvider(providerId);
this.updateModelSnapshot();
void this.refresh({ allowNetwork: false });