fix(extensions): restored provider unregistration

Added the upstream unregisterProvider lifecycle to queued and initialized extension runtimes. Provider removal now clears runtime model/auth state before replacement, while failed factories restore the prior registration queue.

Fixes #7914
This commit is contained in:
roboomp
2026-08-07 15:13:19 +00:00
committed by can1357
parent 336975c46c
commit a09dfd0ba8
8 changed files with 249 additions and 6 deletions
@@ -65,7 +65,7 @@ const BUILT_IN_DISCOVERY_CACHE_TTL_MS = 2 * 60 * 60 * 1000;
const BUILT_IN_DISCOVERY_NON_AUTHORITATIVE_RETRY_MS = 5 * 60 * 1000;
import type { ApiKeyResolver, FetchImpl } from "@oh-my-pi/pi-ai";
import { registerOAuthProvider, unregisterOAuthProviders } from "@oh-my-pi/pi-ai/oauth";
import { registerOAuthProvider, unregisterOAuthProvider, unregisterOAuthProviders } from "@oh-my-pi/pi-ai/oauth";
import type { OAuthCredentials, OAuthLoginCallbacks } from "@oh-my-pi/pi-ai/oauth/types";
import { setCodexAttestationProvider } from "@oh-my-pi/pi-ai/providers/openai-codex-responses";
import { getProviderDefinition } from "@oh-my-pi/pi-ai/registry";
@@ -2514,6 +2514,26 @@ export class ModelRegistry {
this.#reloadStaticModels();
}
/**
* Remove one extension-registered provider and restore its static models.
*/
unregisterProvider(providerName: string): void {
const sourceId = this.#runtimeProviderSourceByName.get(providerName);
if (sourceId) {
const sourceProviders = this.#runtimeProvidersBySource.get(sourceId);
sourceProviders?.delete(providerName);
if (sourceProviders?.size === 0) {
this.#runtimeProvidersBySource.delete(sourceId);
}
this.#runtimeProviderSourceByName.delete(providerName);
}
unregisterOAuthProvider(providerName);
this.#ensureFullSnapshot();
this.#clearRuntimeProviderState(providerName);
this.#lastStaticLoadMtime = null;
this.#reloadStaticModels();
}
/**
* Remove registrations for extension sources that are no longer active.
*/
@@ -73,6 +73,15 @@ export class ExtensionRuntime implements IExtensionRuntime {
flagValues = new Map<string, boolean | string>();
pendingProviderRegistrations: Array<{ name: string; config: ProviderConfig; sourceId: string }> = [];
registerProvider(name: string, config: ProviderConfig, sourceId: string): void {
this.pendingProviderRegistrations.push({ name, config, sourceId });
}
unregisterProvider(name: string): void {
const remaining = this.pendingProviderRegistrations.filter(registration => registration.name !== name);
this.pendingProviderRegistrations.splice(0, this.pendingProviderRegistrations.length, ...remaining);
}
sendMessage(): void {
throw new ExtensionRuntimeNotInitializedError();
}
@@ -290,7 +299,11 @@ class ConcreteExtensionAPI implements ExtensionAPI, IExtensionRuntime {
}
registerProvider(name: string, config: ProviderConfig): void {
this.runtime.pendingProviderRegistrations.push({ name, config, sourceId: this.extension.path });
this.runtime.registerProvider(name, config, this.extension.path);
}
unregisterProvider(name: string): void {
this.runtime.unregisterProvider(name, this.extension.path);
}
}
@@ -313,20 +326,24 @@ function createExtension(extensionPath: string, resolvedPath: string): Extension
/**
* Runs an extension factory with provider registration rollback on failure.
* Records the number of pending provider registrations before the factory runs,
* and restores that checkpoint if the factory throws.
* Restores the complete registration queue when the factory throws because an
* extension may unregister entries queued by an earlier extension.
*/
async function runExtensionFactory(
factory: ExtensionFactory,
api: ExtensionAPI,
runtime: IExtensionRuntime,
): Promise<void> {
const providerRegistrationCheckpoint = runtime.pendingProviderRegistrations.length;
const providerRegistrationCheckpoint = [...runtime.pendingProviderRegistrations];
try {
await factory(api);
} catch (error) {
runtime.pendingProviderRegistrations.length = providerRegistrationCheckpoint;
runtime.pendingProviderRegistrations.splice(
0,
runtime.pendingProviderRegistrations.length,
...providerRegistrationCheckpoint,
);
throw error;
}
}
@@ -530,6 +530,12 @@ export class ExtensionRunner {
this.runtime.setServiceTier = actions.setServiceTier ?? throwUnsupportedServiceTierAction;
this.runtime.getSessionName = actions.getSessionName;
this.runtime.setSessionName = actions.setSessionName;
this.runtime.registerProvider = (name, config, sourceId) => {
this.modelRegistry.registerProvider(name, config, sourceId);
};
this.runtime.unregisterProvider = name => {
this.modelRegistry.unregisterProvider(name);
};
// Context actions (required)
this.#getModel = contextActions.getModel;
@@ -1373,6 +1373,14 @@ export interface ExtensionAPI {
*/
registerProvider(name: string, config: ProviderConfig): void;
/**
* Unregister a provider previously registered by an extension.
*
* Removes extension-provided models and restores overridden built-in models.
* Has no effect when the provider is not registered.
*/
unregisterProvider(name: string): void;
/** Shared event bus for extension communication. */
events: EventBus;
}
@@ -1516,6 +1524,10 @@ export interface ExtensionRuntimeState {
flagValues: Map<string, boolean | string>;
/** Provider registrations queued during extension loading, processed during session initialization */
pendingProviderRegistrations: Array<{ name: string; config: ProviderConfig; sourceId: string }>;
/** Queue a provider registration until initialization, then apply it immediately. */
registerProvider(name: string, config: ProviderConfig, sourceId: string): void;
/** Remove a queued or initialized provider registration. */
unregisterProvider(name: string, sourceId: string): void;
}
/** Action implementations for ExtensionAPI methods. */