fix(extensions): roll back provider registrations on load failure

This commit is contained in:
Mustaqeem66
2026-08-06 20:00:29 +00:00
parent 3a8591a8af
commit 24a334376e
2 changed files with 138 additions and 4 deletions
@@ -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<void> {
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<Extension> {
const extension = createExtension(name, name);
const api = new ConcreteExtensionAPI(PiCodingAgent, extension, runtime, cwd, eventBus);
await factory(api);
await runExtensionFactory(factory, api, runtime);
return extension;
}
@@ -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([]);
});
});