From 879707bd32c6f59d5911d861fda21c082f2ee108 Mon Sep 17 00:00:00 2001 From: roboomp Date: Mon, 27 Jul 2026 10:01:42 +0000 Subject: [PATCH] fix(coding-agent): preserved plan subagent providers - Skipped extension-source reconciliation when restricted sessions intentionally load no extensions. - Added a shared-registry regression covering the provider model, credential, and custom API. Fixes #6783 --- packages/coding-agent/CHANGELOG.md | 5 + packages/coding-agent/src/sdk.ts | 10 +- .../sdk-restricted-extension-provider.test.ts | 112 ++++++++++++++++++ 3 files changed, 123 insertions(+), 4 deletions(-) create mode 100644 packages/coding-agent/test/sdk-restricted-extension-provider.test.ts diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 8278d1171..913b7e357 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -2,6 +2,11 @@ ## [Unreleased] +### Fixed + +- Fixed plan-mode task subagents unregistering extension-provided models, credentials, managers, and custom APIs from the shared parent `ModelRegistry` when restricted sessions intentionally skip extension loading ([#6783](https://github.com/can1357/oh-my-pi/issues/6783)). + + ## [17.1.5] - 2026-07-27 ### Added diff --git a/packages/coding-agent/src/sdk.ts b/packages/coding-agent/src/sdk.ts index 24c63d3e6..2d84dfbc3 100644 --- a/packages/coding-agent/src/sdk.ts +++ b/packages/coding-agent/src/sdk.ts @@ -1985,10 +1985,12 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} // Process provider registrations queued during extension loading. // This must happen before the runner is created so that models registered by // extensions are available for model selection on session resume / fallback. - const activeExtensionSources = extensionsResult.extensions.map(extension => extension.path); - modelRegistry.syncExtensionSources(activeExtensionSources); - for (const sourceId of new Set(activeExtensionSources)) { - modelRegistry.clearSourceRegistrations(sourceId); + if (!restrictToolNames) { + const activeExtensionSources = extensionsResult.extensions.map(extension => extension.path); + modelRegistry.syncExtensionSources(activeExtensionSources); + for (const sourceId of new Set(activeExtensionSources)) { + modelRegistry.clearSourceRegistrations(sourceId); + } } if (extensionsResult.runtime.pendingProviderRegistrations.length > 0) { for (const { name, config, sourceId } of extensionsResult.runtime.pendingProviderRegistrations) { diff --git a/packages/coding-agent/test/sdk-restricted-extension-provider.test.ts b/packages/coding-agent/test/sdk-restricted-extension-provider.test.ts new file mode 100644 index 000000000..dcee45ebb --- /dev/null +++ b/packages/coding-agent/test/sdk-restricted-extension-provider.test.ts @@ -0,0 +1,112 @@ +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; +import * as fs from "node:fs"; +import * as os from "node:os"; +import * as path from "node:path"; +import { createAssistantMessageEventStream, getCustomApi } from "@oh-my-pi/pi-ai"; +import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; +import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; +import { + type CreateAgentSessionOptions, + createAgentSession, + type ExtensionFactory, +} from "@oh-my-pi/pi-coding-agent/sdk"; +import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; +import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; +import { removeSyncWithRetries, Snowflake } from "@oh-my-pi/pi-utils"; + +const providerName = "restricted-session-provider"; +const modelId = "restricted-session-model"; +const apiId = "restricted-session-api"; +const sourceId = ""; + +describe("restricted sessions sharing extension providers", () => { + let tempDir: string; + let authStorage: AuthStorage; + let modelRegistry: ModelRegistry; + let settings: Settings; + + beforeEach(async () => { + tempDir = path.join(os.tmpdir(), `pi-sdk-restricted-provider-${Snowflake.next()}`); + fs.mkdirSync(tempDir, { recursive: true }); + authStorage = await AuthStorage.create(path.join(tempDir, "auth.db")); + modelRegistry = new ModelRegistry(authStorage, path.join(tempDir, "models.yml")); + settings = Settings.isolated(); + settings.setModelRole("default", `${providerName}/${modelId}`); + }); + + afterEach(() => { + modelRegistry.clearSourceRegistrations(sourceId); + authStorage.close(); + removeSyncWithRetries(tempDir); + }); + + const providerExtension: ExtensionFactory = pi => { + pi.registerProvider(providerName, { + baseUrl: "https://runtime.example.com/v1", + apiKey: "RUNTIME_KEY", + api: apiId, + streamSimple: () => createAssistantMessageEventStream(), + models: [ + { + id: modelId, + name: "Restricted Session Model", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128000, + maxTokens: 8192, + }, + ], + }); + }; + + function createOptions(): CreateAgentSessionOptions { + return { + cwd: tempDir, + agentDir: tempDir, + authStorage, + modelRegistry, + settings, + sessionManager: SessionManager.inMemory(), + disableExtensionDiscovery: true, + skills: [], + contextFiles: [], + promptTemplates: [], + slashCommands: [], + enableMCP: false, + enableLsp: false, + skipPythonPreflight: true, + }; + } + + test("does not unregister the parent's provider when extension loading is restricted", async () => { + const { session: parent } = await createAgentSession({ + ...createOptions(), + extensions: [providerExtension], + }); + + try { + expect(parent.model?.provider).toBe(providerName); + expect(modelRegistry.authStorage.hasAuth(providerName)).toBe(true); + expect(getCustomApi(apiId)).toBeDefined(); + + const { session: child } = await createAgentSession({ + ...createOptions(), + model: parent.model, + restrictToolNames: true, + toolNames: ["read"], + }); + + try { + expect(child.model?.provider).toBe(providerName); + expect(modelRegistry.find(providerName, modelId)).toBeDefined(); + expect(modelRegistry.authStorage.hasAuth(providerName)).toBe(true); + expect(getCustomApi(apiId)).toBeDefined(); + } finally { + await child.dispose(); + } + } finally { + await parent.dispose(); + } + }); +});