diff --git a/packages/coding-agent/src/extensibility/extensions/loader.ts b/packages/coding-agent/src/extensibility/extensions/loader.ts index dd1a6a949..3f2db9cc8 100644 --- a/packages/coding-agent/src/extensibility/extensions/loader.ts +++ b/packages/coding-agent/src/extensibility/extensions/loader.ts @@ -311,6 +311,26 @@ 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. + */ +async function runExtensionFactory( + factory: ExtensionFactory, + api: ExtensionAPI, + runtime: IExtensionRuntime, +): Promise { + const providerRegistrationCheckpoint = runtime.pendingProviderRegistrations.length; + + try { + await factory(api); + } catch (error) { + runtime.pendingProviderRegistrations.length = providerRegistrationCheckpoint; + throw error; + } +} + async function loadExtension( extensionPath: string, cwd: string, @@ -331,9 +351,7 @@ async function loadExtension( const extension = createExtension(extensionPath, resolvedPath); const api = new ConcreteExtensionAPI(PiCodingAgent, extension, runtime, cwd, eventBus); - await withHostGuard(async () => { - await factory(api); - }); + await withHostGuard(() => runExtensionFactory(factory, api, runtime)); return { extension, error: null }; } catch (err) { @@ -354,7 +372,7 @@ export async function loadExtensionFromFactory( ): Promise { const extension = createExtension(name, name); const api = new ConcreteExtensionAPI(PiCodingAgent, extension, runtime, cwd, eventBus); - await factory(api); + await runExtensionFactory(factory, api, runtime); return extension; } diff --git a/packages/coding-agent/test/extension-provider-registration-rollback.test.ts b/packages/coding-agent/test/extension-provider-registration-rollback.test.ts new file mode 100644 index 000000000..56e1a5bfc --- /dev/null +++ b/packages/coding-agent/test/extension-provider-registration-rollback.test.ts @@ -0,0 +1,116 @@ +import { describe, expect, test } from "bun:test"; +import { ExtensionRuntime, loadExtensionFromFactory } from "@oh-my-pi/pi-coding-agent/extensibility/extensions/loader"; +import type { ProviderConfig } from "@oh-my-pi/pi-coding-agent/extensibility/extensions/types"; +import { EventBus } from "@oh-my-pi/pi-coding-agent/utils/event-bus"; + +const testProviderConfig: ProviderConfig = { + baseUrl: "https://example.invalid/v1", + apiKey: "TEST_PROVIDER_API_KEY", + api: "openai-completions", + models: [ + { + id: "test-model", + name: "Test Model", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 16_384, + maxTokens: 4_096, + }, + ], +}; + +describe("extension provider registration rollback", () => { + test("removes provider registrations when inline extension initialization fails", async () => { + const runtime = new ExtensionRuntime(); + const events = new EventBus(); + + await expect( + loadExtensionFromFactory( + pi => { + pi.registerProvider("should-not-survive", testProviderConfig); + throw new Error("intentional initialization failure"); + }, + process.cwd(), + events, + runtime, + "broken-inline-extension", + ), + ).rejects.toThrow("intentional initialization failure"); + + expect(runtime.pendingProviderRegistrations).toEqual([]); + }); + + test("preserves provider registrations from earlier successful extensions", async () => { + const runtime = new ExtensionRuntime(); + const events = new EventBus(); + + await loadExtensionFromFactory( + pi => { + pi.registerProvider("working-provider", testProviderConfig); + }, + process.cwd(), + events, + runtime, + "working-extension", + ); + + await expect( + loadExtensionFromFactory( + pi => { + pi.registerProvider("broken-provider", testProviderConfig); + throw new Error("second extension failed"); + }, + process.cwd(), + events, + runtime, + "broken-extension", + ), + ).rejects.toThrow("second extension failed"); + + expect(runtime.pendingProviderRegistrations.map(r => r.name)).toEqual(["working-provider"]); + }); + + test("keeps provider registrations when extension initialization succeeds", async () => { + const runtime = new ExtensionRuntime(); + const events = new EventBus(); + + await loadExtensionFromFactory( + pi => { + pi.registerProvider("provider-one", { + baseUrl: "https://one.example.invalid/v1", + }); + pi.registerProvider("provider-two", { + baseUrl: "https://two.example.invalid/v1", + }); + }, + process.cwd(), + events, + runtime, + "working-extension", + ); + + expect(runtime.pendingProviderRegistrations.map(r => r.name)).toEqual(["provider-one", "provider-two"]); + }); + + test("rolls back every provider added by the failed extension", async () => { + const runtime = new ExtensionRuntime(); + const events = new EventBus(); + + await expect( + loadExtensionFromFactory( + pi => { + pi.registerProvider("broken-provider-one", testProviderConfig); + pi.registerProvider("broken-provider-two", testProviderConfig); + throw new Error("failed after multiple registrations"); + }, + process.cwd(), + events, + runtime, + "broken-multi-provider-extension", + ), + ).rejects.toThrow("failed after multiple registrations"); + + expect(runtime.pendingProviderRegistrations).toEqual([]); + }); +});