fix(extensions): roll back provider registrations on load failure
This commit is contained in:
@@ -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([]);
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user