From 0aba6d4fe7ad58df6aa1bf3902111550449bf932 Mon Sep 17 00:00:00 2001 From: Manoel Aranda Neto Date: Sun, 2 Aug 2026 13:13:36 +0200 Subject: [PATCH 01/12] feat(desktop): add Pi extension support --- .../desktop-pi-rpc-client-factory.test.ts | 5 +- .../desktop-pi-rpc-client-factory.ts | 1 + .../desktop-pi-runtime-factory.test.ts | 10 +- .../desktop-pi-runtime-factory.ts | 1 + products/desktop/docs/PI-EXTENSIONS.md | 51 + .../packages/agent/src/pi/rpc-client.test.ts | 93 ++ .../packages/agent/src/pi/rpc-client.ts | 32 +- .../desktop/packages/agent/src/pi/rpc-host.ts | 6 + .../packages/agent/src/pi/runtime.test.ts | 25 + .../desktop/packages/agent/src/pi/runtime.ts | 27 +- .../desktop/packages/agent/src/pi/types.ts | 42 + .../packages/core/src/pi-runtime/piRunner.ts | 2 + .../pi-runtime/piSessionController.test.ts | 900 ++++++++++++++++++ .../src/pi-runtime/piSessionController.ts | 640 ++++++++++++- .../core/src/pi-runtime/piSessionStore.ts | 38 + .../src/task-detail/taskCreationSaga.test.ts | 1 + .../core/src/task-detail/taskCreationSaga.ts | 2 + .../core/src/task-detail/taskService.test.ts | 38 + .../core/src/task-detail/taskService.ts | 3 +- .../desktop/packages/harness/package.json | 4 + .../harness/src/project-trust.test.ts | 78 ++ .../packages/harness/src/project-trust.ts | 42 + .../packages/harness/src/runtime.test.ts | 43 +- .../desktop/packages/harness/src/runtime.ts | 14 +- .../desktop/packages/harness/tsup.config.ts | 1 + .../host-router/src/pi-session-factory.ts | 46 + .../src/routers/pi-session.router.ts | 71 ++ .../pi-sessions/PiExtensionDialog.test.tsx | 217 +++++ .../pi-sessions/PiExtensionDialog.tsx | 185 ++++ .../pi-sessions/PiExtensionSurfaces.tsx | 57 ++ .../pi-sessions/PiProjectTrustBanner.test.tsx | 76 ++ .../pi-sessions/PiProjectTrustBanner.tsx | 103 ++ .../features/pi-sessions/PiSessionView.tsx | 118 ++- .../pi-sessions/piExtensionEditorText.ts | 8 + .../src/services/pi-session/identifiers.ts | 6 +- .../services/pi-session/pi-session.test.ts | 458 ++++++++- .../src/services/pi-session/pi-session.ts | 377 +++++++- .../src/services/pi-session/schemas.ts | 144 +++ 38 files changed, 3930 insertions(+), 35 deletions(-) create mode 100644 products/desktop/docs/PI-EXTENSIONS.md create mode 100644 products/desktop/packages/harness/src/project-trust.test.ts create mode 100644 products/desktop/packages/harness/src/project-trust.ts create mode 100644 products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.test.tsx create mode 100644 products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.tsx create mode 100644 products/desktop/packages/ui/src/features/pi-sessions/PiExtensionSurfaces.tsx create mode 100644 products/desktop/packages/ui/src/features/pi-sessions/PiProjectTrustBanner.test.tsx create mode 100644 products/desktop/packages/ui/src/features/pi-sessions/PiProjectTrustBanner.tsx create mode 100644 products/desktop/packages/ui/src/features/pi-sessions/piExtensionEditorText.ts diff --git a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.test.ts b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.test.ts index 323237e3efb3..34f2f02e37e9 100644 --- a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.test.ts +++ b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.test.ts @@ -27,12 +27,15 @@ describe("DesktopPiRpcClientFactory", () => { createPiRpcClient.mockReturnValue(client); const factory = new DesktopPiRpcClientFactory(auth, authProxy); - await expect(factory.create({ cwd: "/workspace" })).resolves.toBe(client); + await expect( + factory.create({ cwd: "/workspace", projectTrusted: true }), + ).resolves.toBe(client); expect(authProxy.start).toHaveBeenCalledWith( getLlmGatewayUrl(getCloudUrlFromRegion("eu")), ); expect(createPiRpcClient).toHaveBeenCalledWith({ cwd: "/workspace", + projectTrusted: true, providerOptions: { region: "eu", baseUrl: "http://127.0.0.1:1234", diff --git a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.ts b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.ts index 37d36f54dcf2..24ebe29f9058 100644 --- a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.ts +++ b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-rpc-client-factory.ts @@ -28,6 +28,7 @@ export class DesktopPiRpcClientFactory implements PiRpcClientFactory { cwd: string; model?: string; sessionFile?: string; + projectTrusted?: boolean; }): Promise { const credentials = await this.auth.getOAuthCredentials(); if (!credentials) { diff --git a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.test.ts b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.test.ts index d9b7a19f3fb1..e575d6e9fdae 100644 --- a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.test.ts +++ b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.test.ts @@ -11,10 +11,16 @@ describe("DesktopPiRuntimeFactory", () => { } as unknown as PiRpcClientFactory; const factory = new DesktopPiRuntimeFactory(clientFactory); - const runtime = await factory.create({ cwd: "/workspace" }); + const runtime = await factory.create({ + cwd: "/workspace", + projectTrusted: true, + }); expect(runtime).toBeInstanceOf(PiRuntime); expect(runtime.client).toBe(client); - expect(clientFactory.create).toHaveBeenCalledWith({ cwd: "/workspace" }); + expect(clientFactory.create).toHaveBeenCalledWith({ + cwd: "/workspace", + projectTrusted: true, + }); }); }); diff --git a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.ts b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.ts index b085050ff5d1..c97e3c44a862 100644 --- a/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.ts +++ b/products/desktop/apps/code/src/main/platform-adapters/desktop-pi-runtime-factory.ts @@ -17,6 +17,7 @@ export class DesktopPiRuntimeFactory implements PiRuntimeFactory { cwd: string; model?: string; sessionFile?: string; + projectTrusted?: boolean; }): Promise { const client = await this.clientFactory.create(input); return new PiRuntime(client); diff --git a/products/desktop/docs/PI-EXTENSIONS.md b/products/desktop/docs/PI-EXTENSIONS.md new file mode 100644 index 000000000000..5995aef5e242 --- /dev/null +++ b/products/desktop/docs/PI-EXTENSIONS.md @@ -0,0 +1,51 @@ +# Pi extensions in PostHog Code + +PostHog Code desktop supports Pi's existing extension and package model for local Pi sessions. This is a bridge to Pi's RPC extension protocol, not a general PostHog Code plugin API: extensions run inside the local Pi subprocess and do not load arbitrary React or Electron main-process code. + +Cloud Pi sessions do not load local extensions. + +## Install and discovery + +The bundled Pi CLI is branded as `hog`. Install a package globally with: + +```bash +hog install npm:@scope/package +hog install git:github.com/owner/repository@v1 +hog install https://github.com/owner/repository +``` + +Global installs are recorded in `~/.pi/agent/settings.json`. npm packages are installed under `~/.pi/agent/npm/`, and git packages under `~/.pi/agent/git/`. Standalone global extensions can also be placed at: + +- `~/.pi/agent/extensions/*.ts` +- `~/.pi/agent/extensions/*/index.ts` + +Restart or resume a local Pi session after changing installed resources. Package-provided extensions, skills, prompt templates, and other Pi resources use Pi's normal global discovery rules. + +Project-local resources can be placed in Pi's normal repository locations, including `.pi/extensions/`, `.pi/settings.json`, `.pi/skills/`, and `.pi/prompts/`. They remain disabled until the repository is explicitly trusted in the local Pi session. When project resources are detected, PostHog Code shows a trust control above the composer. + +Trust decisions are persisted with Pi's native project trust store at `~/.pi/agent/trust.json`. Trust is associated with the registered main repository, so the same decision applies to PostHog Code's managed worktrees for that repository. Trusting or revoking trust restarts the local Pi runtime so its resource set is rebuilt while preserving the task's native session. Revoking trust disables project-local resources on that restart. + +## Security warning + +**Pi extensions and packages run with full system permissions as your user.** They can execute arbitrary code, access files outside the current repository, start processes, read available credentials, and use the network. Skills can also instruct the model to perform arbitrary actions. Review the complete source and dependency tree before trusting a repository or installing third-party packages, and pin versions or git revisions where possible. + +Repository trust is opt-in and local-only. Do not trust a repository merely to dismiss the warning; inspect its `.pi` resources and dependencies first. + +## Supported behavior + +The desktop RPC bridge preserves Pi's native extension behavior, including: + +- extension tools and lifecycle hooks +- slash commands +- global and trusted project resources and extension state +- custom text messages +- selection, confirmation, single-line input, and multiline editor dialogs +- notifications, compact session statuses, text widgets above or below the composer, task-view document titles, and composer replacement/prefill + +Dialog requests are shown one at a time per task. They can be cancelled, and timed requests are removed when their Pi timeout expires. Extension failures are non-fatal notifications and do not fail the agent run. + +## RPC versus TUI + +PostHog Code runs Pi in RPC mode, not in Pi's terminal UI. Only behavior represented by Pi's RPC protocol can cross the boundary. In particular, PostHog Code does not support arbitrary Pi TUI components or rendering factories, raw terminal input handlers, custom headers/footers/themes, custom tool-call renderers, editor component replacement, editor autocomplete providers, or synchronous reads of the current composer text. Widgets are text lines only. + +An extension that depends on those TUI-only APIs may still load, but those visual or terminal-specific portions are ignored by Pi RPC mode. Use the supported RPC dialog and fire-and-forget UI methods for an extension intended to work in PostHog Code. diff --git a/products/desktop/packages/agent/src/pi/rpc-client.test.ts b/products/desktop/packages/agent/src/pi/rpc-client.test.ts index eb0abae742ae..eb2e8d57357f 100644 --- a/products/desktop/packages/agent/src/pi/rpc-client.test.ts +++ b/products/desktop/packages/agent/src/pi/rpc-client.test.ts @@ -31,6 +31,42 @@ describe("createPiRpcClient", () => { ).toBeUndefined(); }); + it("passes repository trust over the private bootstrap pipe", async () => { + const directory = await mkdtemp(join(tmpdir(), "pi-project-trust-")); + const hostPath = join(directory, "host.mjs"); + const capturePath = join(directory, "bootstrap.json"); + await writeFile( + hostPath, + ` +import { readFileSync, writeFileSync } from "node:fs"; + +writeFileSync(${JSON.stringify(capturePath)}, readFileSync(3, "utf8")); +process.stdin.resume(); +`, + ); + const client = createPiRpcClient({ + cliPath: hostPath, + cwd: directory, + projectTrusted: true, + providerOptions: { apiKey: "proxy-key" }, + }); + + try { + await client.start(); + await vi.waitFor(async () => { + await expect(readFile(capturePath, "utf8")).resolves.toBe( + JSON.stringify({ + providerOptions: { apiKey: "proxy-key" }, + projectTrusted: true, + }), + ); + }); + } finally { + await client.stop(); + await rm(directory, { recursive: true }); + } + }); + it("runs the RPC host with Electron's Node mode enabled", async () => { const directory = await mkdtemp(join(tmpdir(), "pi-electron-node-mode-")); const hostPath = join(directory, "host.mjs"); @@ -62,6 +98,63 @@ process.stdin.resume(); } }); + it("intercepts extension UI requests and writes responses on the Pi wire", async () => { + const directory = await mkdtemp(join(tmpdir(), "pi-extension-ui-")); + const hostPath = join(directory, "host.mjs"); + const capturePath = join(directory, "response.json"); + await writeFile( + hostPath, + ` +import { closeSync, writeFileSync } from "node:fs"; +import { createInterface } from "node:readline"; + +closeSync(3); +process.stdout.write(JSON.stringify({ + type: "extension_ui_request", + id: "extension-1", + method: "input", + title: "Your name", +}) + "\\n"); +createInterface({ input: process.stdin }).on("line", (line) => { + writeFileSync(${JSON.stringify(capturePath)}, line); +}); +`, + ); + const client = createPiRpcClient({ + cliPath: hostPath, + cwd: directory, + providerOptions: { apiKey: "proxy-key" }, + }); + const request = new Promise((resolve) => client.onEvent(resolve)); + + try { + await client.start(); + await expect(request).resolves.toEqual({ + type: "extension_ui_request", + id: "extension-1", + method: "input", + title: "Your name", + }); + await client.respondToExtensionUI({ + type: "extension_ui_response", + id: "extension-1", + value: "Ada", + }); + await vi.waitFor(async () => { + await expect(readFile(capturePath, "utf8")).resolves.toBe( + JSON.stringify({ + type: "extension_ui_response", + id: "extension-1", + value: "Ada", + }), + ); + }); + } finally { + await client.stop(); + await rm(directory, { recursive: true }); + } + }); + it("uses the private host channel without changing Pi RPC", async () => { const directory = await mkdtemp(join(tmpdir(), "pi-host-channel-")); const hostPath = join(directory, "host.mjs"); diff --git a/products/desktop/packages/agent/src/pi/rpc-client.ts b/products/desktop/packages/agent/src/pi/rpc-client.ts index 757ac8fca901..ab3746f5aa86 100644 --- a/products/desktop/packages/agent/src/pi/rpc-client.ts +++ b/products/desktop/packages/agent/src/pi/rpc-client.ts @@ -8,11 +8,12 @@ import { type RpcClientOptions, } from "@earendil-works/pi-coding-agent"; import { safePiEnvironment } from "./rpc-environment"; -import type { PiQueueSnapshot } from "./types"; +import type { PiExtensionUIResponse, PiQueueSnapshot } from "./types"; export type PiRpcClient = RpcClient & { getQueue(): Promise; clearQueue(): Promise; + respondToExtensionUI(response: PiExtensionUIResponse): Promise; }; export interface PiRpcProviderOptions { @@ -84,6 +85,7 @@ class SecurePiRpcClient extends RpcClient { constructor( private readonly secureOptions: RpcClientOptions, private readonly providerOptions: PiRpcProviderOptions, + private readonly projectTrusted: boolean, ) { super(secureOptions); } @@ -162,7 +164,10 @@ class SecurePiRpcClient extends RpcClient { const bootstrapPipe = child.stdio[3] as Writable | null; bootstrapPipe?.on("error", () => {}); bootstrapPipe?.end( - JSON.stringify({ providerOptions: this.providerOptions }), + JSON.stringify({ + providerOptions: this.providerOptions, + projectTrusted: this.projectTrusted, + }), ); await new Promise((resolve) => setTimeout(resolve, 100)); @@ -182,6 +187,24 @@ class SecurePiRpcClient extends RpcClient { return this.sendHostRequest("clear_queue"); } + respondToExtensionUI(response: PiExtensionUIResponse): Promise { + const child = (this as unknown as RpcClientInternals).process; + const stdin = child?.stdin; + if (!child || !stdin || stdin.destroyed || !stdin.writable) { + return Promise.reject(new Error("Pi RPC client is not writable")); + } + + return new Promise((resolve, reject) => { + stdin.write(`${JSON.stringify(response)}\n`, (error) => { + if (error) { + reject(error); + } else { + resolve(); + } + }); + }); + } + private sendHostRequest( method: PiHostRequest["method"], ): Promise { @@ -265,10 +288,12 @@ export type PiRpcClientOptions = Pick< > & { sessionFile?: string; providerOptions: PiRpcProviderOptions; + projectTrusted?: boolean; }; export function createPiRpcClient(options: PiRpcClientOptions): PiRpcClient { - const { sessionFile, providerOptions, ...rpcOptions } = options; + const { sessionFile, providerOptions, projectTrusted, ...rpcOptions } = + options; const args = sessionFile ? ["--session-file", sessionFile] : []; const cliPath = rpcOptions.cliPath ?? @@ -281,5 +306,6 @@ export function createPiRpcClient(options: PiRpcClientOptions): PiRpcClient { provider: "posthog", }, providerOptions, + projectTrusted ?? false, ); } diff --git a/products/desktop/packages/agent/src/pi/rpc-host.ts b/products/desktop/packages/agent/src/pi/rpc-host.ts index 9516177b4e66..5224a1cf7423 100644 --- a/products/desktop/packages/agent/src/pi/rpc-host.ts +++ b/products/desktop/packages/agent/src/pi/rpc-host.ts @@ -2,6 +2,7 @@ import { readFileSync } from "node:fs"; import { SessionManager } from "@earendil-works/pi-coding-agent"; import { createHarnessRuntime, runRpcMode } from "@posthog/harness"; import type { PosthogProviderOptions } from "@posthog/harness/extensions/posthog-provider/provider"; +import { createPiProjectTrustResolver } from "@posthog/harness/project-trust"; import { POSTHOG_PI_QUEUE_ENTRY_TYPE, readPersistedPiQueue, @@ -10,6 +11,7 @@ import { sanitizePiHostEnvironment } from "./rpc-environment"; interface PiRpcBootstrap { providerOptions?: PosthogProviderOptions; + projectTrusted?: boolean; } interface PiHostRequest { @@ -38,6 +40,10 @@ const sessionManager = sessionFile const runtime = await createHarnessRuntime({ cwd, sessionManager, + projectTrusted: createPiProjectTrustResolver( + cwd, + bootstrap.projectTrusted ?? false, + ), ...providerOptions, }); diff --git a/products/desktop/packages/agent/src/pi/runtime.test.ts b/products/desktop/packages/agent/src/pi/runtime.test.ts index e1c82972fa1b..cf1b87cc50ba 100644 --- a/products/desktop/packages/agent/src/pi/runtime.test.ts +++ b/products/desktop/packages/agent/src/pi/runtime.test.ts @@ -244,6 +244,31 @@ describe("PiRuntime", () => { }); }); + it("routes extension UI and errors outside the conversation stream", () => { + const { client, emit } = createClient(); + const runtime = new PiRuntime(client); + const extensionListener = vi.fn(); + const conversationListener = vi.fn(); + runtime.onExtensionEvent(extensionListener); + runtime.onConversationEvent(conversationListener); + + emit({ + type: "extension_ui_request", + id: "extension-1", + method: "notify", + message: "Done", + } as unknown as AgentSessionEvent); + emit({ + type: "extension_error", + extensionPath: "/extensions/example.ts", + event: "tool_call", + error: "boom", + } as unknown as AgentSessionEvent); + + expect(extensionListener).toHaveBeenCalledTimes(2); + expect(conversationListener).not.toHaveBeenCalled(); + }); + it("normalizes live Pi events before forwarding them", () => { const { client, emit } = createClient(); const runtime = new PiRuntime(client); diff --git a/products/desktop/packages/agent/src/pi/runtime.ts b/products/desktop/packages/agent/src/pi/runtime.ts index 31a2a776a364..9ab4211858b6 100644 --- a/products/desktop/packages/agent/src/pi/runtime.ts +++ b/products/desktop/packages/agent/src/pi/runtime.ts @@ -11,6 +11,7 @@ import { } from "./conversation/translatePiConversation"; import { getPiRpcClientProcess, type PiRpcClient } from "./rpc-client"; import { sendPiRpcCommand } from "./rpc-transport"; +import type { PiExtensionWireEvent } from "./types"; export class PiRuntime { readonly client: PiRpcClient; @@ -22,6 +23,9 @@ export class PiRuntime { private readonly conversationListeners = new Set< (event: AgentConversationEvent) => void >(); + private readonly extensionListeners = new Set< + (event: PiExtensionWireEvent) => void + >(); private readonly pendingUserMessages: Array<{ id: string; message: string; @@ -32,7 +36,9 @@ export class PiRuntime { constructor(client: PiRpcClient) { this.client = client; this.translator = createPiConversationTranslator(); - client.onEvent((event) => this.handleEvent(event)); + client.onEvent((event) => + this.handleEvent(event as AgentSessionEvent | PiExtensionWireEvent), + ); } get process() { @@ -51,6 +57,13 @@ export class PiRuntime { return () => this.conversationListeners.delete(listener); } + onExtensionEvent( + listener: (event: PiExtensionWireEvent) => void, + ): () => void { + this.extensionListeners.add(listener); + return () => this.extensionListeners.delete(listener); + } + async sendCommand(command: RpcCommand): Promise { const isUserMessage = command.type === "prompt" || @@ -113,7 +126,17 @@ export class PiRuntime { } } - private handleEvent(event: AgentSessionEvent): void { + private handleEvent(event: AgentSessionEvent | PiExtensionWireEvent): void { + if ( + event.type === "extension_ui_request" || + event.type === "extension_error" + ) { + for (const listener of this.extensionListeners) { + listener(event); + } + return; + } + for (const listener of this.runtimeListeners) { listener(event); } diff --git a/products/desktop/packages/agent/src/pi/types.ts b/products/desktop/packages/agent/src/pi/types.ts index 1771097372bd..0ed178ba45af 100644 --- a/products/desktop/packages/agent/src/pi/types.ts +++ b/products/desktop/packages/agent/src/pi/types.ts @@ -1,6 +1,8 @@ import type { ThinkingLevel } from "@earendil-works/pi-agent-core"; import type { RpcClient, + RpcExtensionUIRequest, + RpcExtensionUIResponse, RpcSessionState, } from "@earendil-works/pi-coding-agent"; @@ -42,4 +44,44 @@ export interface PiQueueSnapshot { followUp: string[]; } +export type PiExtensionUIRequest = RpcExtensionUIRequest; +export type PiExtensionUIResponse = RpcExtensionUIResponse; + +export interface PiExtensionError { + type: "extension_error"; + extensionPath: string; + event: string; + error: string; +} + +export interface PiExtensionSessionReset { + type: "extension_session_reset"; +} + +export interface PiExtensionDialogExpired { + type: "extension_dialog_expired"; + id: string; +} + +export interface PiExtensionStateSnapshot { + type: "extension_state_snapshot"; + dialogs: Array< + Extract< + PiExtensionUIRequest, + { method: "select" | "confirm" | "input" | "editor" } + > + >; + statuses: Array>; + widgets: Array>; + title?: Extract; + editorText?: Extract; +} + +export type PiExtensionWireEvent = PiExtensionUIRequest | PiExtensionError; +export type PiExtensionEvent = + | PiExtensionWireEvent + | PiExtensionSessionReset + | PiExtensionDialogExpired + | PiExtensionStateSnapshot; + export type PiSessionStats = Awaited>; diff --git a/products/desktop/packages/core/src/pi-runtime/piRunner.ts b/products/desktop/packages/core/src/pi-runtime/piRunner.ts index 40fb9574a8e9..ffb909e84f05 100644 --- a/products/desktop/packages/core/src/pi-runtime/piRunner.ts +++ b/products/desktop/packages/core/src/pi-runtime/piRunner.ts @@ -3,6 +3,7 @@ import type { PiThinkingLevel } from "@posthog/agent/pi/types"; export interface PiRunInput { taskId: string; cwd: string; + projectTrustPath?: string; prompt: string; model?: string; thinkingLevel?: PiThinkingLevel; @@ -11,6 +12,7 @@ export interface PiRunInput { export interface PiResumeInput { taskId: string; cwd: string; + projectTrustPath?: string; } export interface PiRunner { diff --git a/products/desktop/packages/core/src/pi-runtime/piSessionController.test.ts b/products/desktop/packages/core/src/pi-runtime/piSessionController.test.ts index 37df1cae8eb3..dd9f4dcf7e35 100644 --- a/products/desktop/packages/core/src/pi-runtime/piSessionController.test.ts +++ b/products/desktop/packages/core/src/pi-runtime/piSessionController.test.ts @@ -146,6 +146,906 @@ describe("PiSessionController", () => { expect(client.getConversation).not.toHaveBeenCalled(); }); + it("loads repository trust and reconnects after changing it", async () => { + let trusted = false; + const session = createSession(); + session.getProjectTrust = vi.fn(async () => ({ + trusted, + hasProjectResources: true, + })); + session.setProjectTrusted = vi.fn(async (nextTrusted) => { + trusted = nextTrusted; + }); + const controller = createController(session); + + await controller.connect("task-1"); + expect(controller.store.getState().sessions["task-1"].projectTrust).toEqual( + { + trusted: false, + hasProjectResources: true, + }, + ); + controller.store.setState((state) => ({ + sessions: { + ...state.sessions, + "task-1": { + ...state.sessions["task-1"], + queue: { steering: ["queued"], followUp: [] }, + }, + }, + })); + + await controller.setProjectTrusted("task-1", true); + + expect(session.setProjectTrusted).toHaveBeenCalledWith(true); + expect(session.client.prompt).toHaveBeenCalledWith("queued"); + expect(controller.store.getState().sessions["task-1"].projectTrust).toEqual( + { + trusted: true, + hasProjectResources: true, + }, + ); + }); + + it("shares an in-flight repository trust change and rejects an opposite toggle", async () => { + let finishTransition: (() => void) | undefined; + const session = createSession(); + session.setProjectTrusted = vi.fn( + () => + new Promise((resolve) => { + finishTransition = resolve; + }), + ); + const controller = createController(session); + await controller.connect("task-1"); + + const first = controller.setProjectTrusted("task-1", true); + const duplicate = controller.setProjectTrusted("task-1", true); + await expect(controller.setProjectTrusted("task-1", false)).rejects.toThrow( + "already in progress", + ); + expect(session.setProjectTrusted).toHaveBeenCalledOnce(); + + finishTransition?.(); + await expect(Promise.all([first, duplicate])).resolves.toEqual([ + undefined, + undefined, + ]); + }); + + it.each([ + { streaming: true, bashRunning: false }, + { streaming: false, bashRunning: true }, + ])( + "rejects repository trust changes while Pi is busy", + async ({ streaming, bashRunning }) => { + const session = createSession(); + session.setProjectTrusted = vi.fn(async () => {}); + const controller = createController(session); + await controller.connect("task-1"); + const current = controller.store.getState().sessions["task-1"]; + const status = current.status; + if (!status) { + throw new Error("Expected connected Pi status"); + } + controller.store.setState((state) => ({ + sessions: { + ...state.sessions, + "task-1": { + ...state.sessions["task-1"], + status: { + ...status, + isStreaming: streaming, + }, + isBashRunning: bashRunning, + }, + }, + })); + + await expect( + controller.setProjectTrusted("task-1", true), + ).rejects.toBeInstanceOf(PiOperationError); + expect(session.setProjectTrusted).not.toHaveBeenCalled(); + }, + ); + + it("queues extension dialogs and applies replacement, removal, and error events", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + const session = createSession(); + session.respondToExtensionUI = vi.fn(async () => {}); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return () => {}; + }); + const controller = createController(session); + + await controller.connect("task-1"); + onExtensionEvent({ + type: "extension_ui_request", + id: "select-1", + method: "select", + title: "Pick one", + options: ["A", "B"], + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "confirm-1", + method: "confirm", + title: "Continue?", + message: "Are you sure?", + }); + + expect( + controller.store + .getState() + .sessions["task-1"].extensionDialogs.map((dialog) => dialog.id), + ).toEqual(["select-1", "confirm-1"]); + await controller.respondToExtensionUI("task-1", { + type: "extension_ui_response", + id: "select-1", + value: "B", + }); + await controller.cancelExtensionUI("task-1", "select-1"); + expect(session.respondToExtensionUI).toHaveBeenCalledTimes(1); + expect( + controller.store.getState().sessions["task-1"].extensionDialogs[0]?.id, + ).toBe("confirm-1"); + + onExtensionEvent({ + type: "extension_ui_request", + id: "status-1", + method: "setStatus", + statusKey: "build", + statusText: "Running", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "status-2", + method: "setStatus", + statusKey: "build", + statusText: "Done", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "widget-1", + method: "setWidget", + widgetKey: "summary", + widgetLines: ["one", "two"], + widgetPlacement: "belowEditor", + }); + expect(controller.store.getState().sessions["task-1"]).toMatchObject({ + extensionStatuses: { build: "Done" }, + extensionWidgets: { + summary: { lines: ["one", "two"], placement: "belowEditor" }, + }, + }); + + onExtensionEvent({ + type: "extension_ui_request", + id: "status-3", + method: "setStatus", + statusKey: "build", + statusText: undefined, + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "widget-2", + method: "setWidget", + widgetKey: "summary", + widgetLines: undefined, + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "notify-1", + method: "notify", + message: "Extension finished", + notifyType: "warning", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "title-1", + method: "setTitle", + title: "Extension task", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "editor-text-1", + method: "set_editor_text", + text: "replacement draft", + }); + onExtensionEvent({ + type: "extension_error", + extensionPath: "/extensions/example.ts", + event: "tool_call", + error: "boom", + }); + expect(controller.store.getState().sessions["task-1"]).toMatchObject({ + extensionStatuses: {}, + extensionWidgets: {}, + extensionTitle: "Extension task", + extensionEditorText: { + id: "editor-text-1", + text: "replacement draft", + }, + error: undefined, + extensionNotifications: [ + { + id: "notify-1", + notifyType: "warning", + message: "Extension finished", + }, + expect.objectContaining({ + notifyType: "error", + message: expect.stringContaining("boom"), + }), + ], + }); + }); + + it("applies unconsumed state from an authoritative extension snapshot", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + const session = createSession(); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return () => {}; + }); + const controller = createController(session); + await controller.connect("task-1"); + + onExtensionEvent({ + type: "extension_state_snapshot", + dialogs: [ + { + type: "extension_ui_request", + id: "dialog-1", + method: "confirm", + title: "Continue?", + message: "Proceed?", + }, + ], + statuses: [ + { + type: "extension_ui_request", + id: "status-1", + method: "setStatus", + statusKey: "build", + statusText: "Running", + }, + ], + widgets: [ + { + type: "extension_ui_request", + id: "widget-1", + method: "setWidget", + widgetKey: "summary", + widgetLines: ["content"], + }, + ], + title: { + type: "extension_ui_request", + id: "title-1", + method: "setTitle", + title: "Extension title", + }, + editorText: { + type: "extension_ui_request", + id: "editor-1", + method: "set_editor_text", + text: "draft", + }, + }); + + expect(controller.store.getState().sessions["task-1"]).toMatchObject({ + extensionDialogs: [expect.objectContaining({ id: "dialog-1" })], + extensionStatuses: { build: "Running" }, + extensionWidgets: { + summary: { lines: ["content"], placement: "aboveEditor" }, + }, + extensionTitle: "Extension title", + extensionEditorText: { id: "editor-1", text: "draft" }, + }); + }); + + it("consumes editor text locally and retries acknowledgement without replaying it", async () => { + vi.useFakeTimers(); + try { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + let onExtensionError: (error: unknown) => void = () => {}; + const session = createSession(); + session.acknowledgeExtensionEditorText = vi + .fn<() => Promise>() + .mockRejectedValueOnce(new Error("ack lost")) + .mockResolvedValueOnce(); + session.onExtensionEvent = vi.fn((handler, errorHandler) => { + onExtensionEvent = handler; + onExtensionError = errorHandler; + return () => {}; + }); + const controller = createController(session); + const connected = controller.connect("task-1"); + await vi.runAllTimersAsync(); + await connected; + onExtensionEvent({ + type: "extension_ui_request", + id: "editor-1", + method: "set_editor_text", + text: "first draft", + }); + + controller.acknowledgeExtensionEditorText("task-1", "editor-1"); + expect( + controller.store.getState().sessions["task-1"].extensionEditorText, + ).toBeUndefined(); + await vi.waitFor(() => + expect(session.acknowledgeExtensionEditorText).toHaveBeenCalledOnce(), + ); + onExtensionError(new Error("subscription closed")); + await vi.advanceTimersByTimeAsync(100); + await vi.waitFor(() => { + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2); + expect(session.acknowledgeExtensionEditorText).toHaveBeenCalledTimes(2); + }); + + onExtensionEvent({ + type: "extension_state_snapshot", + dialogs: [], + statuses: [], + widgets: [], + editorText: { + type: "extension_ui_request", + id: "editor-1", + method: "set_editor_text", + text: "first draft replay", + }, + }); + expect( + controller.store.getState().sessions["task-1"].extensionEditorText, + ).toBeUndefined(); + controller.acknowledgeExtensionEditorText("task-1", "editor-1"); + expect(session.acknowledgeExtensionEditorText).toHaveBeenCalledTimes(2); + onExtensionEvent({ type: "extension_session_reset" }); + } finally { + vi.useRealTimers(); + } + }); + + it("removes an expired extension dialog without sending a late response", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + const session = createSession(); + session.respondToExtensionUI = vi.fn(async () => {}); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return () => {}; + }); + const controller = createController(session); + await controller.connect("task-1"); + + onExtensionEvent({ + type: "extension_ui_request", + id: "input-1", + method: "input", + title: "Name", + timeout: 50, + }); + onExtensionEvent({ type: "extension_dialog_expired", id: "input-1" }); + + expect( + controller.store.getState().sessions["task-1"].extensionDialogs, + ).toEqual([]); + expect(session.respondToExtensionUI).not.toHaveBeenCalled(); + }); + + it("keeps a dialog available after response failure and retries successfully", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + const session = createSession(); + session.respondToExtensionUI = vi + .fn<() => Promise>() + .mockRejectedValueOnce(new Error("wire unavailable")) + .mockResolvedValueOnce(); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return () => {}; + }); + const controller = createController(session); + await controller.connect("task-1"); + onExtensionEvent({ + type: "extension_ui_request", + id: "confirm-1", + method: "confirm", + title: "Continue?", + message: "Proceed?", + }); + const response = { + type: "extension_ui_response" as const, + id: "confirm-1", + confirmed: true, + }; + + await expect( + controller.respondToExtensionUI("task-1", response), + ).rejects.toThrow("wire unavailable"); + expect( + controller.store.getState().sessions["task-1"].extensionDialogs, + ).toHaveLength(1); + expect( + controller.store.getState().sessions["task-1"].extensionNotifications, + ).toEqual([ + expect.objectContaining({ + notifyType: "error", + message: expect.stringContaining("wire unavailable"), + }), + ]); + + await controller.respondToExtensionUI("task-1", response); + + expect(session.respondToExtensionUI).toHaveBeenCalledTimes(2); + expect( + controller.store.getState().sessions["task-1"].extensionDialogs, + ).toEqual([]); + }); + + it("preserves an in-flight response across subscription error and replay", async () => { + vi.useFakeTimers(); + try { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + let onExtensionError: (error: unknown) => void = () => {}; + let resolveDelivery: () => void = () => {}; + const session = createSession(); + session.respondToExtensionUI = vi.fn( + () => + new Promise((resolve) => { + resolveDelivery = resolve; + }), + ); + session.onExtensionEvent = vi.fn((handler, errorHandler) => { + onExtensionEvent = handler; + onExtensionError = errorHandler; + return () => {}; + }); + const controller = createController(session); + const connected = controller.connect("task-1"); + await vi.runAllTimersAsync(); + await connected; + const request = { + type: "extension_ui_request" as const, + id: "confirm-1", + method: "confirm" as const, + title: "Continue?", + message: "Proceed?", + }; + const response = { + type: "extension_ui_response" as const, + id: "confirm-1", + confirmed: true, + }; + onExtensionEvent(request); + const delivery = controller.respondToExtensionUI("task-1", response); + await vi.waitFor(() => + expect(session.respondToExtensionUI).toHaveBeenCalledOnce(), + ); + + onExtensionError(new Error("subscription closed")); + expect( + controller.store.getState().sessions["task-1"].extensionDialogs, + ).toEqual([request]); + await vi.advanceTimersByTimeAsync(100); + await vi.waitFor(() => + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2), + ); + onExtensionEvent({ + type: "extension_state_snapshot", + dialogs: [request], + statuses: [], + widgets: [], + }); + + resolveDelivery(); + await delivery; + expect( + controller.store.getState().sessions["task-1"].extensionDialogs, + ).toEqual([]); + await controller.respondToExtensionUI("task-1", response); + expect(session.respondToExtensionUI).toHaveBeenCalledOnce(); + onExtensionEvent({ type: "extension_session_reset" }); + } finally { + vi.useRealTimers(); + } + }); + + it("does not send another response when expiry races an in-flight response", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + let resolveDelivery: () => void = () => {}; + const session = createSession(); + session.respondToExtensionUI = vi.fn( + () => + new Promise((resolve) => { + resolveDelivery = resolve; + }), + ); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return () => {}; + }); + const controller = createController(session); + await controller.connect("task-1"); + onExtensionEvent({ + type: "extension_ui_request", + id: "input-1", + method: "input", + title: "Name", + timeout: 50, + }); + + const manual = controller.respondToExtensionUI("task-1", { + type: "extension_ui_response", + id: "input-1", + value: "Ada", + }); + await vi.waitFor(() => + expect(session.respondToExtensionUI).toHaveBeenCalledOnce(), + ); + onExtensionEvent({ type: "extension_dialog_expired", id: "input-1" }); + + expect(session.respondToExtensionUI).toHaveBeenCalledTimes(1); + expect( + controller.store.getState().sessions["task-1"].extensionDialogs, + ).toEqual([]); + resolveDelivery(); + await manual; + }); + + it("keeps extension events subscribed while the task view is disconnected", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + const unsubscribeExtension = vi.fn(); + let resolveDelivery: () => void = () => {}; + const session = createSession(); + session.respondToExtensionUI = vi.fn( + () => + new Promise((resolve) => { + resolveDelivery = resolve; + }), + ); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return unsubscribeExtension; + }); + const controller = createController(session); + + await controller.connect("task-1"); + onExtensionEvent({ + type: "extension_ui_request", + id: "confirm-before-away", + method: "confirm", + title: "Continue?", + message: "Proceed?", + }); + const delivery = controller.respondToExtensionUI("task-1", { + type: "extension_ui_response", + id: "confirm-before-away", + confirmed: true, + }); + await vi.waitFor(() => + expect(session.respondToExtensionUI).toHaveBeenCalledOnce(), + ); + controller.disconnect("task-1"); + onExtensionEvent({ + type: "extension_ui_request", + id: "select-while-away", + method: "select", + title: "Choose", + options: ["A", "B"], + }); + + expect(unsubscribeExtension).not.toHaveBeenCalled(); + resolveDelivery(); + await delivery; + expect( + controller.store + .getState() + .sessions["task-1"].extensionDialogs.map((dialog) => dialog.id), + ).toEqual(["select-while-away"]); + await controller.connect("task-1"); + expect(session.onExtensionEvent).toHaveBeenCalledOnce(); + }); + + it("clears extension state and subscriptions on session replacement", async () => { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + const unsubscribeExtension = vi.fn(); + const session = createSession(); + session.onExtensionEvent = vi.fn((handler) => { + onExtensionEvent = handler; + return unsubscribeExtension; + }); + const controller = createController(session); + + await controller.connect("task-1"); + onExtensionEvent({ + type: "extension_ui_request", + id: "dialog-1", + method: "confirm", + title: "Continue?", + message: "Proceed?", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "status-1", + method: "setStatus", + statusKey: "build", + statusText: "Running", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "widget-1", + method: "setWidget", + widgetKey: "summary", + widgetLines: ["content"], + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "title-1", + method: "setTitle", + title: "Extension title", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "draft-1", + method: "set_editor_text", + text: "draft", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "notify-1", + method: "notify", + message: "notice", + }); + + await controller.connect("task-1", "run-2"); + + expect(unsubscribeExtension).toHaveBeenCalledOnce(); + expect(controller.store.getState().sessions["task-1"]).toMatchObject({ + extensionDialogs: [], + extensionNotifications: [], + extensionStatuses: {}, + extensionWidgets: {}, + extensionTitle: undefined, + extensionEditorText: undefined, + }); + }); + + it("reconnects extension events after a transient subscription error", async () => { + vi.useFakeTimers(); + try { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + let onExtensionError: (error: unknown) => void = () => {}; + const unsubscribeExtension = vi.fn(); + const session = createSession(); + session.onExtensionEvent = vi.fn((handler, errorHandler) => { + onExtensionEvent = handler; + onExtensionError = errorHandler; + return unsubscribeExtension; + }); + const controller = createController(session); + const connected = controller.connect("task-1"); + await vi.runAllTimersAsync(); + await connected; + + onExtensionError(new Error("subscription closed")); + expect(unsubscribeExtension).toHaveBeenCalledOnce(); + await vi.advanceTimersByTimeAsync(100); + await vi.waitFor(() => + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2), + ); + onExtensionEvent({ + type: "extension_ui_request", + id: "dialog-after-reconnect", + method: "input", + title: "Name", + }); + + expect( + controller.store + .getState() + .sessions["task-1"].extensionDialogs.map((dialog) => dialog.id), + ).toEqual(["dialog-after-reconnect"]); + onExtensionEvent({ type: "extension_session_reset" }); + await vi.advanceTimersByTimeAsync(1_000); + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2); + } finally { + vi.useRealTimers(); + } + }); + + it("keeps exponential backoff when reconnects only deliver snapshots", async () => { + vi.useFakeTimers(); + try { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + let onExtensionError: (error: unknown) => void = () => {}; + const session = createSession(); + session.onExtensionEvent = vi.fn((handler, errorHandler) => { + onExtensionEvent = handler; + onExtensionError = errorHandler; + return () => {}; + }); + const controller = createController(session); + const connected = controller.connect("task-1"); + await vi.runAllTimersAsync(); + await connected; + + onExtensionError(new Error("first failure")); + await vi.advanceTimersByTimeAsync(100); + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2); + + onExtensionEvent({ + type: "extension_state_snapshot", + dialogs: [], + statuses: [], + widgets: [], + }); + onExtensionError(new Error("second failure")); + await vi.advanceTimersByTimeAsync(199); + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2); + await vi.advanceTimersByTimeAsync(1); + expect(session.onExtensionEvent).toHaveBeenCalledTimes(3); + + onExtensionEvent({ + type: "extension_state_snapshot", + dialogs: [], + statuses: [], + widgets: [], + }); + onExtensionError(new Error("third failure")); + await vi.advanceTimersByTimeAsync(399); + expect(session.onExtensionEvent).toHaveBeenCalledTimes(3); + await vi.advanceTimersByTimeAsync(1); + expect(session.onExtensionEvent).toHaveBeenCalledTimes(4); + expect( + controller.store.getState().sessions["task-1"].extensionNotifications, + ).toHaveLength(1); + + controller.disconnect("task-1"); + } finally { + vi.useRealTimers(); + } + }); + + it.each([ + "timed dialog expiry", + "status and widget clear", + "session reset", + ] as const)( + "reconciles an empty snapshot after disconnected %s", + async (scenario) => { + vi.useFakeTimers(); + try { + let onExtensionEvent: ( + event: import("@posthog/agent/pi/types").PiExtensionEvent, + ) => void = () => {}; + let onExtensionError: (error: unknown) => void = () => {}; + const session = createSession(); + session.onExtensionEvent = vi.fn((handler, errorHandler) => { + onExtensionEvent = handler; + onExtensionError = errorHandler; + return () => {}; + }); + const controller = createController(session); + const connected = controller.connect("task-1"); + await vi.runAllTimersAsync(); + await connected; + + if (scenario !== "status and widget clear") { + onExtensionEvent({ + type: "extension_ui_request", + id: "dialog-1", + method: "input", + title: "Name", + }); + } + if (scenario !== "timed dialog expiry") { + onExtensionEvent({ + type: "extension_ui_request", + id: "status-1", + method: "setStatus", + statusKey: "build", + statusText: "Running", + }); + onExtensionEvent({ + type: "extension_ui_request", + id: "widget-1", + method: "setWidget", + widgetKey: "summary", + widgetLines: ["content"], + }); + } + if (scenario === "session reset") { + onExtensionEvent({ + type: "extension_ui_request", + id: "title-1", + method: "setTitle", + title: "Old title", + }); + } + + onExtensionError(new Error("subscription closed")); + const preserved = controller.store.getState().sessions["task-1"]; + expect( + preserved.extensionDialogs.length + + Object.keys(preserved.extensionStatuses).length + + Object.keys(preserved.extensionWidgets).length, + ).toBeGreaterThan(0); + await vi.advanceTimersByTimeAsync(100); + await vi.waitFor(() => + expect(session.onExtensionEvent).toHaveBeenCalledTimes(2), + ); + onExtensionEvent({ + type: "extension_state_snapshot", + dialogs: [], + statuses: [], + widgets: [], + }); + + expect(controller.store.getState().sessions["task-1"]).toMatchObject({ + extensionDialogs: [], + extensionStatuses: {}, + extensionWidgets: {}, + extensionTitle: undefined, + extensionEditorText: undefined, + extensionNotifications: [ + expect.objectContaining({ notifyType: "error" }), + ], + }); + onExtensionEvent({ type: "extension_session_reset" }); + } finally { + vi.useRealTimers(); + } + }, + ); + + it("cancels a pending extension reconnect when the view disconnects", async () => { + vi.useFakeTimers(); + try { + let onExtensionError: (error: unknown) => void = () => {}; + const session = createSession(); + session.onExtensionEvent = vi.fn((_handler, errorHandler) => { + onExtensionError = errorHandler; + return () => {}; + }); + const controller = createController(session); + const connected = controller.connect("task-1"); + await vi.runAllTimersAsync(); + await connected; + + onExtensionError(new Error("subscription closed")); + controller.disconnect("task-1"); + await vi.advanceTimersByTimeAsync(1_000); + + expect(session.onExtensionEvent).toHaveBeenCalledOnce(); + } finally { + vi.useRealTimers(); + } + }); + it("uploads cloud follow-up attachments before sending the native message", async () => { const session = createSession(); session.sendUserMessage = vi.fn(async () => {}); diff --git a/products/desktop/packages/core/src/pi-runtime/piSessionController.ts b/products/desktop/packages/core/src/pi-runtime/piSessionController.ts index f86aedb30244..6b164321e767 100644 --- a/products/desktop/packages/core/src/pi-runtime/piSessionController.ts +++ b/products/desktop/packages/core/src/pi-runtime/piSessionController.ts @@ -1,5 +1,7 @@ import type { PiRemoteRpcClient } from "@posthog/agent/pi/remote-rpc-client"; import type { + PiExtensionEvent, + PiExtensionUIResponse, PiNativeModelInfo, PiPersistedSessionConfig, PiQueueSnapshot, @@ -23,6 +25,8 @@ import { createEmptyPiControllerSession, createPiSessionStore, type PiControllerSessionState, + type PiExtensionDialogRequest, + type PiProjectTrustState, type PiSessionError, type PiSessionStore, } from "./piSessionStore"; @@ -54,6 +58,8 @@ export interface PiSession { retry?(): Promise; getQueue(): Promise; clearQueue(): Promise; + getProjectTrust?(): Promise; + setProjectTrusted?(trusted: boolean): Promise; sendUserMessage?( type: "prompt" | "steer" | "follow_up", message: string, @@ -67,6 +73,13 @@ export interface PiSession { onError: (error: unknown) => void, onCloudStatus?: (status: TaskRunStatus) => void, ): () => void; + onExtensionEvent?( + onEvent: (event: PiExtensionEvent) => void, + onError: (error: unknown) => void, + onComplete?: () => void, + ): () => void; + respondToExtensionUI?(response: PiExtensionUIResponse): Promise; + acknowledgeExtensionEditorText?(id: string): Promise; } export interface PiSessionFactory { @@ -88,6 +101,7 @@ type PiOperation = | "bash" | "cancel" | "queue" + | "trust" | "retry" | "restart"; @@ -115,12 +129,43 @@ function normalizeSessionError(error: unknown): { }; } +interface PiExtensionEditorAckState { + attempt: number; + timeout?: ReturnType; +} + +const EXTENSION_RECONNECT_STABLE_MS = 5_000; + @injectable() export class PiSessionController { readonly store: PiSessionStore = createPiSessionStore(); private readonly sessions = new Map>(); private readonly subscriptions = new Map void>(); + private readonly extensionSubscriptions = new Map void>(); + private readonly extensionResponses = new Map< + string, + Map> + >(); + private readonly projectTrustTransitions = new Map< + string, + { trusted: boolean; promise: Promise } + >(); + private readonly extensionVersions = new Map(); + private readonly extensionOneShotIds = new Map>(); + private readonly extensionEditorAcks = new Map< + string, + Map + >(); + private readonly extensionReconnectTimeouts = new Map< + string, + ReturnType + >(); + private readonly extensionReconnectAttempts = new Map(); + private readonly extensionReconnectStabilityTimeouts = new Map< + string, + ReturnType + >(); private readonly liveEvents = new Map(); private readonly connections = new Map>(); private readonly readiness = new Map>(); @@ -142,6 +187,7 @@ export class PiSessionController { ensureConnected(taskId: string, taskRunId?: string): Promise { this.activeTaskIds.add(taskId); this.bindTaskRun(taskId, taskRunId); + this.ensureExtensionSubscription(taskId); this.ensureSubscription(taskId); const existing = this.readiness.get(taskId); @@ -181,6 +227,7 @@ export class PiSessionController { connect(taskId: string, taskRunId?: string): Promise { this.activeTaskIds.add(taskId); this.bindTaskRun(taskId, taskRunId); + this.ensureExtensionSubscription(taskId); this.ensureSubscription(taskId); const existing = this.connections.get(taskId); @@ -201,12 +248,22 @@ export class PiSessionController { disconnect(taskId: string): void { this.cancelAuthRestoration.get(taskId)?.(); - this.resetTransport(taskId); - this.taskRunIds.delete(taskId); + this.advanceSessionVersion(taskId); + this.disposeConversationSubscription(taskId); + this.connections.delete(taskId); + this.readiness.delete(taskId); this.liveEvents.delete(taskId); this.queueRevisions.delete(taskId); this.queuesToRestore.delete(taskId); this.activeTaskIds.delete(taskId); + this.cancelExtensionReconnect(taskId); + this.cancelExtensionReconnectStability(taskId); + + if (!this.extensionSubscriptions.has(taskId)) { + this.sessions.delete(taskId); + this.taskRunIds.delete(taskId); + this.clearExtensionState(taskId); + } } async retry(taskId: string): Promise { @@ -241,6 +298,71 @@ export class PiSessionController { } } + setProjectTrusted(taskId: string, trusted: boolean): Promise { + const existing = this.projectTrustTransitions.get(taskId); + if (existing) { + return existing.trusted === trusted + ? existing.promise + : Promise.reject( + new Error("A repository trust change is already in progress"), + ); + } + + const promise = this.setProjectTrustedInternal(taskId, trusted).finally( + () => { + if (this.projectTrustTransitions.get(taskId)?.promise === promise) { + this.projectTrustTransitions.delete(taskId); + } + }, + ); + this.projectTrustTransitions.set(taskId, { trusted, promise }); + return promise; + } + + private async setProjectTrustedInternal( + taskId: string, + trusted: boolean, + ): Promise { + const current = this.getSession(taskId); + if ( + current.connectionState !== "connected" || + current.status?.isStreaming || + current.isBashRunning + ) { + throw this.recordOperationFailure( + taskId, + "trust", + new Error( + "Wait for Pi to connect and finish before changing repository trust", + ), + ); + } + + const taskRunId = this.taskRunIds.get(taskId); + this.captureQueueForRestore(taskId); + this.updateSession(taskId, { + connectionState: "connecting", + error: undefined, + }); + try { + const session = await this.getPiSession(taskId); + if (!session.setProjectTrusted) { + throw new Error("Pi session does not support repository trust"); + } + await session.setProjectTrusted(trusted); + this.resetTransport(taskId); + await this.ensureConnected(taskId, taskRunId); + } catch (error) { + this.resetTransport(taskId); + try { + await this.ensureConnected(taskId, taskRunId); + } catch { + // Preserve the original trust-transition failure. + } + throw this.recordOperationFailure(taskId, "trust", error); + } + } + async restart(taskId: string): Promise { if (this.getSession(taskId).connectionState === "connecting") { return; @@ -497,6 +619,89 @@ export class PiSessionController { } } + respondToExtensionUI( + taskId: string, + response: PiExtensionUIResponse, + ): Promise { + const request = this.getSession(taskId).extensionDialogs.find( + (dialog) => dialog.id === response.id, + ); + if (!request) { + return Promise.resolve(); + } + + const taskResponses = this.extensionResponses.get(taskId) ?? new Map(); + const existing = taskResponses.get(response.id); + if (existing) { + return existing; + } + + const extensionVersion = this.getExtensionVersion(taskId); + const delivery = this.sendExtensionResponse(taskId, response) + .then(() => { + if (this.getExtensionVersion(taskId) === extensionVersion) { + this.removeExtensionDialog(taskId, response.id); + } + }) + .catch((error) => { + if (this.getExtensionVersion(taskId) === extensionVersion) { + this.addExtensionNotification( + taskId, + `Failed to respond to Pi extension: ${String(error)}`, + "error", + ); + } + throw error; + }) + .finally(() => { + if (taskResponses.get(response.id) === delivery) { + taskResponses.delete(response.id); + if ( + taskResponses.size === 0 && + this.extensionResponses.get(taskId) === taskResponses + ) { + this.extensionResponses.delete(taskId); + } + } + }); + taskResponses.set(response.id, delivery); + this.extensionResponses.set(taskId, taskResponses); + return delivery; + } + + cancelExtensionUI(taskId: string, requestId: string): Promise { + return this.respondToExtensionUI(taskId, { + type: "extension_ui_response", + id: requestId, + cancelled: true, + }); + } + + acknowledgeExtensionNotification(taskId: string, id: string): void { + const session = this.getSession(taskId); + this.updateSession(taskId, { + extensionNotifications: session.extensionNotifications.filter( + (notification) => notification.id !== id, + ), + }); + } + + acknowledgeExtensionEditorText(taskId: string, id: string): void { + if (this.getSession(taskId).extensionEditorText?.id !== id) { + return; + } + this.updateSession(taskId, { extensionEditorText: undefined }); + + const taskAcks = this.extensionEditorAcks.get(taskId) ?? new Map(); + if (taskAcks.has(id)) { + return; + } + const acknowledgement: PiExtensionEditorAckState = { attempt: 0 }; + taskAcks.set(id, acknowledgement); + this.extensionEditorAcks.set(taskId, taskAcks); + this.deliverExtensionEditorAck(taskId, id, acknowledgement); + } + private async ensureConnectedInternal(taskId: string): Promise { const session = await this.getPiSession(taskId); const health = await session.health(); @@ -509,14 +714,15 @@ export class PiSessionController { throw new Error(result.error); } - this.subscriptions.get(taskId)?.(); - this.subscriptions.delete(taskId); + this.disposeConversationSubscription(taskId); this.sessions.delete(taskId); this.connections.delete(taskId); this.ensureSubscription(taskId); } - await this.connect(taskId); + if (this.activeTaskIds.has(taskId)) { + await this.connect(taskId); + } } private ensureSubscription(taskId: string): void { @@ -546,6 +752,175 @@ export class PiSessionController { }); } + private ensureExtensionSubscription(taskId: string): void { + if (this.extensionSubscriptions.has(taskId)) { + return; + } + this.cancelExtensionReconnect(taskId, false); + + let disposed = false; + let unsubscribe: (() => void) | undefined; + const dispose = () => { + disposed = true; + unsubscribe?.(); + }; + this.extensionSubscriptions.set(taskId, dispose); + + void this.getPiSession(taskId) + .then((session) => { + if (disposed || !session.onExtensionEvent) { + this.extensionReconnectAttempts.delete(taskId); + if (this.extensionSubscriptions.get(taskId) === dispose) { + this.extensionSubscriptions.delete(taskId); + } + if (!disposed && !this.activeTaskIds.has(taskId)) { + this.sessions.delete(taskId); + this.taskRunIds.delete(taskId); + this.clearExtensionState(taskId); + } + return; + } + const nextUnsubscribe = session.onExtensionEvent( + (event) => this.handleExtensionEvent(taskId, event, dispose), + (error) => this.endExtensionSubscription(taskId, dispose, error), + () => this.endExtensionSubscription(taskId, dispose), + ); + if (disposed) { + nextUnsubscribe(); + } else { + unsubscribe = nextUnsubscribe; + } + }) + .catch((error) => this.endExtensionSubscription(taskId, dispose, error)); + } + + private endExtensionSubscription( + taskId: string, + dispose: () => void, + error?: unknown, + ): void { + if (this.extensionSubscriptions.get(taskId) !== dispose) { + return; + } + this.extensionSubscriptions.delete(taskId); + dispose(); + this.sessions.delete(taskId); + this.cancelExtensionReconnectStability(taskId); + if (error !== undefined && this.activeTaskIds.has(taskId)) { + if ((this.extensionReconnectAttempts.get(taskId) ?? 0) === 0) { + this.addExtensionNotification( + taskId, + `Pi extension UI disconnected: ${String(error)}`, + "error", + `extension-ui-disconnected:${taskId}`, + ); + } + this.scheduleExtensionReconnect(taskId); + } else { + this.cancelExtensionReconnect(taskId); + this.clearExtensionState(taskId); + } + } + + private scheduleExtensionReconnect(taskId: string): void { + if ( + !this.activeTaskIds.has(taskId) || + this.extensionReconnectTimeouts.has(taskId) + ) { + return; + } + const attempt = this.extensionReconnectAttempts.get(taskId) ?? 0; + const delay = Math.min(100 * 2 ** attempt, 1_000); + this.extensionReconnectAttempts.set(taskId, Math.min(attempt + 1, 4)); + this.extensionReconnectTimeouts.set( + taskId, + setTimeout(() => { + this.extensionReconnectTimeouts.delete(taskId); + if (this.activeTaskIds.has(taskId)) { + this.ensureExtensionSubscription(taskId); + } + }, delay), + ); + } + + private cancelExtensionReconnect(taskId: string, resetAttempts = true): void { + const timeout = this.extensionReconnectTimeouts.get(taskId); + if (timeout) { + clearTimeout(timeout); + this.extensionReconnectTimeouts.delete(taskId); + } + if (resetAttempts) { + this.extensionReconnectAttempts.delete(taskId); + } + } + + private scheduleExtensionReconnectStability( + taskId: string, + dispose: () => void, + ): void { + if (this.extensionReconnectStabilityTimeouts.has(taskId)) { + return; + } + this.extensionReconnectStabilityTimeouts.set( + taskId, + setTimeout(() => { + this.extensionReconnectStabilityTimeouts.delete(taskId); + if (this.extensionSubscriptions.get(taskId) === dispose) { + this.extensionReconnectAttempts.delete(taskId); + } + }, EXTENSION_RECONNECT_STABLE_MS), + ); + } + + private cancelExtensionReconnectStability(taskId: string): void { + const timeout = this.extensionReconnectStabilityTimeouts.get(taskId); + if (timeout) { + clearTimeout(timeout); + this.extensionReconnectStabilityTimeouts.delete(taskId); + } + } + + private deliverExtensionEditorAck( + taskId: string, + id: string, + acknowledgement: PiExtensionEditorAckState, + ): void { + void this.getPiSession(taskId) + .then((session) => session.acknowledgeExtensionEditorText?.(id)) + .then(() => { + const taskAcks = this.extensionEditorAcks.get(taskId); + if (taskAcks?.get(id) !== acknowledgement) { + return; + } + taskAcks.delete(id); + if (taskAcks.size === 0) { + this.extensionEditorAcks.delete(taskId); + } + }) + .catch(() => { + const taskAcks = this.extensionEditorAcks.get(taskId); + if (taskAcks?.get(id) !== acknowledgement) { + return; + } + const delay = Math.min(100 * 2 ** acknowledgement.attempt, 1_000); + acknowledgement.attempt = Math.min(acknowledgement.attempt + 1, 4); + acknowledgement.timeout = setTimeout(() => { + acknowledgement.timeout = undefined; + this.deliverExtensionEditorAck(taskId, id, acknowledgement); + }, delay); + }); + } + + private clearExtensionEditorAcks(taskId: string): void { + const taskAcks = this.extensionEditorAcks.get(taskId); + for (const acknowledgement of taskAcks?.values() ?? []) { + if (acknowledgement.timeout) { + clearTimeout(acknowledgement.timeout); + } + } + this.extensionEditorAcks.delete(taskId); + } + private applyPersistedConfig(taskId: string, session: PiSession): void { const config = session.persistedConfig; if (!config) { @@ -586,11 +961,12 @@ export class PiSessionController { const session = await this.getPiSession(taskId); const queueRevision = this.queueRevisions.get(taskId) ?? 0; const retainedStats = this.getSession(taskId).stats; - const [events, status, queue, stats] = await Promise.all([ + const [events, status, queue, stats, projectTrust] = await Promise.all([ session.getConversation(), session.client.getState(), session.getQueue(), session.client.getSessionStats().catch(() => retainedStats), + session.getProjectTrust?.(), ]); if (this.getSessionVersion(taskId) !== connectedSessionVersion) { return; @@ -652,6 +1028,13 @@ export class PiSessionController { : undefined, authRestoring: currentSession.authRestoring, isBashRunning: false, + projectTrust, + extensionDialogs: currentSession.extensionDialogs, + extensionNotifications: currentSession.extensionNotifications, + extensionStatuses: currentSession.extensionStatuses, + extensionWidgets: currentSession.extensionWidgets, + extensionTitle: currentSession.extensionTitle, + extensionEditorText: currentSession.extensionEditorText, }); await this.restoreQueueIfNeeded(taskId, session, resolvedStatus); @@ -700,6 +1083,224 @@ export class PiSessionController { } } + private handleExtensionEvent( + taskId: string, + event: PiExtensionEvent, + dispose?: () => void, + ): void { + if (event.type === "extension_state_snapshot") { + this.applyExtensionSnapshot(taskId, event); + if (dispose) { + this.scheduleExtensionReconnectStability(taskId, dispose); + } + return; + } + + this.extensionReconnectAttempts.delete(taskId); + this.cancelExtensionReconnectStability(taskId); + + if (event.type === "extension_dialog_expired") { + this.removeExtensionDialog(taskId, event.id); + return; + } + + if (event.type === "extension_session_reset") { + const dispose = this.extensionSubscriptions.get(taskId); + if (dispose) { + this.endExtensionSubscription(taskId, dispose); + } else { + this.clearExtensionState(taskId); + } + return; + } + + if (event.type === "extension_error") { + const extensionName = event.extensionPath.split(/[\\/]/).pop(); + this.addExtensionNotification( + taskId, + `${extensionName ?? event.extensionPath} failed during ${event.event}: ${event.error}`, + "error", + ); + return; + } + + switch (event.method) { + case "select": + case "confirm": + case "input": + case "editor": + this.enqueueExtensionDialog(taskId, event); + return; + case "notify": + this.addExtensionNotification( + taskId, + event.message, + event.notifyType ?? "info", + event.id, + ); + return; + case "setStatus": { + const statuses = { ...this.getSession(taskId).extensionStatuses }; + if (event.statusText === undefined) { + delete statuses[event.statusKey]; + } else { + statuses[event.statusKey] = event.statusText; + } + this.updateSession(taskId, { extensionStatuses: statuses }); + return; + } + case "setWidget": { + const widgets = { ...this.getSession(taskId).extensionWidgets }; + if (event.widgetLines === undefined) { + delete widgets[event.widgetKey]; + } else { + widgets[event.widgetKey] = { + lines: event.widgetLines, + placement: event.widgetPlacement ?? "aboveEditor", + }; + } + this.updateSession(taskId, { extensionWidgets: widgets }); + return; + } + case "setTitle": + this.updateSession(taskId, { extensionTitle: event.title }); + return; + case "set_editor_text": + if (!this.markExtensionOneShot(taskId, event.id)) { + return; + } + this.updateSession(taskId, { + extensionEditorText: { id: event.id, text: event.text }, + }); + return; + } + } + + private applyExtensionSnapshot( + taskId: string, + snapshot: Extract, + ): void { + const current = this.getSession(taskId); + const currentDialogIds = new Set( + current.extensionDialogs.map((dialog) => dialog.id), + ); + const dialogs = snapshot.dialogs.filter( + (dialog) => + currentDialogIds.has(dialog.id) || + this.markExtensionOneShot(taskId, dialog.id), + ); + const statuses: Record = {}; + for (const status of snapshot.statuses) { + if (status.statusText !== undefined) { + statuses[status.statusKey] = status.statusText; + } + } + const widgets: PiControllerSessionState["extensionWidgets"] = {}; + for (const widget of snapshot.widgets) { + if (widget.widgetLines !== undefined) { + widgets[widget.widgetKey] = { + lines: widget.widgetLines, + placement: widget.widgetPlacement ?? "aboveEditor", + }; + } + } + + let editorText: PiControllerSessionState["extensionEditorText"]; + if (snapshot.editorText) { + if (current.extensionEditorText?.id === snapshot.editorText.id) { + editorText = current.extensionEditorText; + } else if (this.markExtensionOneShot(taskId, snapshot.editorText.id)) { + editorText = { + id: snapshot.editorText.id, + text: snapshot.editorText.text, + }; + } + } + + this.updateSession(taskId, { + extensionDialogs: dialogs, + extensionStatuses: statuses, + extensionWidgets: widgets, + extensionTitle: snapshot.title?.title, + extensionEditorText: editorText, + }); + } + + private enqueueExtensionDialog( + taskId: string, + request: PiExtensionDialogRequest, + ): void { + const session = this.getSession(taskId); + if (!this.markExtensionOneShot(taskId, request.id)) { + return; + } + this.updateSession(taskId, { + extensionDialogs: [...session.extensionDialogs, request], + }); + } + + private markExtensionOneShot(taskId: string, id: string): boolean { + const ids = this.extensionOneShotIds.get(taskId) ?? new Set(); + if (ids.has(id)) { + return false; + } + ids.add(id); + this.extensionOneShotIds.set(taskId, ids); + return true; + } + + private removeExtensionDialog(taskId: string, requestId: string): void { + const session = this.getSession(taskId); + this.updateSession(taskId, { + extensionDialogs: session.extensionDialogs.filter( + (dialog) => dialog.id !== requestId, + ), + }); + } + + private async sendExtensionResponse( + taskId: string, + response: PiExtensionUIResponse, + ): Promise { + const session = await this.getPiSession(taskId); + if (!session.respondToExtensionUI) { + throw new Error("Pi session does not support extension UI responses"); + } + await session.respondToExtensionUI(response); + } + + private clearExtensionState(taskId: string): void { + this.extensionVersions.set(taskId, this.getExtensionVersion(taskId) + 1); + this.extensionResponses.delete(taskId); + this.extensionOneShotIds.delete(taskId); + this.clearExtensionEditorAcks(taskId); + this.updateSession(taskId, { + extensionDialogs: [], + extensionNotifications: [], + extensionStatuses: {}, + extensionWidgets: {}, + extensionTitle: undefined, + extensionEditorText: undefined, + }); + } + + private addExtensionNotification( + taskId: string, + message: string, + notifyType: "info" | "warning" | "error", + id: string = globalThis.crypto.randomUUID(), + ): void { + const session = this.getSession(taskId); + this.updateSession(taskId, { + extensionNotifications: [ + ...session.extensionNotifications.filter( + (notification) => notification.id !== id, + ), + { id, message, notifyType }, + ], + }); + } + private handleEvent(taskId: string, event: AgentConversationEvent): void { if (event.type === "queue_update") { const queue = { @@ -944,6 +1545,7 @@ export class PiSessionController { bash: "Failed to run Pi bash command", cancel: "Failed to stop Pi", queue: "Failed to update queued message", + trust: "Failed to change repository trust", retry: "Failed to reconnect to Pi", restart: "Failed to restart Pi", }; @@ -1169,7 +1771,12 @@ export class PiSessionController { taskId, taskRunId, ); - this.disconnect(taskId); + this.resetTransport(taskId); + this.taskRunIds.delete(taskId); + this.liveEvents.delete(taskId); + this.queueRevisions.delete(taskId); + this.queuesToRestore.delete(taskId); + this.activeTaskIds.delete(taskId); await this.ensureConnected(taskId, resumedRun.id); return this.getPiSession(taskId); } @@ -1180,7 +1787,7 @@ export class PiSessionController { return; } - if (currentTaskRunId) { + if (currentTaskRunId || this.sessions.has(taskId)) { this.resetTransport(taskId); this.liveEvents.delete(taskId); } @@ -1189,11 +1796,20 @@ export class PiSessionController { private resetTransport(taskId: string): void { this.advanceSessionVersion(taskId); - this.subscriptions.get(taskId)?.(); - this.subscriptions.delete(taskId); + this.cancelExtensionReconnect(taskId); + this.cancelExtensionReconnectStability(taskId); + this.disposeConversationSubscription(taskId); + this.extensionSubscriptions.get(taskId)?.(); + this.extensionSubscriptions.delete(taskId); this.sessions.delete(taskId); this.connections.delete(taskId); this.readiness.delete(taskId); + this.clearExtensionState(taskId); + } + + private disposeConversationSubscription(taskId: string): void { + this.subscriptions.get(taskId)?.(); + this.subscriptions.delete(taskId); } private getPiSession(taskId: string): Promise { @@ -1223,6 +1839,10 @@ export class PiSessionController { return this.sessionVersions.get(taskId) ?? 0; } + private getExtensionVersion(taskId: string): number { + return this.extensionVersions.get(taskId) ?? 0; + } + private advanceSessionVersion(taskId: string): void { this.sessionVersions.set(taskId, this.getSessionVersion(taskId) + 1); } diff --git a/products/desktop/packages/core/src/pi-runtime/piSessionStore.ts b/products/desktop/packages/core/src/pi-runtime/piSessionStore.ts index a1dd61b0e192..9a4d88d32418 100644 --- a/products/desktop/packages/core/src/pi-runtime/piSessionStore.ts +++ b/products/desktop/packages/core/src/pi-runtime/piSessionStore.ts @@ -1,5 +1,6 @@ import type { PiCommand, + PiExtensionUIRequest, PiNativeModelInfo, PiQueueSnapshot, PiSessionStats, @@ -15,6 +16,32 @@ import type { } from "@posthog/shared"; import { createStore, type StoreApi } from "zustand/vanilla"; +export type PiExtensionDialogRequest = Extract< + PiExtensionUIRequest, + { method: "select" | "confirm" | "input" | "editor" } +>; + +export interface PiExtensionNotification { + id: string; + message: string; + notifyType: "info" | "warning" | "error"; +} + +export interface PiExtensionWidget { + lines: string[]; + placement: "aboveEditor" | "belowEditor"; +} + +export interface PiExtensionEditorText { + id: string; + text: string; +} + +export interface PiProjectTrustState { + trusted: boolean; + hasProjectResources: boolean; +} + export interface PiSessionError { id: string; scope: "connection" | "operation"; @@ -41,6 +68,13 @@ export interface PiControllerSessionState { error?: PiSessionError; authRestoring: boolean; isBashRunning: boolean; + projectTrust?: PiProjectTrustState; + extensionDialogs: PiExtensionDialogRequest[]; + extensionNotifications: PiExtensionNotification[]; + extensionStatuses: Record; + extensionWidgets: Record; + extensionTitle?: string; + extensionEditorText?: PiExtensionEditorText; } export interface PiSessionState { @@ -65,5 +99,9 @@ export function createEmptyPiControllerSession(): PiControllerSessionState { queue: { steering: [], followUp: [] }, authRestoring: false, isBashRunning: false, + extensionDialogs: [], + extensionNotifications: [], + extensionStatuses: {}, + extensionWidgets: {}, }; } diff --git a/products/desktop/packages/core/src/task-detail/taskCreationSaga.test.ts b/products/desktop/packages/core/src/task-detail/taskCreationSaga.test.ts index 1d3849110da0..872d8534abe9 100644 --- a/products/desktop/packages/core/src/task-detail/taskCreationSaga.test.ts +++ b/products/desktop/packages/core/src/task-detail/taskCreationSaga.test.ts @@ -374,6 +374,7 @@ describe("TaskCreationSaga", () => { expect(piRunner.create).toHaveBeenCalledWith({ taskId: "task-123", cwd: "/tmp/scratch/task-123", + projectTrustPath: "/tmp/scratch/task-123", prompt: "Draft a launch email", model: "claude-sonnet", thinkingLevel: "medium", diff --git a/products/desktop/packages/core/src/task-detail/taskCreationSaga.ts b/products/desktop/packages/core/src/task-detail/taskCreationSaga.ts index da3f55a9caec..bffba3e95a8e 100644 --- a/products/desktop/packages/core/src/task-detail/taskCreationSaga.ts +++ b/products/desktop/packages/core/src/task-detail/taskCreationSaga.ts @@ -531,6 +531,8 @@ export class TaskCreationSaga extends Saga< await this.deps.piRunner.create({ taskId: task.id, cwd: agentCwd ?? "", + projectTrustPath: + workspace?.folderPath ?? repoPath ?? scratchCwd ?? undefined, prompt: input.content ?? "", model: input.model, thinkingLevel, diff --git a/products/desktop/packages/core/src/task-detail/taskService.test.ts b/products/desktop/packages/core/src/task-detail/taskService.test.ts index 9b147bd0598e..5b614c0c0f04 100644 --- a/products/desktop/packages/core/src/task-detail/taskService.test.ts +++ b/products/desktop/packages/core/src/task-detail/taskService.test.ts @@ -55,6 +55,44 @@ function makeService(): TaskService { } describe("TaskService.openTask", () => { + it("passes the main repository path when resuming Pi in a worktree", async () => { + const task = { + id: "task-1", + runtime: "pi", + latest_run: { id: "run-1", environment: "local", status: "running" }, + }; + const workspace = { + folderPath: "/repo", + worktreePath: "/worktrees/task-1", + }; + const piRunner = { + create: vi.fn(), + resume: vi.fn(async () => {}), + stop: vi.fn(), + } as unknown as PiRunner; + const service = new TaskService( + { + getAuthenticatedClient: vi.fn(async () => ({ + getTask: vi.fn(async () => task), + })), + getWorkspace: vi.fn(async () => workspace), + } as unknown as ITaskCreationHost, + {} as SessionService, + {} as TaskCreationEffects, + piRunner, + rootLogger, + ); + + await expect(service.openTask("task-1")).resolves.toMatchObject({ + success: true, + }); + expect(piRunner.resume).toHaveBeenCalledWith({ + taskId: "task-1", + cwd: "/worktrees/task-1", + projectTrustPath: "/repo", + }); + }); + it("opens a completed cloud Pi run without resuming it", async () => { const completedRun = { id: "run-1", diff --git a/products/desktop/packages/core/src/task-detail/taskService.ts b/products/desktop/packages/core/src/task-detail/taskService.ts index 52c76216a6ef..fbcf5e8aeb7e 100644 --- a/products/desktop/packages/core/src/task-detail/taskService.ts +++ b/products/desktop/packages/core/src/task-detail/taskService.ts @@ -257,6 +257,7 @@ export class TaskService { await this.piRunner.resume({ taskId, cwd: existingWorkspace.worktreePath ?? existingWorkspace.folderPath, + projectTrustPath: existingWorkspace.folderPath, }); } @@ -277,7 +278,7 @@ export class TaskService { if (runtime === "pi") { try { const cwd = await this.host.ensureScratchDir(taskId); - await this.piRunner.resume({ taskId, cwd }); + await this.piRunner.resume({ taskId, cwd, projectTrustPath: cwd }); return { success: true, data: { task, workspace: null }, diff --git a/products/desktop/packages/harness/package.json b/products/desktop/packages/harness/package.json index a158c46ba3e6..9ddb5277d084 100644 --- a/products/desktop/packages/harness/package.json +++ b/products/desktop/packages/harness/package.json @@ -20,6 +20,10 @@ "types": "./dist/runtime.d.ts", "import": "./dist/runtime.js" }, + "./project-trust": { + "types": "./dist/project-trust.d.ts", + "import": "./dist/project-trust.js" + }, "./extensions": { "types": "./dist/extensions/registry.d.ts", "import": "./dist/extensions/registry.js" diff --git a/products/desktop/packages/harness/src/project-trust.test.ts b/products/desktop/packages/harness/src/project-trust.test.ts new file mode 100644 index 000000000000..372305ef5539 --- /dev/null +++ b/products/desktop/packages/harness/src/project-trust.test.ts @@ -0,0 +1,78 @@ +import { mkdir, mkdtemp, rm } from "node:fs/promises"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { afterEach, describe, expect, it } from "vitest"; +import { + createPiProjectTrustResolver, + readPiProjectTrust, + writePiProjectTrust, +} from "./project-trust"; + +const temporaryDirectories: string[] = []; + +async function temporaryDirectory(): Promise { + const directory = await mkdtemp(join(tmpdir(), "posthog-pi-trust-")); + temporaryDirectories.push(directory); + return directory; +} + +afterEach(async () => { + await Promise.all( + temporaryDirectories + .splice(0) + .map((directory) => rm(directory, { force: true, recursive: true })), + ); +}); + +describe("Pi project trust", () => { + it("persists trust for a repository and applies it to managed worktrees", async () => { + const agentDir = await temporaryDirectory(); + const workspaceRoot = await temporaryDirectory(); + const repository = join(workspaceRoot, "repository"); + const worktree = join(workspaceRoot, "worktrees", "task-1"); + await mkdir(repository); + await mkdir(join(worktree, ".pi", "extensions"), { recursive: true }); + + expect(readPiProjectTrust(repository, worktree, agentDir)).toEqual({ + trusted: false, + hasProjectResources: true, + }); + + writePiProjectTrust(repository, true, agentDir); + + expect(readPiProjectTrust(repository, worktree, agentDir)).toEqual({ + trusted: true, + hasProjectResources: true, + }); + }); + + it("does not carry initial repository trust into an unrelated session cwd", async () => { + const agentDir = await temporaryDirectory(); + const initialRepository = await temporaryDirectory(); + const unrelatedRepository = await temporaryDirectory(); + const resolveTrust = createPiProjectTrustResolver( + initialRepository, + true, + agentDir, + ); + + expect(resolveTrust(initialRepository)).toBe(true); + expect(resolveTrust(unrelatedRepository)).toBe(false); + + writePiProjectTrust(unrelatedRepository, true, agentDir); + expect(resolveTrust(unrelatedRepository)).toBe(true); + }); + + it("honors Pi trust decisions inherited from an ancestor", async () => { + const agentDir = await temporaryDirectory(); + const parent = await temporaryDirectory(); + const repository = join(parent, "repository"); + await mkdir(repository); + + writePiProjectTrust(parent, true, agentDir); + + expect(readPiProjectTrust(repository, repository, agentDir).trusted).toBe( + true, + ); + }); +}); diff --git a/products/desktop/packages/harness/src/project-trust.ts b/products/desktop/packages/harness/src/project-trust.ts new file mode 100644 index 000000000000..954089c06113 --- /dev/null +++ b/products/desktop/packages/harness/src/project-trust.ts @@ -0,0 +1,42 @@ +import { resolve } from "node:path"; +import { + getAgentDir, + hasTrustRequiringProjectResources, + ProjectTrustStore, +} from "@earendil-works/pi-coding-agent"; + +export interface PiProjectTrust { + trusted: boolean; + hasProjectResources: boolean; +} + +export function readPiProjectTrust( + projectTrustPath: string, + runtimeCwd: string = projectTrustPath, + agentDir: string = getAgentDir(), +): PiProjectTrust { + return { + trusted: new ProjectTrustStore(agentDir).get(projectTrustPath) === true, + hasProjectResources: hasTrustRequiringProjectResources(runtimeCwd), + }; +} + +export function writePiProjectTrust( + projectTrustPath: string, + trusted: boolean, + agentDir: string = getAgentDir(), +): void { + new ProjectTrustStore(agentDir).set(projectTrustPath, trusted); +} + +export function createPiProjectTrustResolver( + initialCwd: string, + initialTrusted: boolean, + agentDir: string = getAgentDir(), +): (runtimeCwd: string) => boolean { + const resolvedInitialCwd = resolve(initialCwd); + return (runtimeCwd) => + resolve(runtimeCwd) === resolvedInitialCwd + ? initialTrusted + : readPiProjectTrust(runtimeCwd, runtimeCwd, agentDir).trusted; +} diff --git a/products/desktop/packages/harness/src/runtime.test.ts b/products/desktop/packages/harness/src/runtime.test.ts index d127d073ecc0..5a4333264804 100644 --- a/products/desktop/packages/harness/src/runtime.test.ts +++ b/products/desktop/packages/harness/src/runtime.test.ts @@ -1,5 +1,5 @@ import { existsSync } from "node:fs"; -import { mkdtemp, readFile, rm, writeFile } from "node:fs/promises"; +import { mkdir, mkdtemp, readFile, rm, writeFile } from "node:fs/promises"; import { tmpdir } from "node:os"; import { join } from "node:path"; import { InMemoryCredentialStore } from "@earendil-works/pi-ai"; @@ -69,6 +69,47 @@ describe("createHarnessRuntime", () => { }, ); + it.each([ + { label: "false", projectTrusted: false, loaded: false }, + { label: "true", projectTrusted: true, loaded: true }, + { + label: "resolved true for the runtime cwd", + projectTrusted: (_cwd: string) => true, + loaded: true, + }, + ])( + "loads project-local extensions when project trust is $label", + async ({ projectTrusted, loaded }) => { + vi.stubEnv("PI_OFFLINE", "1"); + const pi = await import("@earendil-works/pi-coding-agent"); + const cwd = await temporaryDirectory(); + const agentDir = await temporaryDirectory(); + const extensionPath = join(cwd, ".pi", "extensions", "project.ts"); + await mkdir(join(cwd, ".pi", "extensions"), { recursive: true }); + await writeFile( + extensionPath, + "export default function projectExtension() {}\n", + ); + + const runtime = await createHarnessRuntime({ + agentDir, + credentialStore: new InMemoryCredentialStore(), + cwd, + projectTrusted, + sessionManager: pi.SessionManager.inMemory(cwd), + }); + + try { + const extensionPaths = runtime.services.resourceLoader + .getExtensions() + .extensions.map((extension) => extension.path); + expect(extensionPaths.includes(extensionPath)).toBe(loaded); + } finally { + await runtime.dispose(); + } + }, + ); + it("restores the session model before calculating context usage", async () => { vi.stubEnv("PI_OFFLINE", "1"); const pi = await import("@earendil-works/pi-coding-agent"); diff --git a/products/desktop/packages/harness/src/runtime.ts b/products/desktop/packages/harness/src/runtime.ts index 906b53b616c2..5ff56e558c81 100644 --- a/products/desktop/packages/harness/src/runtime.ts +++ b/products/desktop/packages/harness/src/runtime.ts @@ -43,6 +43,7 @@ async function createCredentialStore( export type HarnessRuntimeOptions = HarnessExtensionOptions & { credentialStore?: CredentialStore; posthogOAuthCredentials?: PosthogOAuthCredentials; + projectTrusted?: boolean | ((cwd: string) => boolean); } & Partial< Pick< PiRuntimeTarget, @@ -67,8 +68,12 @@ export type HarnessRuntimeOptions = HarnessExtensionOptions & { export async function createHarnessRuntime( options: HarnessRuntimeOptions = {}, ): Promise { - const { credentialStore, posthogOAuthCredentials, ...runtimeOptions } = - options; + const { + credentialStore, + posthogOAuthCredentials, + projectTrusted, + ...runtimeOptions + } = options; // Pi reads its application branding when the SDK is first evaluated. Keep // every runtime import below dynamic so this always happens first. installHogBrandEnv(); @@ -112,7 +117,10 @@ export async function createHarnessRuntime( settingsManager: options.settingsManager ?? pi.SettingsManager.create(runtimeCwd, runtimeAgentDir, { - projectTrusted: false, + projectTrusted: + typeof projectTrusted === "function" + ? projectTrusted(runtimeCwd) + : (projectTrusted ?? false), }), resourceLoaderOptions: { ...runtimeOptions.resourceLoaderOptions, diff --git a/products/desktop/packages/harness/tsup.config.ts b/products/desktop/packages/harness/tsup.config.ts index 0ebc718d44f9..2e9509969100 100644 --- a/products/desktop/packages/harness/tsup.config.ts +++ b/products/desktop/packages/harness/tsup.config.ts @@ -6,6 +6,7 @@ export default defineConfig({ "src/index.ts", "src/cli.ts", "src/runtime.ts", + "src/project-trust.ts", "src/extensions/registry.ts", "src/extensions/hog-branding/extension.ts", "src/extensions/hog-branding/index.ts", diff --git a/products/desktop/packages/host-router/src/pi-session-factory.ts b/products/desktop/packages/host-router/src/pi-session-factory.ts index 79f133ef66aa..853322f1cbda 100644 --- a/products/desktop/packages/host-router/src/pi-session-factory.ts +++ b/products/desktop/packages/host-router/src/pi-session-factory.ts @@ -44,6 +44,52 @@ class TrpcPiSession implements PiSession { return this.hostClient.piSession.clearQueue.mutate({ taskId: this.taskId }); } + getProjectTrust() { + return this.hostClient.piSession.getProjectTrust.query({ + taskId: this.taskId, + }); + } + + setProjectTrusted(trusted: boolean) { + return this.hostClient.piSession.setProjectTrusted.mutate({ + taskId: this.taskId, + trusted, + }); + } + + respondToExtensionUI( + response: Parameters>[0], + ) { + return this.hostClient.piSession.respondToExtensionUI.mutate({ + taskId: this.taskId, + response, + }); + } + + acknowledgeExtensionEditorText(id: string) { + return this.hostClient.piSession.acknowledgeExtensionEditorText.mutate({ + taskId: this.taskId, + id, + }); + } + + onExtensionEvent( + onEvent: Parameters>[0], + onError: Parameters>[1], + onComplete?: Parameters>[2], + ): () => void { + const subscription = this.hostClient.piSession.onExtensionEvent.subscribe( + { taskId: this.taskId }, + { + onData: (event) => onEvent(event as Parameters[0]), + onError, + onComplete, + }, + ); + + return () => subscription.unsubscribe(); + } + onConversationEvent( onEvent: Parameters[0], onError: Parameters[1], diff --git a/products/desktop/packages/host-router/src/routers/pi-session.router.ts b/products/desktop/packages/host-router/src/routers/pi-session.router.ts index 23a6705c68a4..7408c9f4c5c6 100644 --- a/products/desktop/packages/host-router/src/routers/pi-session.router.ts +++ b/products/desktop/packages/host-router/src/routers/pi-session.router.ts @@ -2,6 +2,10 @@ import { publicProcedure, router } from "@posthog/host-trpc/trpc"; import { PI_SESSION_SERVICE } from "@posthog/workspace-server/services/pi-session/identifiers"; import type { PiSessionService } from "@posthog/workspace-server/services/pi-session/pi-session"; import { + piExtensionEditorTextAckInput, + piExtensionEventSchema, + piExtensionUIResponseInput, + piProjectTrustOutput, piQueueSnapshotOutput, piRpcResponseSchema, piSessionConfigInput, @@ -11,6 +15,7 @@ import { piSessionStartOutput, piSessionTaskInput, resumePiSessionInput, + setPiProjectTrustInput, startPiSessionInput, } from "@posthog/workspace-server/services/pi-session/schemas"; @@ -64,6 +69,37 @@ export const piSessionRouter = router({ getService(ctx.container).clearQueue(input.taskId), ), + getProjectTrust: publicProcedure + .input(piSessionTaskInput) + .output(piProjectTrustOutput) + .query(({ ctx, input }) => + getService(ctx.container).getProjectTrust(input.taskId), + ), + + setProjectTrusted: publicProcedure + .input(setPiProjectTrustInput) + .mutation(({ ctx, input }) => + getService(ctx.container).setProjectTrusted(input.taskId, input.trusted), + ), + + respondToExtensionUI: publicProcedure + .input(piExtensionUIResponseInput) + .mutation(({ ctx, input }) => + getService(ctx.container).respondToExtensionUI( + input.taskId, + input.response, + ), + ), + + acknowledgeExtensionEditorText: publicProcedure + .input(piExtensionEditorTextAckInput) + .mutation(({ ctx, input }) => + getService(ctx.container).acknowledgeExtensionEditorText( + input.taskId, + input.id, + ), + ), + onEvent: publicProcedure .input(piSessionTaskInput) .subscription(async function* (opts) { @@ -75,4 +111,39 @@ export const piSessionRouter = router({ } } }), + + onExtensionEvent: publicProcedure + .input(piSessionTaskInput) + .subscription(async function* (opts) { + const service = getService(opts.ctx.container); + const iterator = service + .toIterable("extensionEvent", { signal: opts.signal }) + [Symbol.asyncIterator](); + let next = iterator.next(); + const snapshot = service.getExtensionStateSnapshot(opts.input.taskId); + + try { + yield piExtensionEventSchema.parse(snapshot); + + while (true) { + const result = await next; + if (result.done) { + return; + } + if (result.value.taskId === opts.input.taskId) { + const event = piExtensionEventSchema.parse(result.value.event); + if (event.type === "extension_session_reset") { + yield event; + return; + } + next = iterator.next(); + yield event; + } else { + next = iterator.next(); + } + } + } finally { + await iterator.return?.(); + } + }), }); diff --git a/products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.test.tsx b/products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.test.tsx new file mode 100644 index 000000000000..dcb9d32bdde6 --- /dev/null +++ b/products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.test.tsx @@ -0,0 +1,217 @@ +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; +import { + buildPiExtensionResponse, + PiExtensionDialog, +} from "./PiExtensionDialog"; +import { PiExtensionStatuses, PiExtensionWidgets } from "./PiExtensionSurfaces"; +import { piExtensionEditorTextToContent } from "./piExtensionEditorText"; + +describe("Pi extension presenters", () => { + it("builds matching response wire shapes", () => { + expect( + buildPiExtensionResponse( + { + type: "extension_ui_request", + id: "confirm-1", + method: "confirm", + title: "Continue?", + message: "Proceed?", + }, + true, + ), + ).toEqual({ + type: "extension_ui_response", + id: "confirm-1", + confirmed: true, + }); + expect( + buildPiExtensionResponse( + { + type: "extension_ui_request", + id: "confirm-2", + method: "confirm", + title: "Continue?", + message: "Proceed?", + }, + false, + ), + ).toEqual({ + type: "extension_ui_response", + id: "confirm-2", + confirmed: false, + }); + expect( + buildPiExtensionResponse( + { + type: "extension_ui_request", + id: "editor-1", + method: "editor", + title: "Edit", + }, + "updated", + ), + ).toEqual({ + type: "extension_ui_response", + id: "editor-1", + value: "updated", + }); + }); + + it("submits labelled input with Enter and allows retry after failure", async () => { + const user = userEvent.setup(); + const onRespond = vi + .fn<() => Promise>() + .mockRejectedValueOnce(new Error("wire failed")) + .mockResolvedValueOnce(); + const onCancel = vi.fn(async () => {}); + + render( + , + ); + + const input = screen.getByLabelText("Response"); + await user.type(input, "Ada{Enter}"); + await waitFor(() => expect(onRespond).toHaveBeenCalledTimes(1)); + expect(onRespond).toHaveBeenLastCalledWith({ + type: "extension_ui_response", + id: "input-1", + value: "Ada", + }); + + await user.click(screen.getByRole("button", { name: "Submit" })); + expect(onRespond).toHaveBeenCalledTimes(2); + }); + + it("submits an explicit negative confirmation", async () => { + const onRespond = vi.fn(async () => {}); + + render( + {})} + />, + ); + + await userEvent.click(screen.getByRole("button", { name: "Decline" })); + expect(onRespond).toHaveBeenCalledWith({ + type: "extension_ui_response", + id: "confirm-1", + confirmed: false, + }); + }); + + it("guards concurrent cancellation while delivery is pending", async () => { + let resolveCancel: () => void = () => {}; + const onCancel = vi.fn( + () => + new Promise((resolve) => { + resolveCancel = resolve; + }), + ); + + render( + {})} + onCancel={onCancel} + />, + ); + + const form = screen.getByRole("form", { name: "Continue? response" }); + fireEvent.click(screen.getByRole("button", { name: "Cancel" })); + fireEvent.submit(form); + fireEvent.click(screen.getByRole("button", { name: "Cancel" })); + + expect(onCancel).toHaveBeenCalledTimes(1); + expect(screen.getByRole("button", { name: "Submitting…" })).toHaveAttribute( + "aria-disabled", + "true", + ); + resolveCancel(); + await waitFor(() => + expect(screen.getByRole("button", { name: "Confirm" })).toHaveAttribute( + "aria-disabled", + "false", + ), + ); + }); + + it("keeps Enter as a newline in the multiline editor", async () => { + const user = userEvent.setup(); + const onRespond = vi.fn(async () => {}); + + render( + {})} + />, + ); + + const editor = screen.getByLabelText("Response"); + await user.click(editor); + await user.keyboard("{Enter}second"); + + expect(onRespond).not.toHaveBeenCalled(); + expect(editor).toHaveValue("first\nsecond"); + }); + + it("keeps editor replacement tags as literal plain text", () => { + const text = ''; + + expect(piExtensionEditorTextToContent(text)).toEqual({ + segments: [{ type: "text", text }], + }); + }); + + it("renders compact statuses and only widgets for the requested placement", () => { + render( + <> + + + , + ); + + expect(screen.getByText("Running")).toBeInTheDocument(); + expect(screen.getByRole("status")).toHaveAttribute("aria-live", "polite"); + expect(screen.getByText("Above content")).toBeInTheDocument(); + expect(screen.queryByText("Below content")).not.toBeInTheDocument(); + }); +}); diff --git a/products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.tsx b/products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.tsx new file mode 100644 index 000000000000..e521d452c467 --- /dev/null +++ b/products/desktop/packages/ui/src/features/pi-sessions/PiExtensionDialog.tsx @@ -0,0 +1,185 @@ +import type { PiExtensionUIResponse } from "@posthog/agent/pi/types"; +import type { PiExtensionDialogRequest } from "@posthog/core/pi-runtime/piSessionStore"; +import { + Button, + Dialog, + DialogBody, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, + Input, + Label, + Textarea, +} from "@posthog/quill"; +import { type FormEvent, useId, useRef, useState } from "react"; + +export function buildPiExtensionResponse( + request: PiExtensionDialogRequest, + value: string | boolean, +): PiExtensionUIResponse { + if (request.method === "confirm") { + return { + type: "extension_ui_response", + id: request.id, + confirmed: value === true, + }; + } + return { + type: "extension_ui_response", + id: request.id, + value: typeof value === "string" ? value : "", + }; +} + +interface PiExtensionDialogProps { + request: PiExtensionDialogRequest; + onRespond: (response: PiExtensionUIResponse) => Promise; + onCancel: () => Promise; +} + +export function PiExtensionDialog({ + request, + onRespond, + onCancel, +}: PiExtensionDialogProps) { + const [value, setValue] = useState( + request.method === "editor" ? (request.prefill ?? "") : "", + ); + const [submitting, setSubmitting] = useState(false); + const submittingRef = useRef(false); + const fieldId = useId(); + + const complete = async (response?: PiExtensionUIResponse): Promise => { + if (submittingRef.current) { + return; + } + submittingRef.current = true; + setSubmitting(true); + try { + if (response) { + await onRespond(response); + } else { + await onCancel(); + } + } catch { + // The controller retains the dialog and reports delivery failures by toast. + } finally { + submittingRef.current = false; + setSubmitting(false); + } + }; + + const submit = (event: FormEvent): void => { + event.preventDefault(); + if (request.method === "select") { + return; + } + void complete( + buildPiExtensionResponse( + request, + request.method === "confirm" ? true : value, + ), + ); + }; + + const description = + request.method === "select" + ? "Choose one of the available options." + : request.method === "confirm" + ? request.message + : request.method === "editor" + ? "Enter or edit the response." + : "Enter a response."; + + return ( + !open && void complete()}> + +
+ + {request.title} + {description} + + + {request.method === "select" ? ( +
+ {request.options.map((option) => ( + + ))} +
+ ) : request.method === "confirm" ? null : ( +
+ + {request.method === "editor" ? ( +