From 33bc0a7b83cde2a3cd2b568ed3107d9bdf7a768a Mon Sep 17 00:00:00 2001 From: Christian Klotz Date: Fri, 31 Jul 2026 00:41:53 +0300 Subject: [PATCH 1/5] feat(client): add runtime-neutral session client Split harness shutdown initiation from settlement so callbacks can request shutdown without deadlocking active work. --- package-lock.json | 19 + package.json | 4 +- packages/agent/src/harness/agent-harness.ts | 15 +- .../agent/test/harness/agent-harness.test.ts | 64 ++- packages/client/CHANGELOG.md | 7 + packages/client/README.md | 36 ++ packages/client/package.json | 35 ++ packages/client/src/client.ts | 398 ++++++++++++++++++ packages/client/src/errors.ts | 39 ++ packages/client/src/index.ts | 11 + packages/client/src/listeners.ts | 9 + packages/client/src/session-client.ts | 63 +++ packages/client/src/transport.ts | 18 + packages/client/src/types.ts | 51 +++ .../client/test/client-connection.test.ts | 306 ++++++++++++++ packages/client/test/client-state.test.ts | 200 +++++++++ packages/client/test/support.ts | 136 ++++++ packages/client/tsconfig.build.json | 12 + packages/client/tsconfig.test.json | 13 + packages/client/vitest.config.ts | 15 + scripts/browser-smoke-entry.ts | 2 + scripts/local-release.mjs | 1 + scripts/publish.mjs | 1 + tsconfig.json | 2 + 24 files changed, 1441 insertions(+), 16 deletions(-) create mode 100644 packages/client/CHANGELOG.md create mode 100644 packages/client/README.md create mode 100644 packages/client/package.json create mode 100644 packages/client/src/client.ts create mode 100644 packages/client/src/errors.ts create mode 100644 packages/client/src/index.ts create mode 100644 packages/client/src/listeners.ts create mode 100644 packages/client/src/session-client.ts create mode 100644 packages/client/src/transport.ts create mode 100644 packages/client/src/types.ts create mode 100644 packages/client/test/client-connection.test.ts create mode 100644 packages/client/test/client-state.test.ts create mode 100644 packages/client/test/support.ts create mode 100644 packages/client/tsconfig.build.json create mode 100644 packages/client/tsconfig.test.json create mode 100644 packages/client/vitest.config.ts diff --git a/package-lock.json b/package-lock.json index e55c44da23e..1149569d982 100644 --- a/package-lock.json +++ b/package-lock.json @@ -790,6 +790,10 @@ "resolved": "packages/ai", "link": true }, + "node_modules/@earendil-works/pi-client": { + "resolved": "packages/client", + "link": true + }, "node_modules/@earendil-works/pi-coding-agent": { "resolved": "packages/coding-agent", "link": true @@ -5501,6 +5505,21 @@ "dev": true, "license": "MIT" }, + "packages/client": { + "name": "@earendil-works/pi-client", + "version": "0.83.0", + "license": "MIT", + "dependencies": { + "@earendil-works/pi-protocol": "^0.83.0" + }, + "devDependencies": { + "shx": "0.4.0", + "vitest": "4.1.9" + }, + "engines": { + "node": ">=22.19.0" + } + }, "packages/coding-agent": { "name": "@earendil-works/pi-coding-agent", "version": "0.83.0", diff --git a/package.json b/package.json index 107d54bce7e..d6320579681 100644 --- a/package.json +++ b/package.json @@ -13,8 +13,8 @@ ], "scripts": { "clean": "npm run clean --workspaces", - "build": "cd packages/tui && npm run build && cd ../ai && npm run build && cd ../agent && npm run build && cd ../storage/sqlite-node && npm run build && cd ../../protocol && npm run build && cd ../coding-agent && npm run build && cd ../server && npm run build", - "build:offline": "cd packages/tui && npm run build && cd ../ai && npm run build:offline && cd ../agent && npm run build && cd ../storage/sqlite-node && npm run build && cd ../../protocol && npm run build && cd ../coding-agent && npm run build && cd ../server && npm run build", + "build": "cd packages/tui && npm run build && cd ../ai && npm run build && cd ../agent && npm run build && cd ../storage/sqlite-node && npm run build && cd ../../protocol && npm run build && cd ../client && npm run build && cd ../coding-agent && npm run build && cd ../server && npm run build", + "build:offline": "cd packages/tui && npm run build && cd ../ai && npm run build:offline && cd ../agent && npm run build && cd ../storage/sqlite-node && npm run build && cd ../../protocol && npm run build && cd ../client && npm run build && cd ../coding-agent && npm run build && cd ../server && npm run build", "check": "biome check --write --error-on-warnings . && npm run check:pinned-deps && npm run check:ts-imports && npm run check:shrinkwrap && npm run check:install-lock:coding-agent && tsgo --noEmit && npm run check:browser-smoke", "check:browser-smoke": "node scripts/check-browser-smoke.mjs", "check:pinned-deps": "node scripts/check-pinned-deps.mjs", diff --git a/packages/agent/src/harness/agent-harness.ts b/packages/agent/src/harness/agent-harness.ts index 2c9938ba3a0..dc8244e44ac 100644 --- a/packages/agent/src/harness/agent-harness.ts +++ b/packages/agent/src/harness/agent-harness.ts @@ -1100,12 +1100,9 @@ export class AgentHarness< this.streamOptions = cloneStreamOptions(streamOptions); } - /** - * Permanently stop this harness instance without deleting its durable session. - * Clears queued work, aborts the active operation, and waits for it to settle. - */ - async shutdown(): Promise { - if (this.shutdownPromise) return this.shutdownPromise; + /** Permanently stop this harness instance without deleting its durable session. */ + requestShutdown(): void { + if (this.isShutdown) return; this.isShutdown = true; this.pendingSessionWrites = []; this.steerQueue = []; @@ -1113,7 +1110,11 @@ export class AgentHarness< this.nextTurnQueue = []; this.activeAbortController?.abort(); this.shutdownPromise = this.waitForTasks(); - return this.shutdownPromise; + } + + /** Waits for work active when shutdown was requested to settle. */ + waitForShutdown(): Promise { + return this.shutdownPromise ?? Promise.resolve(); } async abort(): Promise { diff --git a/packages/agent/test/harness/agent-harness.test.ts b/packages/agent/test/harness/agent-harness.test.ts index 9b1ec714b0a..d59d0624496 100644 --- a/packages/agent/test/harness/agent-harness.test.ts +++ b/packages/agent/test/harness/agent-harness.test.ts @@ -156,10 +156,12 @@ describe("AgentHarness", () => { await harness.nextTurn("queued next turn"); let firstShutdownSettled = false; - const firstShutdown = harness.shutdown().then(() => { + harness.requestShutdown(); + const firstShutdown = harness.waitForShutdown().then(() => { firstShutdownSettled = true; }); - const secondShutdown = harness.shutdown(); + harness.requestShutdown(); + const secondShutdown = harness.waitForShutdown(); await Promise.resolve(); expect(signal?.aborted).toBe(true); @@ -177,6 +179,49 @@ describe("AgentHarness", () => { }); }); + it("allows a hook to request shutdown without deadlocking its operation", async () => { + const registration = newFaux(); + let providerCalls = 0; + registration.setResponses([ + () => { + providerCalls++; + return fauxAssistantMessage("must not run"); + }, + ]); + const harness = new AgentHarness({ + models, + session: new Session(new InMemorySessionStorage()), + model: registration.getModel(), + }); + harness.on("before_agent_start", () => { + harness.requestShutdown(); + return undefined; + }); + + await expect(harness.prompt("hello")).rejects.toMatchObject({ code: "invalid_state" }); + await expect(harness.waitForShutdown()).resolves.toBeUndefined(); + expect(providerCalls).toBe(0); + }); + + it("allows a subscriber to request shutdown without deadlocking its operation", async () => { + const registration = newFaux(); + registration.setResponses([() => fauxAssistantMessage("reply")]); + const harness = new AgentHarness({ + models, + session: new Session(new InMemorySessionStorage()), + model: registration.getModel(), + }); + let subscriberCalls = 0; + harness.subscribe(() => { + subscriberCalls++; + harness.requestShutdown(); + }); + + await expect(harness.prompt("hello")).resolves.toMatchObject({ role: "assistant", stopReason: "aborted" }); + await expect(harness.waitForShutdown()).resolves.toBeUndefined(); + expect(subscriberCalls).toBeGreaterThan(1); + }); + it("does not start a provider request when shutdown occurs during before_agent_start", async () => { const registration = newFaux(); const entered = deferred(); @@ -202,7 +247,8 @@ describe("AgentHarness", () => { await entered.promise; let shutdownSettled = false; - const shutdown = harness.shutdown().then(() => { + harness.requestShutdown(); + const shutdown = harness.waitForShutdown().then(() => { shutdownSettled = true; }); await Promise.resolve(); @@ -235,7 +281,8 @@ describe("AgentHarness", () => { await entered.promise; let shutdownSettled = false; - const shutdown = harness.shutdown().then(() => { + harness.requestShutdown(); + const shutdown = harness.waitForShutdown().then(() => { shutdownSettled = true; }); await Promise.resolve(); @@ -271,7 +318,8 @@ describe("AgentHarness", () => { await entered.promise; let shutdownSettled = false; - const shutdown = harness.shutdown().then(() => { + harness.requestShutdown(); + const shutdown = harness.waitForShutdown().then(() => { shutdownSettled = true; }); await Promise.resolve(); @@ -319,7 +367,8 @@ describe("AgentHarness", () => { ]; await storage.allWritesStarted.promise; - const shutdown = harness.shutdown(); + harness.requestShutdown(); + const shutdown = harness.waitForShutdown(); const firstSettlement = await Promise.race([ shutdown.then(() => "shutdown" as const), new Promise<"writes-pending">((resolve) => setImmediate(() => resolve("writes-pending"))), @@ -340,7 +389,8 @@ describe("AgentHarness", () => { model: getModel("anthropic", "claude-sonnet-4-5"), }); - await harness.shutdown(); + harness.requestShutdown(); + await harness.waitForShutdown(); const messages = (await session.getEntries()).flatMap((entry) => entry.type === "message" ? [entry.message] : [], diff --git a/packages/client/CHANGELOG.md b/packages/client/CHANGELOG.md new file mode 100644 index 00000000000..b94aff2dd32 --- /dev/null +++ b/packages/client/CHANGELOG.md @@ -0,0 +1,7 @@ +# Changelog + +## [Unreleased] + +### Added + +- Added the experimental transport-neutral `PiClient` and multi-session `PiSessionClient` APIs. diff --git a/packages/client/README.md b/packages/client/README.md new file mode 100644 index 00000000000..e251429bb4b --- /dev/null +++ b/packages/client/README.md @@ -0,0 +1,36 @@ +# @earendil-works/pi-client + +Transport-neutral client for remote pi sessions. `PiClient` exchanges length-prefixed CBOR messages through a small `ByteTransport` interface. The package has no Node-specific imports. + +```ts +import { PiClient, type ByteTransportFactory } from "@earendil-works/pi-client"; + +const transportFactory: ByteTransportFactory = async (handlers) => { + // Connect using WebSocket, Unix socket, or another ordered byte transport. + return { + async send(chunk) { + // Deliver chunks in invocation order and honor backpressure. + }, + close() {}, + }; +}; + +const client = new PiClient({ token: bearerToken, transportFactory }); +await client.connect(); +const session = await client.createSession({ cwd: "/workspace" }); +const unsubscribe = session.subscribe((snapshot) => render(snapshot)); +await session.prompt("Inspect this project"); +unsubscribe(); +``` + +Call `handlers.onData(chunk)` for inbound bytes, `handlers.onClose()` for an orderly terminal close, and `handlers.onError(error)` for transport failures. A factory must create a fresh transport for every connection attempt. + +`PiClient` does not reconnect automatically. Call `reconnect()` after disconnection. One connection can attach several `PiSessionClient` handles. Requests are correlated by ID. Server snapshots and successful response snapshots are authoritative, while progress events do not mutate snapshot state optimistically. + +`subscribe()` observes authoritative snapshots. `onEvent()` observes protocol events. Both return an unsubscribe function. A detached session handle remains readable, but commands throw `PiSessionDetachedError` until it is attached again. + +## Limits and security + +`PiClientOptions.maxFrameLength` bounds inbound and outbound CBOR payloads. Configure matching limits on the client and server. Transports should separately bound queued outbound bytes and preserve send order. + +Treat peers as untrusted. Use a secure transport where required and protect the protocol bearer token. diff --git a/packages/client/package.json b/packages/client/package.json new file mode 100644 index 00000000000..5ca572253b7 --- /dev/null +++ b/packages/client/package.json @@ -0,0 +1,35 @@ +{ + "name": "@earendil-works/pi-client", + "version": "0.83.0", + "description": "Transport-neutral client for remote pi sessions over framed CBOR bytes", + "type": "module", + "main": "./dist/index.js", + "types": "./dist/index.d.ts", + "exports": { + ".": { + "types": "./dist/index.d.ts", + "import": "./dist/index.js" + }, + "./package.json": "./package.json" + }, + "sideEffects": false, + "files": ["dist", "README.md", "CHANGELOG.md"], + "scripts": { + "clean": "shx rm -rf dist", + "build": "tsgo -p tsconfig.build.json", + "test": "vitest --run", + "typecheck": "tsgo -p tsconfig.test.json", + "prepublishOnly": "npm run clean && npm run build" + }, + "keywords": ["pi", "client", "protocol", "cbor", "binary"], + "author": "Earendil Works", + "license": "MIT", + "repository": { + "type": "git", + "url": "git+https://github.com/earendil-works/pi.git", + "directory": "packages/client" + }, + "engines": { "node": ">=22.19.0" }, + "dependencies": { "@earendil-works/pi-protocol": "^0.83.0" }, + "devDependencies": { "shx": "0.4.0", "vitest": "4.1.9" } +} diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts new file mode 100644 index 00000000000..d79d9741ca0 --- /dev/null +++ b/packages/client/src/client.ts @@ -0,0 +1,398 @@ +import { + type Command, + type CommandResult, + DEFAULT_MAX_FRAME_LENGTH, + encodeClientMessage, + PROTOCOL_VERSION, + ProtocolValidationError, + type ResultForCommand, + type ServerEvent, + type ServerMessage, + ServerMessageDecoder, + type ServerSnapshot, + type SessionSnapshot, + type SessionSummary, +} from "@earendil-works/pi-protocol"; +import { PiDisconnectedError, PiError, PiSessionDetachedError, toDisconnectedError, toError } from "./errors.ts"; +import { notifyListeners } from "./listeners.ts"; +import { PiSessionClient } from "./session-client.ts"; +import type { ByteTransport, ByteTransportHandlers } from "./transport.ts"; +import type { + ConnectionState, + ConnectionStateChange, + CreateSessionOptions, + PendingRequest, + PiClientOptions, + Unsubscribe, +} from "./types.ts"; + +const MAX_UINT32 = 0xffff_ffff; +export class PiClient { + private readonly options: PiClientOptions; + private readonly maxFrameLength: number; + private transport: ByteTransport | undefined; + private decoder: ServerMessageDecoder | undefined; + private connectionSequence = 0; + private stateValue: ConnectionState = "disconnected"; + private snapshotValue: ServerSnapshot | undefined; + private readonly sessionSnapshots = new Map(); + private readonly attachedSessionIds = new Set(); + private readonly sessionHandles = new Map(); + private readonly snapshotListeners = new Set<(snapshot: ServerSnapshot) => void>(); + private readonly eventListeners = new Set<(event: ServerEvent) => void>(); + private readonly stateListeners = new Set<(change: ConnectionStateChange) => void>(); + private readonly sessionSnapshotListeners = new Map void>>(); + private readonly sessionEventListeners = new Map void>>(); + private readonly pendingRequests = new Map(); + private requestSequence = 0; + private connectResolve: ((snapshot: ServerSnapshot) => void) | undefined; + private connectReject: ((error: Error) => void) | undefined; + + constructor(options: PiClientOptions) { + this.options = options; + this.maxFrameLength = options.maxFrameLength ?? DEFAULT_MAX_FRAME_LENGTH; + if (!Number.isSafeInteger(this.maxFrameLength) || this.maxFrameLength <= 0 || this.maxFrameLength > MAX_UINT32) { + throw new TypeError(`PiClient maxFrameLength must be an integer between 1 and ${MAX_UINT32}`); + } + } + + get connectionState(): ConnectionState { + return this.stateValue; + } + get connected(): boolean { + return this.stateValue === "connected"; + } + get snapshot(): ServerSnapshot | undefined { + return this.snapshotValue; + } + get sessions(): readonly SessionSummary[] { + return this.snapshotValue?.sessions ?? []; + } + + connect(): Promise { + if (this.stateValue !== "disconnected") { + return Promise.reject(new PiDisconnectedError(`PiClient is already ${this.stateValue}`)); + } + this.setConnectionState("connecting"); + this.snapshotValue = undefined; + this.sessionSnapshots.clear(); + this.attachedSessionIds.clear(); + this.decoder = new ServerMessageDecoder({ maxFrameLength: this.maxFrameLength }); + const connectionId = ++this.connectionSequence; + const connected = new Promise((resolve, reject) => { + this.connectResolve = resolve; + this.connectReject = reject; + }); + const handlers: ByteTransportHandlers = { + onData: (chunk) => { + if (!this.isCurrentConnection(connectionId)) return; + if (!this.transport) { + this.protocolFailure( + new ProtocolValidationError("Received server data before the client hello was sent"), + ); + return; + } + this.handleChunk(chunk); + }, + onClose: () => { + if (this.isCurrentConnection(connectionId)) this.handleTransportClose(); + }, + onError: (error) => { + if (this.isCurrentConnection(connectionId)) this.handleTransportError(error); + }, + }; + void this.openTransport(connectionId, handlers); + return connected; + } + + reconnect(): Promise { + return this.connect(); + } + disconnect(reason = "Client disconnected"): void { + if (this.stateValue === "disconnected") return; + const transport = this.transport; + this.failConnection(new PiDisconnectedError(reason)); + transport?.close(); + } + subscribe(listener: (snapshot: ServerSnapshot) => void): Unsubscribe { + this.snapshotListeners.add(listener); + return () => this.snapshotListeners.delete(listener); + } + onEvent(listener: (event: ServerEvent) => void): Unsubscribe { + this.eventListeners.add(listener); + return () => this.eventListeners.delete(listener); + } + onConnectionStateChange(listener: (change: ConnectionStateChange) => void): Unsubscribe { + this.stateListeners.add(listener); + return () => this.stateListeners.delete(listener); + } + getSession(sessionId: string): PiSessionClient { + let handle = this.sessionHandles.get(sessionId); + if (!handle) { + handle = new PiSessionClient(this, sessionId); + this.sessionHandles.set(sessionId, handle); + } + return handle; + } + getSessionSnapshot(sessionId: string): SessionSnapshot | undefined { + return this.sessionSnapshots.get(sessionId); + } + isSessionAttached(sessionId: string): boolean { + return this.attachedSessionIds.has(sessionId); + } + async listSessions(): Promise { + return (await this.request({ command: "list" })).sessions; + } + async createSession(options: CreateSessionOptions = {}): Promise { + const result = await this.request({ command: "create", ...options }); + return this.getSession(result.session.id); + } + async attachSession(sessionId: string): Promise { + const previous = this.sessionSnapshots.get(sessionId); + this.sessionSnapshots.delete(sessionId); + try { + await this.request({ command: "attach", sessionId }); + return this.getSession(sessionId); + } catch (error) { + if (previous && !this.sessionSnapshots.has(sessionId)) this.sessionSnapshots.set(sessionId, previous); + throw error; + } + } + async detachSession(sessionId: string): Promise { + await this.request({ command: "detach", sessionId }); + } + + request(command: TCommand): Promise> { + const transport = this.transport; + if (this.stateValue !== "connected" || !transport) return Promise.reject(new PiDisconnectedError()); + const id = `request-${++this.requestSequence}`; + let frame: Uint8Array; + try { + frame = encodeClientMessage( + { type: "request", id, request: command }, + { maxFrameLength: this.maxFrameLength }, + ); + } catch (error) { + return Promise.reject(toError(error)); + } + const promise = new Promise((resolve, reject) => { + this.pendingRequests.set(id, { command, resolve, reject }); + }); + this.sendFrame(transport, frame); + return promise as Promise>; + } + + subscribeSession(sessionId: string, listener: (snapshot: SessionSnapshot) => void): Unsubscribe { + let listeners = this.sessionSnapshotListeners.get(sessionId); + if (!listeners) { + listeners = new Set(); + this.sessionSnapshotListeners.set(sessionId, listeners); + } + listeners.add(listener); + return () => { + listeners.delete(listener); + if (listeners.size === 0) this.sessionSnapshotListeners.delete(sessionId); + }; + } + onSessionEvent(sessionId: string, listener: (event: ServerEvent) => void): Unsubscribe { + let listeners = this.sessionEventListeners.get(sessionId); + if (!listeners) { + listeners = new Set(); + this.sessionEventListeners.set(sessionId, listeners); + } + listeners.add(listener); + return () => { + listeners.delete(listener); + if (listeners.size === 0) this.sessionEventListeners.delete(sessionId); + }; + } + assertAttached(sessionId: string): void { + if (this.stateValue !== "connected") throw new PiDisconnectedError(); + if (!this.attachedSessionIds.has(sessionId)) throw new PiSessionDetachedError(sessionId); + } + + private async openTransport(connectionId: number, handlers: ByteTransportHandlers): Promise { + let transport: ByteTransport; + try { + transport = await this.options.transportFactory(handlers); + } catch (error) { + if (this.isCurrentConnection(connectionId)) this.failConnection(toDisconnectedError(error)); + return; + } + if (!this.isCurrentConnection(connectionId)) { + transport.close(); + return; + } + this.transport = transport; + try { + await transport.send( + encodeClientMessage( + { type: "hello", version: PROTOCOL_VERSION, token: this.options.token }, + { maxFrameLength: this.maxFrameLength }, + ), + ); + } catch (error) { + if (this.isCurrentConnection(connectionId)) { + this.failConnection(toDisconnectedError(error)); + transport.close(); + } + } + } + private sendFrame(transport: ByteTransport, frame: Uint8Array): void { + let sending: Promise; + try { + sending = transport.send(frame); + } catch (error) { + this.handleTransportError(toError(error)); + return; + } + void sending.catch((error: unknown) => { + if (this.transport === transport) this.handleTransportError(toError(error)); + }); + } + private handleChunk(chunk: Uint8Array): void { + let messages: ServerMessage[]; + try { + messages = this.decoder?.push(chunk) ?? []; + } catch (error) { + this.protocolFailure(toError(error)); + return; + } + for (const message of messages) { + if (this.stateValue === "disconnected") return; + this.handleMessage(message); + } + } + private handleMessage(message: ServerMessage): void { + if (this.stateValue === "connecting") { + if (message.type === "hello_error") { + const transport = this.transport; + this.failConnection(new PiError(message.error)); + transport?.close(); + return; + } + if (message.type !== "hello") { + this.protocolFailure(new ProtocolValidationError("Expected server hello as first message")); + return; + } + this.setConnectionState("connected"); + this.applyServerSnapshot(message.snapshot); + const resolve = this.connectResolve; + this.connectResolve = undefined; + this.connectReject = undefined; + resolve?.(message.snapshot); + return; + } + if (this.stateValue !== "connected") return; + if (message.type === "hello" || message.type === "hello_error") { + this.protocolFailure(new ProtocolValidationError("Unexpected handshake message")); + return; + } + if (message.type === "event") { + this.applyEvent(message.event); + return; + } + const pending = this.pendingRequests.get(message.id); + if (!pending) { + this.protocolFailure(new ProtocolValidationError("Response has no matching request")); + return; + } + this.pendingRequests.delete(message.id); + if (!message.ok) { + pending.reject(new PiError(message.error)); + return; + } + if (message.result.command !== pending.command.command) { + const error = new ProtocolValidationError( + `Response command ${message.result.command} does not match ${pending.command.command}`, + ); + pending.reject(error); + this.protocolFailure(error); + return; + } + this.applyResult(message.result); + pending.resolve(message.result); + } + private applyResult(result: CommandResult): void { + if (result.command === "list") return; + if (result.command === "detach") { + this.attachedSessionIds.delete(result.sessionId); + const snapshot = this.sessionSnapshots.get(result.sessionId); + if (snapshot) this.applySessionSnapshot({ ...snapshot, attached: false }, true); + return; + } + this.applySessionSnapshot(result.session); + } + private applyEvent(event: ServerEvent): void { + if (event.type === "server_snapshot") this.applyServerSnapshot(event.snapshot); + if (event.type === "session_snapshot") this.applySessionSnapshot(event.snapshot); + if (event.type === "session_removed") { + this.sessionSnapshots.delete(event.sessionId); + this.attachedSessionIds.delete(event.sessionId); + } + notifyListeners(this.eventListeners, event); + const sessionId = getEventSessionId(event); + if (sessionId) notifyListeners(this.sessionEventListeners.get(sessionId), event); + } + private applyServerSnapshot(snapshot: ServerSnapshot): void { + if (this.snapshotValue && snapshot.revision < this.snapshotValue.revision) return; + this.snapshotValue = snapshot; + this.attachedSessionIds.clear(); + for (const session of snapshot.sessions) if (session.attached) this.attachedSessionIds.add(session.id); + notifyListeners(this.snapshotListeners, snapshot); + } + private applySessionSnapshot(snapshot: SessionSnapshot, force = false): void { + const current = this.sessionSnapshots.get(snapshot.id); + if (!force && current && snapshot.revision < current.revision) return; + this.sessionSnapshots.set(snapshot.id, snapshot); + if (snapshot.attached) this.attachedSessionIds.add(snapshot.id); + else this.attachedSessionIds.delete(snapshot.id); + notifyListeners(this.sessionSnapshotListeners.get(snapshot.id), snapshot); + } + private handleTransportClose(): void { + let error: Error = new PiDisconnectedError("Byte transport closed"); + try { + this.decoder?.end(); + } catch (decoderError) { + error = toError(decoderError); + } + this.failConnection(error); + } + private handleTransportError(error: Error): void { + const transport = this.transport; + this.failConnection(toDisconnectedError(error)); + transport?.close(); + } + private protocolFailure(error: Error): void { + const transport = this.transport; + this.failConnection(error); + transport?.close(); + } + private failConnection(error: Error): void { + if (this.stateValue === "disconnected") return; + const reject = this.connectReject; + const pending = [...this.pendingRequests.values()]; + this.transport = undefined; + this.decoder = undefined; + this.connectResolve = undefined; + this.connectReject = undefined; + this.pendingRequests.clear(); + this.attachedSessionIds.clear(); + reject?.(error); + for (const request of pending) request.reject(error); + this.setConnectionState("disconnected", error); + } + private isCurrentConnection(connectionId: number): boolean { + return connectionId === this.connectionSequence && this.stateValue !== "disconnected"; + } + private setConnectionState(state: ConnectionState, error?: Error): void { + this.stateValue = state; + notifyListeners(this.stateListeners, error ? { state, error } : { state }); + } +} + +function getEventSessionId(event: ServerEvent): string | undefined { + if (event.type === "session_snapshot") return event.snapshot.id; + if (event.type === "session_progress" || event.type === "session_removed") return event.sessionId; + return undefined; +} diff --git a/packages/client/src/errors.ts b/packages/client/src/errors.ts new file mode 100644 index 00000000000..54b6d16be1e --- /dev/null +++ b/packages/client/src/errors.ts @@ -0,0 +1,39 @@ +import type { JsonValue, ProtocolError, ProtocolErrorCode } from "@earendil-works/pi-protocol"; + +export class PiError extends Error { + readonly code: ProtocolErrorCode; + readonly details: JsonValue | undefined; + + constructor(error: ProtocolError) { + super(error.message); + this.name = "PiError"; + this.code = error.code; + this.details = error.details; + } +} + +export class PiDisconnectedError extends Error { + constructor(message = "Pi client is disconnected") { + super(message); + this.name = "PiDisconnectedError"; + } +} + +export class PiSessionDetachedError extends Error { + readonly sessionId: string; + + constructor(sessionId: string) { + super(`Session ${sessionId} is not attached`); + this.name = "PiSessionDetachedError"; + this.sessionId = sessionId; + } +} + +export function toError(error: unknown): Error { + return error instanceof Error ? error : new Error(String(error)); +} + +export function toDisconnectedError(error: unknown): PiDisconnectedError { + const cause = toError(error); + return cause instanceof PiDisconnectedError ? cause : new PiDisconnectedError(cause.message); +} diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts new file mode 100644 index 00000000000..d7ddefdb5a8 --- /dev/null +++ b/packages/client/src/index.ts @@ -0,0 +1,11 @@ +export { PiClient } from "./client.ts"; +export { PiDisconnectedError, PiError, PiSessionDetachedError } from "./errors.ts"; +export { PiSessionClient } from "./session-client.ts"; +export type { ByteTransport, ByteTransportFactory, ByteTransportHandlers } from "./transport.ts"; +export type { + ConnectionState, + ConnectionStateChange, + CreateSessionOptions, + PiClientOptions, + Unsubscribe, +} from "./types.ts"; diff --git a/packages/client/src/listeners.ts b/packages/client/src/listeners.ts new file mode 100644 index 00000000000..a94680c4c17 --- /dev/null +++ b/packages/client/src/listeners.ts @@ -0,0 +1,9 @@ +export function notifyListeners(listeners: Iterable<(value: T) => void> | undefined, value: T): void { + for (const listener of listeners ?? []) { + try { + listener(value); + } catch { + // Consumer callbacks cannot affect protocol or transport state. + } + } +} diff --git a/packages/client/src/session-client.ts b/packages/client/src/session-client.ts new file mode 100644 index 00000000000..dccaccb97fd --- /dev/null +++ b/packages/client/src/session-client.ts @@ -0,0 +1,63 @@ +import type { ModelRef, ServerEvent, SessionSnapshot, ThinkingLevel } from "@earendil-works/pi-protocol"; +import type { SessionClientHost, Unsubscribe } from "./types.ts"; + +export class PiSessionClient { + readonly id: string; + private readonly client: SessionClientHost; + + constructor(client: SessionClientHost, id: string) { + this.client = client; + this.id = id; + } + + get attached(): boolean { + return this.client.isSessionAttached(this.id); + } + + get snapshot(): SessionSnapshot | undefined { + return this.client.getSessionSnapshot(this.id); + } + + subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe { + return this.client.subscribeSession(this.id, listener); + } + + onEvent(listener: (event: ServerEvent) => void): Unsubscribe { + return this.client.onSessionEvent(this.id, listener); + } + + async detach(): Promise { + this.client.assertAttached(this.id); + await this.client.detachSession(this.id); + } + + async prompt(text: string): Promise { + this.client.assertAttached(this.id); + const result = await this.client.request({ command: "prompt", sessionId: this.id, text }); + return result.session; + } + + async steer(text: string): Promise { + this.client.assertAttached(this.id); + const result = await this.client.request({ command: "steer", sessionId: this.id, text }); + return result.session; + } + + async abort(): Promise { + this.client.assertAttached(this.id); + const result = await this.client.request({ command: "abort", sessionId: this.id }); + return result.session; + } + + async setModel(model: ModelRef): Promise { + this.client.assertAttached(this.id); + const result = await this.client.request({ command: "set_model", sessionId: this.id, model }); + return result.session; + } + + async setThinking(thinkingLevel: ThinkingLevel): Promise { + this.client.assertAttached(this.id); + const result = await this.client.request({ command: "set_thinking", sessionId: this.id, thinkingLevel }); + return result.session; + } +} diff --git a/packages/client/src/transport.ts b/packages/client/src/transport.ts new file mode 100644 index 00000000000..71b7489c360 --- /dev/null +++ b/packages/client/src/transport.ts @@ -0,0 +1,18 @@ +export interface ByteTransport { + /** Sends one byte chunk. Calls must be delivered in invocation order. */ + send(chunk: Uint8Array): Promise; + /** Closes the transport. Implementations must make repeated calls harmless. */ + close(): void; +} + +export interface ByteTransportHandlers { + /** Delivers an arbitrary inbound byte chunk. */ + onData(chunk: Uint8Array): void; + /** Reports an orderly terminal close. */ + onClose(): void; + /** Reports a terminal transport failure. */ + onError(error: Error): void; +} + +/** Creates a fresh connected transport for each PiClient connection attempt. Exactly one terminal handler is expected. */ +export type ByteTransportFactory = (handlers: ByteTransportHandlers) => ByteTransport | Promise; diff --git a/packages/client/src/types.ts b/packages/client/src/types.ts new file mode 100644 index 00000000000..aba7b496687 --- /dev/null +++ b/packages/client/src/types.ts @@ -0,0 +1,51 @@ +import type { + Command, + CommandResult, + ModelRef, + ResultForCommand, + ServerEvent, + ServerSnapshot, + SessionSnapshot, + ThinkingLevel, +} from "@earendil-works/pi-protocol"; +import type { ByteTransportFactory } from "./transport.ts"; + +export type ConnectionState = "disconnected" | "connecting" | "connected"; + +export interface ConnectionStateChange { + state: ConnectionState; + error?: Error; +} + +export type Unsubscribe = () => void; + +export interface PiClientOptions { + token: string; + transportFactory: ByteTransportFactory; + maxFrameLength?: number; +} + +export interface CreateSessionOptions { + cwd?: string; + name?: string; + model?: ModelRef; + thinkingLevel?: ThinkingLevel; +} + +export interface SessionClientHost { + isSessionAttached(sessionId: string): boolean; + getSessionSnapshot(sessionId: string): SessionSnapshot | undefined; + detachSession(sessionId: string): Promise; + request(command: TCommand): Promise>; + subscribeSession(sessionId: string, listener: (snapshot: SessionSnapshot) => void): Unsubscribe; + onSessionEvent(sessionId: string, listener: (event: ServerEvent) => void): Unsubscribe; + assertAttached(sessionId: string): void; +} + +export interface PendingRequest { + command: Command; + resolve(result: CommandResult): void; + reject(error: Error): void; +} + +export type ServerSnapshotListener = (snapshot: ServerSnapshot) => void; diff --git a/packages/client/test/client-connection.test.ts b/packages/client/test/client-connection.test.ts new file mode 100644 index 00000000000..d76d129539b --- /dev/null +++ b/packages/client/test/client-connection.test.ts @@ -0,0 +1,306 @@ +import { + type ClientMessage, + encodeServerMessage, + PROTOCOL_VERSION, + type ServerSnapshot, +} from "@earendil-works/pi-protocol"; +import { describe, expect, test } from "vitest"; +import { type ByteTransportFactory, PiClient, PiDisconnectedError, PiSessionDetachedError } from "../src/index.ts"; +import { + baseServerSnapshot, + collectRequests, + connectClient, + createClient, + MemoryByteServer, + sessionSnapshot, +} from "./support.ts"; + +describe("PiClient", () => { + test("sends a framed version and bearer token before accepting a fragmented server hello", async () => { + const server = new MemoryByteServer(); + const received: ClientMessage[] = []; + server.onMessage((message) => { + received.push(message); + if (message.type === "hello") { + server.send( + { + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }, + 3, + ); + } + }); + const client = createClient(server); + + await expect(client.connect()).resolves.toEqual(baseServerSnapshot); + expect(received[0]).toEqual({ type: "hello", version: PROTOCOL_VERSION, token: "bearer-secret" }); + expect(server.sentByClient[0]).toBeInstanceOf(Uint8Array); + expect(client.connectionState).toBe("connected"); + }); + + test("rejects server data delivered before sending the client hello", async () => { + let closeCount = 0; + let sendCount = 0; + const client = new PiClient({ + token: "bearer-secret", + transportFactory: (handlers) => { + handlers.onData( + encodeServerMessage({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }), + ); + return { + async send() { + sendCount++; + }, + close() { + closeCount++; + }, + }; + }, + }); + + await expect(client.connect()).rejects.toMatchObject({ + name: "ProtocolValidationError", + message: "Received server data before the client hello was sent", + }); + expect(client.connectionState).toBe("disconnected"); + expect(sendCount).toBe(0); + expect(closeCount).toBe(1); + }); + + test("isolates subscriber failures from handshake and transport state", async () => { + const server = new MemoryByteServer(); + server.onMessage((message) => { + if (message.type === "hello") { + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }); + } + }); + const client = createClient(server); + client.subscribe(() => { + throw new Error("consumer failure"); + }); + + await expect(client.connect()).resolves.toEqual(baseServerSnapshot); + expect(client.connectionState).toBe("connected"); + }); + + test("rejects a typed handshake authentication error", async () => { + const server = new MemoryByteServer(); + server.onMessage(() => { + server.send({ + type: "hello_error", + error: { code: "auth", message: "Invalid token" }, + }); + }); + const client = createClient(server, "wrong"); + + await expect(client.connect()).rejects.toMatchObject({ + name: "PiError", + code: "auth", + message: "Invalid token", + }); + expect(client.connectionState).toBe("disconnected"); + expect(server.clientCloseCount).toBe(1); + }); + + test("correlates coalesced out-of-order responses", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const listed = client.listSessions(); + const attached = client.attachSession("session-1"); + expect(requests).toHaveLength(2); + + const attachRequest = requests.find((request) => request.request.command === "attach"); + const listRequest = requests.find((request) => request.request.command === "list"); + if (!attachRequest || !listRequest) throw new Error("Missing requests"); + server.sendTogether([ + { + type: "response", + id: attachRequest.id, + ok: true, + result: { command: "attach", session: sessionSnapshot("session-1") }, + }, + { + type: "response", + id: listRequest.id, + ok: true, + result: { command: "list", sessions: [] }, + }, + ]); + + await expect(listed).resolves.toEqual([]); + await expect(attached).resolves.toMatchObject({ id: "session-1", attached: true }); + }); + + test("reduces only authoritative snapshots and supports unsubscribe", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const initial = sessionSnapshot("session-1", { revision: 1, phase: "idle" }); + server.send({ type: "event", event: { type: "session_snapshot", snapshot: initial } }); + const handle = client.getSession("session-1"); + const observed: number[] = []; + const progressTypes: string[] = []; + const unsubscribe = handle.subscribe((snapshot) => observed.push(snapshot.revision)); + const unsubscribeEvents = handle.onEvent((event) => progressTypes.push(event.type)); + server.send({ + type: "event", + event: { + type: "session_progress", + sessionId: "session-1", + progress: { + type: "assistant_delta", + messageId: "assistant-1", + contentIndex: 0, + kind: "text", + delta: "hi", + }, + }, + }); + expect(progressTypes).toEqual(["session_progress"]); + expect(handle.snapshot).toEqual(initial); + + const prompting = handle.prompt("hello"); + expect(handle.snapshot).toEqual(initial); + const promptRequest = requests.find((request) => request.request.command === "prompt"); + if (!promptRequest) throw new Error("Missing prompt request"); + const updated = sessionSnapshot("session-1", { revision: 2, phase: "turn" }); + server.send({ + type: "response", + id: promptRequest.id, + ok: true, + result: { command: "prompt", session: updated }, + }); + await expect(prompting).resolves.toEqual(updated); + expect(handle.snapshot).toEqual(updated); + expect(observed).toEqual([2]); + + unsubscribe(); + unsubscribeEvents(); + server.send({ + type: "event", + event: { type: "session_snapshot", snapshot: sessionSnapshot("session-1", { revision: 3 }) }, + }); + expect(observed).toEqual([2]); + }); + + test("keeps multiple session handles independent and enforces detach", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + server.onMessage((message) => { + if (message.type !== "request") return; + const request = message.request; + if (request.command === "attach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { command: "attach", session: sessionSnapshot(request.sessionId) }, + }); + } + if (request.command === "detach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { command: "detach", sessionId: request.sessionId }, + }); + } + }); + + const first = await client.attachSession("session-1"); + const second = await client.attachSession("session-2"); + expect(first.attached).toBe(true); + expect(second.attached).toBe(true); + await first.detach(); + expect(first.attached).toBe(false); + expect(second.attached).toBe(true); + await expect(first.abort()).rejects.toBeInstanceOf(PiSessionDetachedError); + }); + + test("rejects pending requests on close and reconnects through a fresh factory result", async () => { + const first = new MemoryByteServer(); + const second = new MemoryByteServer(); + let connection = 0; + for (const server of [first, second]) { + server.onMessage((message) => { + if (message.type === "hello") { + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: `connection-${connection}`, + snapshot: { ...baseServerSnapshot, revision: connection }, + }); + } + }); + } + const transportFactory: ByteTransportFactory = (handlers) => + (connection++ === 0 ? first : second).connect(handlers); + const client = new PiClient({ token: "bearer-secret", transportFactory }); + const states: string[] = []; + client.onConnectionStateChange(({ state }) => states.push(state)); + await client.connect(); + const pending = client.listSessions(); + first.close(); + await expect(pending).rejects.toBeInstanceOf(PiDisconnectedError); + expect(client.connectionState).toBe("disconnected"); + + await expect(client.reconnect()).resolves.toMatchObject({ revision: 2 }); + expect(client.connectionState).toBe("connected"); + expect(states).toEqual(["connecting", "connected", "disconnected", "connecting", "connected"]); + }); + + test("supports synchronous reconnect from a disconnection listener", async () => { + const first = new MemoryByteServer(); + const second = new MemoryByteServer(); + let connection = 0; + for (const server of [first, second]) { + server.onMessage((message) => { + if (message.type !== "hello") return; + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: `connection-${connection}`, + snapshot: { ...baseServerSnapshot, revision: connection }, + }); + }); + } + const client = new PiClient({ + token: "bearer-secret", + transportFactory: (handlers) => (connection++ === 0 ? first : second).connect(handlers), + }); + await client.connect(); + let reconnect: Promise | undefined; + client.onConnectionStateChange(({ state }) => { + if (state === "disconnected") reconnect = client.reconnect(); + }); + + first.close(); + expect(reconnect).toBeDefined(); + await expect(reconnect).resolves.toMatchObject({ revision: 2 }); + expect(client.connectionState).toBe("connected"); + }); + + test("rejects pending requests on transport errors", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const pending = client.listSessions(); + server.error(new Error("read failed")); + + await expect(pending).rejects.toMatchObject({ name: "PiDisconnectedError", message: "read failed" }); + expect(client.connectionState).toBe("disconnected"); + }); +}); diff --git a/packages/client/test/client-state.test.ts b/packages/client/test/client-state.test.ts new file mode 100644 index 00000000000..0f894cf2bb1 --- /dev/null +++ b/packages/client/test/client-state.test.ts @@ -0,0 +1,200 @@ +import { encodeCbor, encodeFrame, PROTOCOL_VERSION, ProtocolValidationError } from "@earendil-works/pi-protocol"; +import { describe, expect, test } from "vitest"; +import { PiClient } from "../src/index.ts"; +import { baseServerSnapshot, collectRequests, connectClient, MemoryByteServer, sessionSnapshot } from "./support.ts"; + +describe("PiClient", () => { + test("enforces the configured frame limit for outbound and inbound messages", async () => { + const server = new MemoryByteServer(); + server.onMessage((message) => { + if (message.type === "hello") { + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }); + } + }); + const client = new PiClient({ + token: "bearer-secret", + maxFrameLength: 512, + transportFactory: (handlers) => server.connect(handlers), + }); + await client.connect(); + const sentBefore = server.sentByClient.length; + await expect( + client.request({ command: "prompt", sessionId: "session-1", text: "x".repeat(1_000) }), + ).rejects.toBeInstanceOf(ProtocolValidationError); + expect(server.sentByClient).toHaveLength(sentBefore); + + server.sendRaw(new Uint8Array([0, 0, 2, 1])); + expect(client.connectionState).toBe("disconnected"); + }); + + test("disconnects on invalid protocol data", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + server.sendRaw(encodeFrame(encodeCbor({ type: "event", event: { type: "session_removed", sessionId: 1 } }))); + expect(client.connectionState).toBe("disconnected"); + }); + + test("reports truncated framing when the transport closes", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const pending = client.listSessions(); + server.sendRaw(new Uint8Array([0, 0, 0, 2, 1])); + server.close(); + + await expect(pending).rejects.toMatchObject({ + name: "ProtocolValidationError", + message: expect.stringMatching(/truncated/i), + }); + expect(client.connectionState).toBe("disconnected"); + }); + + test("rejects a mismatched response instead of leaving its request pending", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const listed = client.listSessions(); + expect(requests).toMatchObject([{ request: { command: "list" } }]); + server.send({ + type: "response", + id: requests[0]!.id, + ok: true, + result: { command: "attach", session: sessionSnapshot("session-1") }, + }); + + await expect(listed).rejects.toMatchObject({ + name: "ProtocolValidationError", + message: "Response command attach does not match list", + }); + expect(client.connectionState).toBe("disconnected"); + }); + + test("does not let a delayed command response replace a newer event snapshot", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const initial = sessionSnapshot("session-1", { revision: 1, thinkingLevel: "off" }); + server.send({ type: "event", event: { type: "session_snapshot", snapshot: initial } }); + const handle = client.getSession("session-1"); + const requests = collectRequests(server); + const changing = handle.setThinking("high"); + const request = requests.find((candidate) => candidate.request.command === "set_thinking"); + if (!request) throw new Error("Missing set_thinking request"); + server.send({ + type: "event", + event: { + type: "session_snapshot", + snapshot: sessionSnapshot("session-1", { revision: 3, thinkingLevel: "high" }), + }, + }); + server.send({ + type: "response", + id: request.id, + ok: true, + result: { + command: "set_thinking", + session: sessionSnapshot("session-1", { revision: 2, thinkingLevel: "medium" }), + }, + }); + + await changing; + expect(handle.snapshot).toMatchObject({ revision: 3, thinkingLevel: "high" }); + }); + + test("does not let an attach response replace a newer snapshot from the reacquired runtime", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + server.send({ + type: "event", + event: { + type: "session_snapshot", + snapshot: sessionSnapshot("session-1", { revision: 10, attached: false }), + }, + }); + server.onMessage((message) => { + if (message.type !== "request" || message.request.command !== "attach") return; + server.send({ + type: "event", + event: { + type: "session_snapshot", + snapshot: sessionSnapshot("session-1", { revision: 3, thinkingLevel: "high" }), + }, + }); + server.send({ + type: "response", + id: message.id, + ok: true, + result: { + command: "attach", + session: sessionSnapshot("session-1", { revision: 2, thinkingLevel: "medium" }), + }, + }); + }); + + const handle = await client.attachSession("session-1"); + expect(handle.snapshot).toMatchObject({ revision: 3, thinkingLevel: "high" }); + }); + + test("accepts a lower revision after detaching and reacquiring the same session", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + let attachCount = 0; + server.onMessage((message) => { + if (message.type !== "request") return; + if (message.request.command === "attach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { + command: "attach", + session: sessionSnapshot("session-1", { revision: attachCount++ === 0 ? 10 : 0 }), + }, + }); + } + if (message.request.command === "detach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { command: "detach", sessionId: "session-1" }, + }); + } + }); + + const first = await client.attachSession("session-1"); + expect(first.snapshot?.revision).toBe(10); + await first.detach(); + const reopened = await client.attachSession("session-1"); + expect(reopened.snapshot?.revision).toBe(0); + }); + + test("rejects frame limits outside the unsigned 32-bit range", () => { + const server = new MemoryByteServer(); + expect( + () => + new PiClient({ + token: "secret", + maxFrameLength: 0x1_0000_0000, + transportFactory: (handlers) => server.connect(handlers), + }), + ).toThrow(/maxFrameLength/); + }); + + test("surfaces typed request errors", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const attaching = client.attachSession("locked"); + server.send({ + type: "response", + id: requests[0]?.id ?? "missing", + ok: false, + error: { code: "session_locked", message: "Already attached" }, + }); + await expect(attaching).rejects.toMatchObject({ name: "PiError", code: "session_locked" }); + }); +}); diff --git a/packages/client/test/support.ts b/packages/client/test/support.ts new file mode 100644 index 00000000000..03818670913 --- /dev/null +++ b/packages/client/test/support.ts @@ -0,0 +1,136 @@ +import { + type ClientMessage, + ClientMessageDecoder, + encodeServerMessage, + PROTOCOL_VERSION, + type RequestEnvelope, + type ServerMessage, + type ServerSnapshot, + type SessionSnapshot, +} from "@earendil-works/pi-protocol"; +import type { ByteTransport, ByteTransportHandlers } from "../src/index.ts"; +import { PiClient } from "../src/index.ts"; + +export class MemoryByteServer { + private handlers: ByteTransportHandlers | undefined; + private readonly decoder = new ClientMessageDecoder(); + private readonly messageListeners = new Set<(message: ClientMessage) => void>(); + public readonly sentByClient: Uint8Array[] = []; + public clientCloseCount = 0; + + connect(handlers: ByteTransportHandlers): ByteTransport { + this.handlers = handlers; + let closed = false; + return { + send: async (chunk) => { + if (closed) throw new Error("Transport is closed"); + this.sentByClient.push(chunk.slice()); + for (const message of this.decoder.push(chunk)) { + for (const listener of this.messageListeners) listener(message); + } + }, + close: () => { + if (closed) return; + closed = true; + this.clientCloseCount++; + }, + }; + } + + onMessage(listener: (message: ClientMessage) => void): () => void { + this.messageListeners.add(listener); + return () => this.messageListeners.delete(listener); + } + + send(message: ServerMessage, splitAt?: number): void { + const frame = encodeServerMessage(message); + if (splitAt === undefined) { + this.sendRaw(frame); + return; + } + this.sendRaw(frame.subarray(0, splitAt)); + this.sendRaw(frame.subarray(splitAt)); + } + + sendTogether(messages: ServerMessage[]): void { + const frames = messages.map((message) => encodeServerMessage(message)); + const length = frames.reduce((total, frame) => total + frame.byteLength, 0); + const chunk = new Uint8Array(length); + let offset = 0; + for (const frame of frames) { + chunk.set(frame, offset); + offset += frame.byteLength; + } + this.sendRaw(chunk); + } + + sendRaw(chunk: Uint8Array): void { + this.handlers?.onData(chunk); + } + + close(): void { + this.handlers?.onClose(); + } + + error(error: Error): void { + this.handlers?.onError(error); + } +} + +export const baseServerSnapshot: ServerSnapshot = { + serverId: "server-1", + protocolVersion: PROTOCOL_VERSION, + revision: 1, + sessions: [], + models: [], +}; + +export function sessionSnapshot(id: string, overrides: Partial = {}): SessionSnapshot { + return { + id, + cwd: "/workspace", + createdAt: 1, + updatedAt: 1, + phase: "idle", + model: { provider: "faux", id: "model" }, + thinkingLevel: "off", + attached: true, + locked: true, + revision: 1, + transcript: [], + queuedSteer: [], + queuedSteerCount: 0, + ...overrides, + }; +} + +export function createClient(server: MemoryByteServer, token = "bearer-secret"): PiClient { + return new PiClient({ + token, + transportFactory: (handlers) => server.connect(handlers), + }); +} + +export async function connectClient(server: MemoryByteServer, token = "bearer-secret"): Promise { + const client = createClient(server, token); + server.onMessage((message) => { + if (message.type === "hello") { + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }); + } + }); + await client.connect(); + return client; +} + +export function collectRequests(server: MemoryByteServer): RequestEnvelope[] { + const requests: RequestEnvelope[] = []; + server.onMessage((message) => { + if (message.type === "request") requests.push(message); + }); + return requests; +} diff --git a/packages/client/tsconfig.build.json b/packages/client/tsconfig.build.json new file mode 100644 index 00000000000..0e72ae51a6c --- /dev/null +++ b/packages/client/tsconfig.build.json @@ -0,0 +1,12 @@ +{ + "extends": "../../tsconfig.base.json", + "compilerOptions": { + "outDir": "./dist", + "rootDir": "./src", + "paths": { + "@earendil-works/pi-protocol": ["../protocol/dist/index.d.ts"] + } + }, + "include": ["src/**/*.ts"], + "exclude": ["node_modules", "dist", "**/*.d.ts", "src/**/*.d.ts"] +} diff --git a/packages/client/tsconfig.test.json b/packages/client/tsconfig.test.json new file mode 100644 index 00000000000..b5e61b3d6c1 --- /dev/null +++ b/packages/client/tsconfig.test.json @@ -0,0 +1,13 @@ +{ + "extends": "../../tsconfig.base.json", + "compilerOptions": { + "noEmit": true, + "module": "NodeNext", + "moduleResolution": "NodeNext", + "types": ["node", "vitest"], + "paths": { + "@earendil-works/pi-protocol": ["../protocol/src/index.ts"] + } + }, + "include": ["src/**/*.ts", "test/**/*.ts"] +} diff --git a/packages/client/vitest.config.ts b/packages/client/vitest.config.ts new file mode 100644 index 00000000000..7e6abe3f97f --- /dev/null +++ b/packages/client/vitest.config.ts @@ -0,0 +1,15 @@ +import { fileURLToPath } from "node:url"; +import { defineConfig } from "vitest/config"; + +export default defineConfig({ + test: { + globals: true, + environment: "node", + reporters: process.env.GITHUB_ACTIONS ? ["dot", "github-actions"] : ["dot"], + }, + resolve: { + alias: { + "@earendil-works/pi-protocol": fileURLToPath(new URL("../protocol/src/index.ts", import.meta.url)), + }, + }, +}); diff --git a/scripts/browser-smoke-entry.ts b/scripts/browser-smoke-entry.ts index 64927b6e526..bfb81190175 100644 --- a/scripts/browser-smoke-entry.ts +++ b/scripts/browser-smoke-entry.ts @@ -1,3 +1,4 @@ +import { PiClient } from "@earendil-works/pi-client"; import { createAssistantMessageEventStream, Type } from "@earendil-works/pi-ai"; import { complete, getModel, getProviders, streamSimple } from "@earendil-works/pi-ai/compat"; import { @@ -59,6 +60,7 @@ console.log( new FileError("not_found", "missing").code, toError("boom").message, typeof streamProxy, + typeof PiClient, PROTOCOL_VERSION, decodeCbor(encodeCbor({ browser: true })), ); diff --git a/scripts/local-release.mjs b/scripts/local-release.mjs index 3b275fc884f..dd046fadaa0 100644 --- a/scripts/local-release.mjs +++ b/scripts/local-release.mjs @@ -10,6 +10,7 @@ const packages = [ { directory: "packages/tui", name: "@earendil-works/pi-tui" }, { directory: "packages/agent", name: "@earendil-works/pi-agent-core" }, { directory: "packages/protocol", name: "@earendil-works/pi-protocol" }, + { directory: "packages/client", name: "@earendil-works/pi-client" }, { directory: "packages/storage/sqlite-node", name: "@earendil-works/pi-storage-sqlite-node" }, { directory: "packages/coding-agent", name: "@earendil-works/pi-coding-agent" }, ]; diff --git a/scripts/publish.mjs b/scripts/publish.mjs index 967f61373cb..0053f6b1c43 100644 --- a/scripts/publish.mjs +++ b/scripts/publish.mjs @@ -8,6 +8,7 @@ const packages = [ { directory: "packages/ai", name: "@earendil-works/pi-ai" }, { directory: "packages/agent", name: "@earendil-works/pi-agent-core" }, { directory: "packages/protocol", name: "@earendil-works/pi-protocol" }, + { directory: "packages/client", name: "@earendil-works/pi-client" }, { directory: "packages/storage/sqlite-node", name: "@earendil-works/pi-storage-sqlite-node" }, { directory: "packages/tui", name: "@earendil-works/pi-tui" }, { directory: "packages/coding-agent", name: "@earendil-works/pi-coding-agent" }, diff --git a/tsconfig.json b/tsconfig.json index 4d36db2f76a..a409cdf1107 100644 --- a/tsconfig.json +++ b/tsconfig.json @@ -18,6 +18,8 @@ "@earendil-works/pi-coding-agent/*": ["./packages/coding-agent/src/*"], "@earendil-works/pi-protocol": ["./packages/protocol/src/index.ts"], "@earendil-works/pi-protocol/*": ["./packages/protocol/src/*"], + "@earendil-works/pi-client": ["./packages/client/src/index.ts"], + "@earendil-works/pi-client/*": ["./packages/client/src/*"], "@earendil-works/pi-server": ["./packages/server/src/index.ts"], "@earendil-works/pi-server/*": ["./packages/server/src/*"], "typebox": ["./node_modules/typebox"], From 7e121c67cff54c44f963787a103179b3c7baadae Mon Sep 17 00:00:00 2001 From: Christian Klotz Date: Fri, 31 Jul 2026 01:11:12 +0300 Subject: [PATCH 2/5] refactor(client): tighten lifecycle state modeling --- packages/agent/src/harness/agent-harness.ts | 5 +- .../agent/test/harness/agent-harness.test.ts | 13 + packages/client/README.md | 2 + packages/client/src/client.ts | 422 ++++++++---------- packages/client/src/index.ts | 1 + packages/client/src/listeners.ts | 18 +- packages/client/src/pending-requests.ts | 33 ++ packages/client/src/session-client.ts | 63 +-- packages/client/src/state-store.ts | 141 ++++++ packages/client/src/types.ts | 32 +- .../client/test/client-connection.test.ts | 27 ++ tsconfig.base.json | 2 +- 12 files changed, 470 insertions(+), 289 deletions(-) create mode 100644 packages/client/src/pending-requests.ts create mode 100644 packages/client/src/state-store.ts diff --git a/packages/agent/src/harness/agent-harness.ts b/packages/agent/src/harness/agent-harness.ts index dc8244e44ac..53b1f334706 100644 --- a/packages/agent/src/harness/agent-harness.ts +++ b/packages/agent/src/harness/agent-harness.ts @@ -1114,7 +1114,10 @@ export class AgentHarness< /** Waits for work active when shutdown was requested to settle. */ waitForShutdown(): Promise { - return this.shutdownPromise ?? Promise.resolve(); + if (!this.shutdownPromise) { + return Promise.reject(new AgentHarnessError("invalid_state", "Shutdown has not been requested")); + } + return this.shutdownPromise; } async abort(): Promise { diff --git a/packages/agent/test/harness/agent-harness.test.ts b/packages/agent/test/harness/agent-harness.test.ts index d59d0624496..c10222ecff1 100644 --- a/packages/agent/test/harness/agent-harness.test.ts +++ b/packages/agent/test/harness/agent-harness.test.ts @@ -131,6 +131,19 @@ describe("AgentHarness", () => { expect(harness.getFollowUpMode()).toBe("one-at-a-time"); }); + it("rejects waiting before shutdown is requested", async () => { + const harness = new AgentHarness({ + models, + session: new Session(new InMemorySessionStorage()), + model: getModel("anthropic", "claude-sonnet-4-5"), + }); + + await expect(harness.waitForShutdown()).rejects.toMatchObject({ + code: "invalid_state", + message: "Shutdown has not been requested", + }); + }); + it("shuts down active work permanently and idempotently", async () => { const registration = newFaux(); const entered = deferred(); diff --git a/packages/client/README.md b/packages/client/README.md index e251429bb4b..1e2221404a8 100644 --- a/packages/client/README.md +++ b/packages/client/README.md @@ -34,3 +34,5 @@ Call `handlers.onData(chunk)` for inbound bytes, `handlers.onClose()` for an ord `PiClientOptions.maxFrameLength` bounds inbound and outbound CBOR payloads. Configure matching limits on the client and server. Transports should separately bound queued outbound bytes and preserve send order. Treat peers as untrusted. Use a secure transport where required and protect the protocol bearer token. + +Subscriber exceptions are isolated from protocol state. Set `onListenerError` in `PiClientOptions` to report them to application logging or diagnostics. diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index d79d9741ca0..88abfa75eeb 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -1,6 +1,5 @@ import { type Command, - type CommandResult, DEFAULT_MAX_FRAME_LENGTH, encodeClientMessage, PROTOCOL_VERSION, @@ -10,294 +9,288 @@ import { type ServerMessage, ServerMessageDecoder, type ServerSnapshot, - type SessionSnapshot, type SessionSummary, } from "@earendil-works/pi-protocol"; import { PiDisconnectedError, PiError, PiSessionDetachedError, toDisconnectedError, toError } from "./errors.ts"; import { notifyListeners } from "./listeners.ts"; -import { PiSessionClient } from "./session-client.ts"; +import { PendingRequests } from "./pending-requests.ts"; +import { PiSessionClient, type SessionClientOperations } from "./session-client.ts"; +import { ClientStateStore } from "./state-store.ts"; import type { ByteTransport, ByteTransportHandlers } from "./transport.ts"; import type { ConnectionState, ConnectionStateChange, CreateSessionOptions, - PendingRequest, PiClientOptions, Unsubscribe, } from "./types.ts"; const MAX_UINT32 = 0xffff_ffff; + +type Connection = + | { state: "disconnected" } + | { + state: "connecting"; + id: number; + decoder: ServerMessageDecoder; + deferred: PromiseWithResolvers; + transport?: ByteTransport; + } + | { state: "connected"; id: number; decoder: ServerMessageDecoder; transport: ByteTransport }; + export class PiClient { - private readonly options: PiClientOptions; - private readonly maxFrameLength: number; - private transport: ByteTransport | undefined; - private decoder: ServerMessageDecoder | undefined; - private connectionSequence = 0; - private stateValue: ConnectionState = "disconnected"; - private snapshotValue: ServerSnapshot | undefined; - private readonly sessionSnapshots = new Map(); - private readonly attachedSessionIds = new Set(); - private readonly sessionHandles = new Map(); - private readonly snapshotListeners = new Set<(snapshot: ServerSnapshot) => void>(); - private readonly eventListeners = new Set<(event: ServerEvent) => void>(); - private readonly stateListeners = new Set<(change: ConnectionStateChange) => void>(); - private readonly sessionSnapshotListeners = new Map void>>(); - private readonly sessionEventListeners = new Map void>>(); - private readonly pendingRequests = new Map(); - private requestSequence = 0; - private connectResolve: ((snapshot: ServerSnapshot) => void) | undefined; - private connectReject: ((error: Error) => void) | undefined; + readonly #options: PiClientOptions; + readonly #maxFrameLength: number; + readonly #stateStore: ClientStateStore; + readonly #pendingRequests = new PendingRequests(); + readonly #sessionHandles = new Map(); + readonly #stateListeners = new Set<(change: ConnectionStateChange) => void>(); + #connection: Connection = { state: "disconnected" }; + #connectionSequence = 0; constructor(options: PiClientOptions) { - this.options = options; - this.maxFrameLength = options.maxFrameLength ?? DEFAULT_MAX_FRAME_LENGTH; - if (!Number.isSafeInteger(this.maxFrameLength) || this.maxFrameLength <= 0 || this.maxFrameLength > MAX_UINT32) { + this.#options = options; + this.#maxFrameLength = options.maxFrameLength ?? DEFAULT_MAX_FRAME_LENGTH; + if ( + !Number.isSafeInteger(this.#maxFrameLength) || + this.#maxFrameLength <= 0 || + this.#maxFrameLength > MAX_UINT32 + ) { throw new TypeError(`PiClient maxFrameLength must be an integer between 1 and ${MAX_UINT32}`); } + this.#stateStore = new ClientStateStore(options.onListenerError); } get connectionState(): ConnectionState { - return this.stateValue; + return this.#connection.state; } + get connected(): boolean { - return this.stateValue === "connected"; + return this.#connection.state === "connected"; } + get snapshot(): ServerSnapshot | undefined { - return this.snapshotValue; + return this.#stateStore.snapshot; } + get sessions(): readonly SessionSummary[] { - return this.snapshotValue?.sessions ?? []; + return this.#stateStore.sessions; } connect(): Promise { - if (this.stateValue !== "disconnected") { - return Promise.reject(new PiDisconnectedError(`PiClient is already ${this.stateValue}`)); + if (this.#connection.state !== "disconnected") { + return Promise.reject(new PiDisconnectedError(`PiClient is already ${this.#connection.state}`)); } - this.setConnectionState("connecting"); - this.snapshotValue = undefined; - this.sessionSnapshots.clear(); - this.attachedSessionIds.clear(); - this.decoder = new ServerMessageDecoder({ maxFrameLength: this.maxFrameLength }); - const connectionId = ++this.connectionSequence; - const connected = new Promise((resolve, reject) => { - this.connectResolve = resolve; - this.connectReject = reject; - }); - const handlers: ByteTransportHandlers = { - onData: (chunk) => { - if (!this.isCurrentConnection(connectionId)) return; - if (!this.transport) { - this.protocolFailure( - new ProtocolValidationError("Received server data before the client hello was sent"), - ); - return; - } - this.handleChunk(chunk); - }, + this.#stateStore.reset(); + const id = ++this.#connectionSequence; + const deferred = Promise.withResolvers(); + this.#connection = { + state: "connecting", + id, + decoder: new ServerMessageDecoder({ maxFrameLength: this.#maxFrameLength }), + deferred, + }; + this.#notifyConnectionState(); + const handlers = { + onData: (chunk) => this.#handleTransportData(id, chunk), onClose: () => { - if (this.isCurrentConnection(connectionId)) this.handleTransportClose(); + if (this.#isCurrentConnection(id)) this.#handleTransportClose(); }, onError: (error) => { - if (this.isCurrentConnection(connectionId)) this.handleTransportError(error); + if (this.#isCurrentConnection(id)) this.#handleTransportError(error); }, - }; - void this.openTransport(connectionId, handlers); - return connected; + } satisfies ByteTransportHandlers; + void this.#openTransport(id, handlers); + return deferred.promise; } reconnect(): Promise { return this.connect(); } + disconnect(reason = "Client disconnected"): void { - if (this.stateValue === "disconnected") return; - const transport = this.transport; - this.failConnection(new PiDisconnectedError(reason)); + if (this.#connection.state === "disconnected") return; + const transport = this.#connection.transport; + this.#failConnection(new PiDisconnectedError(reason)); transport?.close(); } + subscribe(listener: (snapshot: ServerSnapshot) => void): Unsubscribe { - this.snapshotListeners.add(listener); - return () => this.snapshotListeners.delete(listener); + return this.#stateStore.subscribe(listener); } + onEvent(listener: (event: ServerEvent) => void): Unsubscribe { - this.eventListeners.add(listener); - return () => this.eventListeners.delete(listener); + return this.#stateStore.onEvent(listener); } + onConnectionStateChange(listener: (change: ConnectionStateChange) => void): Unsubscribe { - this.stateListeners.add(listener); - return () => this.stateListeners.delete(listener); + this.#stateListeners.add(listener); + return () => this.#stateListeners.delete(listener); } + getSession(sessionId: string): PiSessionClient { - let handle = this.sessionHandles.get(sessionId); - if (!handle) { - handle = new PiSessionClient(this, sessionId); - this.sessionHandles.set(sessionId, handle); - } + let handle = this.#sessionHandles.get(sessionId); + if (handle) return handle; + const operations: SessionClientOperations = { + isAttached: () => this.#stateStore.isSessionAttached(sessionId), + getSnapshot: () => this.#stateStore.getSessionSnapshot(sessionId), + subscribe: (listener) => this.#stateStore.subscribeSession(sessionId, listener), + onEvent: (listener) => this.#stateStore.onSessionEvent(sessionId, listener), + request: (command) => { + this.#assertSessionAttached(sessionId); + return this.request(command); + }, + }; + handle = new PiSessionClient(sessionId, operations); + this.#sessionHandles.set(sessionId, handle); return handle; } - getSessionSnapshot(sessionId: string): SessionSnapshot | undefined { - return this.sessionSnapshots.get(sessionId); - } - isSessionAttached(sessionId: string): boolean { - return this.attachedSessionIds.has(sessionId); - } + async listSessions(): Promise { return (await this.request({ command: "list" })).sessions; } + async createSession(options: CreateSessionOptions = {}): Promise { const result = await this.request({ command: "create", ...options }); return this.getSession(result.session.id); } + async attachSession(sessionId: string): Promise { - const previous = this.sessionSnapshots.get(sessionId); - this.sessionSnapshots.delete(sessionId); + const previous = this.#stateStore.forgetSessionSnapshot(sessionId); try { await this.request({ command: "attach", sessionId }); return this.getSession(sessionId); } catch (error) { - if (previous && !this.sessionSnapshots.has(sessionId)) this.sessionSnapshots.set(sessionId, previous); + if (previous) this.#stateStore.restoreSessionSnapshot(previous); throw error; } } - async detachSession(sessionId: string): Promise { - await this.request({ command: "detach", sessionId }); - } request(command: TCommand): Promise> { - const transport = this.transport; - if (this.stateValue !== "connected" || !transport) return Promise.reject(new PiDisconnectedError()); - const id = `request-${++this.requestSequence}`; + const connection = this.#connection; + if (connection.state !== "connected") return Promise.reject(new PiDisconnectedError()); let frame: Uint8Array; + const pending = this.#pendingRequests.create(command); try { frame = encodeClientMessage( - { type: "request", id, request: command }, - { maxFrameLength: this.maxFrameLength }, + { type: "request", id: pending.id, request: command }, + { maxFrameLength: this.#maxFrameLength }, ); } catch (error) { - return Promise.reject(toError(error)); - } - const promise = new Promise((resolve, reject) => { - this.pendingRequests.set(id, { command, resolve, reject }); - }); - this.sendFrame(transport, frame); - return promise as Promise>; - } - - subscribeSession(sessionId: string, listener: (snapshot: SessionSnapshot) => void): Unsubscribe { - let listeners = this.sessionSnapshotListeners.get(sessionId); - if (!listeners) { - listeners = new Set(); - this.sessionSnapshotListeners.set(sessionId, listeners); - } - listeners.add(listener); - return () => { - listeners.delete(listener); - if (listeners.size === 0) this.sessionSnapshotListeners.delete(sessionId); - }; - } - onSessionEvent(sessionId: string, listener: (event: ServerEvent) => void): Unsubscribe { - let listeners = this.sessionEventListeners.get(sessionId); - if (!listeners) { - listeners = new Set(); - this.sessionEventListeners.set(sessionId, listeners); + this.#pendingRequests.take(pending.id)?.reject(toError(error)); + return pending.promise; } - listeners.add(listener); - return () => { - listeners.delete(listener); - if (listeners.size === 0) this.sessionEventListeners.delete(sessionId); - }; - } - assertAttached(sessionId: string): void { - if (this.stateValue !== "connected") throw new PiDisconnectedError(); - if (!this.attachedSessionIds.has(sessionId)) throw new PiSessionDetachedError(sessionId); + this.#sendFrame(connection.transport, frame); + return pending.promise; } - private async openTransport(connectionId: number, handlers: ByteTransportHandlers): Promise { + async #openTransport(connectionId: number, handlers: ByteTransportHandlers): Promise { let transport: ByteTransport; try { - transport = await this.options.transportFactory(handlers); + transport = await this.#options.transportFactory(handlers); } catch (error) { - if (this.isCurrentConnection(connectionId)) this.failConnection(toDisconnectedError(error)); + if (this.#isCurrentConnection(connectionId)) this.#failConnection(toDisconnectedError(error)); return; } - if (!this.isCurrentConnection(connectionId)) { + const connection = this.#connection; + if (connection.state !== "connecting" || connection.id !== connectionId) { transport.close(); return; } - this.transport = transport; + this.#connection = { ...connection, transport }; try { await transport.send( encodeClientMessage( - { type: "hello", version: PROTOCOL_VERSION, token: this.options.token }, - { maxFrameLength: this.maxFrameLength }, + { type: "hello", version: PROTOCOL_VERSION, token: this.#options.token }, + { maxFrameLength: this.#maxFrameLength }, ), ); } catch (error) { - if (this.isCurrentConnection(connectionId)) { - this.failConnection(toDisconnectedError(error)); + if (this.#isCurrentConnection(connectionId)) { + this.#failConnection(toDisconnectedError(error)); transport.close(); } } } - private sendFrame(transport: ByteTransport, frame: Uint8Array): void { + + #sendFrame(transport: ByteTransport, frame: Uint8Array): void { let sending: Promise; try { sending = transport.send(frame); } catch (error) { - this.handleTransportError(toError(error)); + this.#handleTransportError(toError(error)); return; } void sending.catch((error: unknown) => { - if (this.transport === transport) this.handleTransportError(toError(error)); + const connection = this.#connection; + if (connection.state !== "disconnected" && connection.transport === transport) { + this.#handleTransportError(toError(error)); + } }); } - private handleChunk(chunk: Uint8Array): void { + + #handleTransportData(connectionId: number, chunk: Uint8Array): void { + const connection = this.#connection; + if (connection.state === "disconnected" || connection.id !== connectionId) return; + if (connection.state === "connecting" && !connection.transport) { + this.#protocolFailure(new ProtocolValidationError("Received server data before the client hello was sent")); + return; + } let messages: ServerMessage[]; try { - messages = this.decoder?.push(chunk) ?? []; + messages = connection.decoder.push(chunk); } catch (error) { - this.protocolFailure(toError(error)); + this.#protocolFailure(toError(error)); return; } for (const message of messages) { - if (this.stateValue === "disconnected") return; - this.handleMessage(message); + if (this.#connection.state === "disconnected") return; + this.#handleMessage(message); } } - private handleMessage(message: ServerMessage): void { - if (this.stateValue === "connecting") { + + #handleMessage(message: ServerMessage): void { + const connection = this.#connection; + if (connection.state === "connecting") { if (message.type === "hello_error") { - const transport = this.transport; - this.failConnection(new PiError(message.error)); - transport?.close(); + this.#failAndClose(new PiError(message.error)); return; } if (message.type !== "hello") { - this.protocolFailure(new ProtocolValidationError("Expected server hello as first message")); + this.#protocolFailure(new ProtocolValidationError("Expected server hello as first message")); + return; + } + if (!connection.transport) { + this.#protocolFailure( + new ProtocolValidationError("Received server hello before the client hello was sent"), + ); return; } - this.setConnectionState("connected"); - this.applyServerSnapshot(message.snapshot); - const resolve = this.connectResolve; - this.connectResolve = undefined; - this.connectReject = undefined; - resolve?.(message.snapshot); + this.#connection = { + state: "connected", + id: connection.id, + decoder: connection.decoder, + transport: connection.transport, + }; + this.#stateStore.applyServerSnapshot(message.snapshot); + this.#notifyConnectionState(); + connection.deferred.resolve(message.snapshot); return; } - if (this.stateValue !== "connected") return; + if (connection.state !== "connected") return; if (message.type === "hello" || message.type === "hello_error") { - this.protocolFailure(new ProtocolValidationError("Unexpected handshake message")); + this.#protocolFailure(new ProtocolValidationError("Unexpected handshake message")); return; } if (message.type === "event") { - this.applyEvent(message.event); + this.#stateStore.applyEvent(message.event); return; } - const pending = this.pendingRequests.get(message.id); + const pending = this.#pendingRequests.take(message.id); if (!pending) { - this.protocolFailure(new ProtocolValidationError("Response has no matching request")); + this.#protocolFailure(new ProtocolValidationError("Response has no matching request")); return; } - this.pendingRequests.delete(message.id); if (!message.ok) { pending.reject(new PiError(message.error)); return; @@ -307,92 +300,61 @@ export class PiClient { `Response command ${message.result.command} does not match ${pending.command.command}`, ); pending.reject(error); - this.protocolFailure(error); + this.#protocolFailure(error); return; } - this.applyResult(message.result); + this.#stateStore.applyResult(message.result); pending.resolve(message.result); } - private applyResult(result: CommandResult): void { - if (result.command === "list") return; - if (result.command === "detach") { - this.attachedSessionIds.delete(result.sessionId); - const snapshot = this.sessionSnapshots.get(result.sessionId); - if (snapshot) this.applySessionSnapshot({ ...snapshot, attached: false }, true); - return; - } - this.applySessionSnapshot(result.session); - } - private applyEvent(event: ServerEvent): void { - if (event.type === "server_snapshot") this.applyServerSnapshot(event.snapshot); - if (event.type === "session_snapshot") this.applySessionSnapshot(event.snapshot); - if (event.type === "session_removed") { - this.sessionSnapshots.delete(event.sessionId); - this.attachedSessionIds.delete(event.sessionId); - } - notifyListeners(this.eventListeners, event); - const sessionId = getEventSessionId(event); - if (sessionId) notifyListeners(this.sessionEventListeners.get(sessionId), event); - } - private applyServerSnapshot(snapshot: ServerSnapshot): void { - if (this.snapshotValue && snapshot.revision < this.snapshotValue.revision) return; - this.snapshotValue = snapshot; - this.attachedSessionIds.clear(); - for (const session of snapshot.sessions) if (session.attached) this.attachedSessionIds.add(session.id); - notifyListeners(this.snapshotListeners, snapshot); - } - private applySessionSnapshot(snapshot: SessionSnapshot, force = false): void { - const current = this.sessionSnapshots.get(snapshot.id); - if (!force && current && snapshot.revision < current.revision) return; - this.sessionSnapshots.set(snapshot.id, snapshot); - if (snapshot.attached) this.attachedSessionIds.add(snapshot.id); - else this.attachedSessionIds.delete(snapshot.id); - notifyListeners(this.sessionSnapshotListeners.get(snapshot.id), snapshot); - } - private handleTransportClose(): void { + + #handleTransportClose(): void { + const connection = this.#connection; + if (connection.state === "disconnected") return; let error: Error = new PiDisconnectedError("Byte transport closed"); try { - this.decoder?.end(); + connection.decoder.end(); } catch (decoderError) { error = toError(decoderError); } - this.failConnection(error); + this.#failConnection(error); } - private handleTransportError(error: Error): void { - const transport = this.transport; - this.failConnection(toDisconnectedError(error)); - transport?.close(); + + #handleTransportError(error: Error): void { + this.#failAndClose(toDisconnectedError(error)); + } + + #protocolFailure(error: Error): void { + this.#failAndClose(error); } - private protocolFailure(error: Error): void { - const transport = this.transport; - this.failConnection(error); + + #failAndClose(error: Error): void { + const connection = this.#connection; + const transport = connection.state === "disconnected" ? undefined : connection.transport; + this.#failConnection(error); transport?.close(); } - private failConnection(error: Error): void { - if (this.stateValue === "disconnected") return; - const reject = this.connectReject; - const pending = [...this.pendingRequests.values()]; - this.transport = undefined; - this.decoder = undefined; - this.connectResolve = undefined; - this.connectReject = undefined; - this.pendingRequests.clear(); - this.attachedSessionIds.clear(); - reject?.(error); - for (const request of pending) request.reject(error); - this.setConnectionState("disconnected", error); + + #failConnection(error: Error): void { + const connection = this.#connection; + if (connection.state === "disconnected") return; + this.#connection = { state: "disconnected" }; + this.#stateStore.clearAttachments(); + if (connection.state === "connecting") connection.deferred.reject(error); + this.#pendingRequests.rejectAll(error); + this.#notifyConnectionState(error); } - private isCurrentConnection(connectionId: number): boolean { - return connectionId === this.connectionSequence && this.stateValue !== "disconnected"; + + #assertSessionAttached(sessionId: string): void { + if (this.#connection.state !== "connected") throw new PiDisconnectedError(); + if (!this.#stateStore.isSessionAttached(sessionId)) throw new PiSessionDetachedError(sessionId); } - private setConnectionState(state: ConnectionState, error?: Error): void { - this.stateValue = state; - notifyListeners(this.stateListeners, error ? { state, error } : { state }); + + #isCurrentConnection(connectionId: number): boolean { + return this.#connection.state !== "disconnected" && this.#connection.id === connectionId; } -} -function getEventSessionId(event: ServerEvent): string | undefined { - if (event.type === "session_snapshot") return event.snapshot.id; - if (event.type === "session_progress" || event.type === "session_removed") return event.sessionId; - return undefined; + #notifyConnectionState(error?: Error): void { + const state = this.#connection.state; + notifyListeners(this.#stateListeners, error ? { state, error } : { state }, this.#options.onListenerError); + } } diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts index d7ddefdb5a8..0b1235b6020 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -6,6 +6,7 @@ export type { ConnectionState, ConnectionStateChange, CreateSessionOptions, + ListenerErrorHandler, PiClientOptions, Unsubscribe, } from "./types.ts"; diff --git a/packages/client/src/listeners.ts b/packages/client/src/listeners.ts index a94680c4c17..9a850b34eeb 100644 --- a/packages/client/src/listeners.ts +++ b/packages/client/src/listeners.ts @@ -1,9 +1,21 @@ -export function notifyListeners(listeners: Iterable<(value: T) => void> | undefined, value: T): void { +import { toError } from "./errors.ts"; +import type { ListenerErrorHandler } from "./types.ts"; + +export function notifyListeners( + listeners: Iterable<(value: T) => void> | undefined, + value: T, + onError?: ListenerErrorHandler, +): void { for (const listener of listeners ?? []) { try { listener(value); - } catch { - // Consumer callbacks cannot affect protocol or transport state. + } catch (error) { + if (!onError) continue; + try { + onError(toError(error)); + } catch { + // Diagnostics cannot affect protocol or transport state. + } } } } diff --git a/packages/client/src/pending-requests.ts b/packages/client/src/pending-requests.ts new file mode 100644 index 00000000000..6fbe1fc0292 --- /dev/null +++ b/packages/client/src/pending-requests.ts @@ -0,0 +1,33 @@ +import type { Command, CommandResult, ResultForCommand } from "@earendil-works/pi-protocol"; + +interface PendingRequest { + command: Command; + resolve(result: CommandResult): void; + reject(error: Error): void; +} + +export class PendingRequests { + readonly #requests = new Map(); + #sequence = 0; + + create( + command: TCommand, + ): { id: string; promise: Promise> } { + const id = `request-${++this.#sequence}`; + const { promise, resolve, reject } = Promise.withResolvers(); + this.#requests.set(id, { command, resolve, reject }); + return { id, promise: promise as Promise> }; + } + + take(id: string): PendingRequest | undefined { + const request = this.#requests.get(id); + if (request) this.#requests.delete(id); + return request; + } + + rejectAll(error: Error): void { + const requests = [...this.#requests.values()]; + this.#requests.clear(); + for (const request of requests) request.reject(error); + } +} diff --git a/packages/client/src/session-client.ts b/packages/client/src/session-client.ts index dccaccb97fd..a5c1dba920c 100644 --- a/packages/client/src/session-client.ts +++ b/packages/client/src/session-client.ts @@ -1,63 +1,74 @@ -import type { ModelRef, ServerEvent, SessionSnapshot, ThinkingLevel } from "@earendil-works/pi-protocol"; -import type { SessionClientHost, Unsubscribe } from "./types.ts"; +import type { + Command, + ModelRef, + ResultForCommand, + ServerEvent, + SessionSnapshot, + ThinkingLevel, +} from "@earendil-works/pi-protocol"; +import type { Unsubscribe } from "./types.ts"; + +type SessionCommand = Extract; + +export interface SessionClientOperations { + isAttached(): boolean; + getSnapshot(): SessionSnapshot | undefined; + subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe; + onEvent(listener: (event: ServerEvent) => void): Unsubscribe; + request(command: TCommand): Promise>; +} export class PiSessionClient { readonly id: string; - private readonly client: SessionClientHost; + readonly #operations: SessionClientOperations; - constructor(client: SessionClientHost, id: string) { - this.client = client; + /** @internal Construct session handles through `PiClient`. */ + constructor(id: string, operations: SessionClientOperations) { this.id = id; + this.#operations = operations; } get attached(): boolean { - return this.client.isSessionAttached(this.id); + return this.#operations.isAttached(); } get snapshot(): SessionSnapshot | undefined { - return this.client.getSessionSnapshot(this.id); + return this.#operations.getSnapshot(); } subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe { - return this.client.subscribeSession(this.id, listener); + return this.#operations.subscribe(listener); } onEvent(listener: (event: ServerEvent) => void): Unsubscribe { - return this.client.onSessionEvent(this.id, listener); + return this.#operations.onEvent(listener); } async detach(): Promise { - this.client.assertAttached(this.id); - await this.client.detachSession(this.id); + await this.#request({ command: "detach", sessionId: this.id }); } async prompt(text: string): Promise { - this.client.assertAttached(this.id); - const result = await this.client.request({ command: "prompt", sessionId: this.id, text }); - return result.session; + return (await this.#request({ command: "prompt", sessionId: this.id, text })).session; } async steer(text: string): Promise { - this.client.assertAttached(this.id); - const result = await this.client.request({ command: "steer", sessionId: this.id, text }); - return result.session; + return (await this.#request({ command: "steer", sessionId: this.id, text })).session; } async abort(): Promise { - this.client.assertAttached(this.id); - const result = await this.client.request({ command: "abort", sessionId: this.id }); - return result.session; + return (await this.#request({ command: "abort", sessionId: this.id })).session; } async setModel(model: ModelRef): Promise { - this.client.assertAttached(this.id); - const result = await this.client.request({ command: "set_model", sessionId: this.id, model }); - return result.session; + return (await this.#request({ command: "set_model", sessionId: this.id, model })).session; } async setThinking(thinkingLevel: ThinkingLevel): Promise { - this.client.assertAttached(this.id); - const result = await this.client.request({ command: "set_thinking", sessionId: this.id, thinkingLevel }); - return result.session; + return (await this.#request({ command: "set_thinking", sessionId: this.id, thinkingLevel })).session; + } + + #request(command: TCommand): Promise> { + return this.#operations.request(command); } } diff --git a/packages/client/src/state-store.ts b/packages/client/src/state-store.ts new file mode 100644 index 00000000000..097dece20f6 --- /dev/null +++ b/packages/client/src/state-store.ts @@ -0,0 +1,141 @@ +import type { + CommandResult, + ServerEvent, + ServerSnapshot, + SessionSnapshot, + SessionSummary, +} from "@earendil-works/pi-protocol"; +import { notifyListeners } from "./listeners.ts"; +import type { ListenerErrorHandler, Unsubscribe } from "./types.ts"; + +export class ClientStateStore { + readonly #sessionSnapshots = new Map(); + readonly #attachedSessionIds = new Set(); + readonly #snapshotListeners = new Set<(snapshot: ServerSnapshot) => void>(); + readonly #eventListeners = new Set<(event: ServerEvent) => void>(); + readonly #sessionSnapshotListeners = new Map void>>(); + readonly #sessionEventListeners = new Map void>>(); + readonly #onListenerError: ListenerErrorHandler | undefined; + #snapshot: ServerSnapshot | undefined; + + constructor(onListenerError?: ListenerErrorHandler) { + this.#onListenerError = onListenerError; + } + + get snapshot(): ServerSnapshot | undefined { + return this.#snapshot; + } + + get sessions(): readonly SessionSummary[] { + return this.#snapshot?.sessions ?? []; + } + + reset(): void { + this.#snapshot = undefined; + this.#sessionSnapshots.clear(); + this.#attachedSessionIds.clear(); + } + + clearAttachments(): void { + this.#attachedSessionIds.clear(); + } + + getSessionSnapshot(sessionId: string): SessionSnapshot | undefined { + return this.#sessionSnapshots.get(sessionId); + } + + isSessionAttached(sessionId: string): boolean { + return this.#attachedSessionIds.has(sessionId); + } + + forgetSessionSnapshot(sessionId: string): SessionSnapshot | undefined { + const previous = this.#sessionSnapshots.get(sessionId); + this.#sessionSnapshots.delete(sessionId); + return previous; + } + + restoreSessionSnapshot(snapshot: SessionSnapshot): void { + if (!this.#sessionSnapshots.has(snapshot.id)) this.#sessionSnapshots.set(snapshot.id, snapshot); + } + + subscribe(listener: (snapshot: ServerSnapshot) => void): Unsubscribe { + this.#snapshotListeners.add(listener); + return () => this.#snapshotListeners.delete(listener); + } + + onEvent(listener: (event: ServerEvent) => void): Unsubscribe { + this.#eventListeners.add(listener); + return () => this.#eventListeners.delete(listener); + } + + subscribeSession(sessionId: string, listener: (snapshot: SessionSnapshot) => void): Unsubscribe { + return addMappedListener(this.#sessionSnapshotListeners, sessionId, listener); + } + + onSessionEvent(sessionId: string, listener: (event: ServerEvent) => void): Unsubscribe { + return addMappedListener(this.#sessionEventListeners, sessionId, listener); + } + + applyResult(result: CommandResult): void { + if (result.command === "list") return; + if (result.command === "detach") { + this.#attachedSessionIds.delete(result.sessionId); + const snapshot = this.#sessionSnapshots.get(result.sessionId); + if (snapshot) this.#applySessionSnapshot({ ...snapshot, attached: false }, true); + return; + } + this.#applySessionSnapshot(result.session); + } + + applyEvent(event: ServerEvent): void { + if (event.type === "server_snapshot") this.applyServerSnapshot(event.snapshot); + if (event.type === "session_snapshot") this.#applySessionSnapshot(event.snapshot); + if (event.type === "session_removed") { + this.#sessionSnapshots.delete(event.sessionId); + this.#attachedSessionIds.delete(event.sessionId); + } + notifyListeners(this.#eventListeners, event, this.#onListenerError); + const sessionId = getEventSessionId(event); + if (sessionId) notifyListeners(this.#sessionEventListeners.get(sessionId), event, this.#onListenerError); + } + + applyServerSnapshot(snapshot: ServerSnapshot): void { + if (this.#snapshot && snapshot.revision < this.#snapshot.revision) return; + this.#snapshot = snapshot; + this.#attachedSessionIds.clear(); + for (const session of snapshot.sessions) if (session.attached) this.#attachedSessionIds.add(session.id); + notifyListeners(this.#snapshotListeners, snapshot, this.#onListenerError); + } + + #applySessionSnapshot(snapshot: SessionSnapshot, force = false): void { + const current = this.#sessionSnapshots.get(snapshot.id); + if (!force && current && snapshot.revision < current.revision) return; + this.#sessionSnapshots.set(snapshot.id, snapshot); + if (snapshot.attached) this.#attachedSessionIds.add(snapshot.id); + else this.#attachedSessionIds.delete(snapshot.id); + notifyListeners(this.#sessionSnapshotListeners.get(snapshot.id), snapshot, this.#onListenerError); + } +} + +function addMappedListener( + listenersById: Map void>>, + id: string, + listener: (value: T) => void, +): Unsubscribe { + let listeners = listenersById.get(id); + if (!listeners) { + listeners = new Set(); + listenersById.set(id, listeners); + } + listeners.add(listener); + return () => { + listeners.delete(listener); + if (listeners.size === 0) listenersById.delete(id); + }; +} + +function getEventSessionId(event: ServerEvent): string | undefined { + if (event.type === "session_snapshot") return event.snapshot.id; + if (event.type === "session_progress" || event.type === "session_removed") return event.sessionId; + return undefined; +} diff --git a/packages/client/src/types.ts b/packages/client/src/types.ts index aba7b496687..0a5cf6662b7 100644 --- a/packages/client/src/types.ts +++ b/packages/client/src/types.ts @@ -1,13 +1,4 @@ -import type { - Command, - CommandResult, - ModelRef, - ResultForCommand, - ServerEvent, - ServerSnapshot, - SessionSnapshot, - ThinkingLevel, -} from "@earendil-works/pi-protocol"; +import type { ModelRef, ThinkingLevel } from "@earendil-works/pi-protocol"; import type { ByteTransportFactory } from "./transport.ts"; export type ConnectionState = "disconnected" | "connecting" | "connected"; @@ -18,11 +9,14 @@ export interface ConnectionStateChange { } export type Unsubscribe = () => void; +export type ListenerErrorHandler = (error: Error) => void; export interface PiClientOptions { token: string; transportFactory: ByteTransportFactory; maxFrameLength?: number; + /** Reports subscriber failures without allowing them to corrupt client state. */ + onListenerError?: ListenerErrorHandler; } export interface CreateSessionOptions { @@ -31,21 +25,3 @@ export interface CreateSessionOptions { model?: ModelRef; thinkingLevel?: ThinkingLevel; } - -export interface SessionClientHost { - isSessionAttached(sessionId: string): boolean; - getSessionSnapshot(sessionId: string): SessionSnapshot | undefined; - detachSession(sessionId: string): Promise; - request(command: TCommand): Promise>; - subscribeSession(sessionId: string, listener: (snapshot: SessionSnapshot) => void): Unsubscribe; - onSessionEvent(sessionId: string, listener: (event: ServerEvent) => void): Unsubscribe; - assertAttached(sessionId: string): void; -} - -export interface PendingRequest { - command: Command; - resolve(result: CommandResult): void; - reject(error: Error): void; -} - -export type ServerSnapshotListener = (snapshot: ServerSnapshot) => void; diff --git a/packages/client/test/client-connection.test.ts b/packages/client/test/client-connection.test.ts index d76d129539b..9c842740c2b 100644 --- a/packages/client/test/client-connection.test.ts +++ b/packages/client/test/client-connection.test.ts @@ -96,6 +96,33 @@ describe("PiClient", () => { expect(client.connectionState).toBe("connected"); }); + test("reports subscriber failures without changing connection state", async () => { + const server = new MemoryByteServer(); + const listenerErrors: Error[] = []; + server.onMessage((message) => { + if (message.type === "hello") { + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }); + } + }); + const client = new PiClient({ + token: "bearer-secret", + transportFactory: (handlers) => server.connect(handlers), + onListenerError: (error) => listenerErrors.push(error), + }); + client.subscribe(() => { + throw new Error("consumer failure"); + }); + + await expect(client.connect()).resolves.toEqual(baseServerSnapshot); + expect(listenerErrors).toEqual([expect.objectContaining({ message: "consumer failure" })]); + expect(client.connectionState).toBe("connected"); + }); + test("rejects a typed handshake authentication error", async () => { const server = new MemoryByteServer(); server.onMessage(() => { diff --git a/tsconfig.base.json b/tsconfig.base.json index 57e97d6e361..2f338f05805 100644 --- a/tsconfig.base.json +++ b/tsconfig.base.json @@ -2,7 +2,7 @@ "compilerOptions": { "target": "ES2022", "module": "Node16", - "lib": ["ES2022"], + "lib": ["ES2024"], "strict": true, "erasableSyntaxOnly": true, "esModuleInterop": true, From ee3156715d3c7def800c6aaaa2bd664a2f315d38 Mon Sep 17 00:00:00 2001 From: Christian Klotz Date: Fri, 31 Jul 2026 11:17:51 +0300 Subject: [PATCH 3/5] refactor(client): simplify lifecycle architecture --- packages/client/CHANGELOG.md | 2 +- packages/client/README.md | 6 +- packages/client/src/client.ts | 437 ++++++++---------- packages/client/src/connection.ts | 239 ++++++++++ packages/client/src/errors.ts | 4 +- packages/client/src/index.ts | 4 +- packages/client/src/listeners.ts | 21 - packages/client/src/pending-requests.ts | 33 -- packages/client/src/session-client.ts | 74 --- .../client/src/{state-store.ts => state.ts} | 43 +- packages/client/test/client-state.test.ts | 200 -------- ...-connection.test.ts => connection.test.ts} | 237 +++++----- packages/client/test/requests.test.ts | 68 +++ packages/client/test/sessions.test.ts | 74 +++ packages/client/test/state.test.ts | 119 +++++ packages/client/test/support.ts | 20 +- 16 files changed, 860 insertions(+), 721 deletions(-) create mode 100644 packages/client/src/connection.ts delete mode 100644 packages/client/src/listeners.ts delete mode 100644 packages/client/src/pending-requests.ts delete mode 100644 packages/client/src/session-client.ts rename packages/client/src/{state-store.ts => state.ts} (83%) delete mode 100644 packages/client/test/client-state.test.ts rename packages/client/test/{client-connection.test.ts => connection.test.ts} (65%) create mode 100644 packages/client/test/requests.test.ts create mode 100644 packages/client/test/sessions.test.ts create mode 100644 packages/client/test/state.test.ts diff --git a/packages/client/CHANGELOG.md b/packages/client/CHANGELOG.md index b94aff2dd32..5dab7bdb813 100644 --- a/packages/client/CHANGELOG.md +++ b/packages/client/CHANGELOG.md @@ -4,4 +4,4 @@ ### Added -- Added the experimental transport-neutral `PiClient` and multi-session `PiSessionClient` APIs. +- Added the experimental transport-neutral `PiClient` and multi-session `PiSessionHandle` APIs with structured `PiServerError` responses. diff --git a/packages/client/README.md b/packages/client/README.md index 1e2221404a8..cc11595bfc8 100644 --- a/packages/client/README.md +++ b/packages/client/README.md @@ -25,9 +25,11 @@ unsubscribe(); Call `handlers.onData(chunk)` for inbound bytes, `handlers.onClose()` for an orderly terminal close, and `handlers.onError(error)` for transport failures. A factory must create a fresh transport for every connection attempt. -`PiClient` does not reconnect automatically. Call `reconnect()` after disconnection. One connection can attach several `PiSessionClient` handles. Requests are correlated by ID. Server snapshots and successful response snapshots are authoritative, while progress events do not mutate snapshot state optimistically. +`PiClient` does not reconnect automatically. Call `reconnect()` after disconnection. One connection can attach several sessions. Requests are correlated by ID. Server snapshots and successful response snapshots are authoritative, while progress events do not mutate snapshot state optimistically. Read cached session summaries from `client.snapshot?.sessions`; call `listSessions()` to request a refreshed list from the server. -`subscribe()` observes authoritative snapshots. `onEvent()` observes protocol events. Both return an unsubscribe function. A detached session handle remains readable, but commands throw `PiSessionDetachedError` until it is attached again. +`createSession()` and `attachSession()` return a `PiSessionHandle`; handles cannot be constructed directly. A returned handle is attached and remains a stable client-side reference for that session. Explicit detach, server removal, or disconnection makes a retained handle unavailable for commands. Its latest snapshot remains readable after detach or disconnection unless the server removes the session. Calling `attachSession()` again reacquires the session and returns the existing handle. Commands fail with `PiDisconnectedError` while the client is disconnected and `PiSessionDetachedError` when the client is connected but the session is detached. + +`subscribe()` observes authoritative snapshots. `onEvent()` observes protocol events. Both return an unsubscribe function. Structured errors returned by the server are exposed as `PiServerError`. ## Limits and security diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index 88abfa75eeb..7318d662896 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -1,22 +1,21 @@ import { type Command, - DEFAULT_MAX_FRAME_LENGTH, + type CommandResult, + type EventEnvelope, encodeClientMessage, - PROTOCOL_VERSION, + type ModelRef, ProtocolValidationError, + type ResponseEnvelope, type ResultForCommand, type ServerEvent, - type ServerMessage, - ServerMessageDecoder, type ServerSnapshot, + type SessionSnapshot, type SessionSummary, + type ThinkingLevel, } from "@earendil-works/pi-protocol"; -import { PiDisconnectedError, PiError, PiSessionDetachedError, toDisconnectedError, toError } from "./errors.ts"; -import { notifyListeners } from "./listeners.ts"; -import { PendingRequests } from "./pending-requests.ts"; -import { PiSessionClient, type SessionClientOperations } from "./session-client.ts"; -import { ClientStateStore } from "./state-store.ts"; -import type { ByteTransport, ByteTransportHandlers } from "./transport.ts"; +import { Connection } from "./connection.ts"; +import { PiDisconnectedError, PiServerError, PiSessionDetachedError, toError } from "./errors.ts"; +import { ClientState } from "./state.ts"; import type { ConnectionState, ConnectionStateChange, @@ -25,40 +24,110 @@ import type { Unsubscribe, } from "./types.ts"; -const MAX_UINT32 = 0xffff_ffff; +type SessionCommand = Extract; -type Connection = - | { state: "disconnected" } - | { - state: "connecting"; - id: number; - decoder: ServerMessageDecoder; - deferred: PromiseWithResolvers; - transport?: ByteTransport; - } - | { state: "connected"; id: number; decoder: ServerMessageDecoder; transport: ByteTransport }; +export interface PiSessionHandle { + readonly id: string; + readonly attached: boolean; + readonly snapshot: SessionSnapshot | undefined; + subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe; + onEvent(listener: (event: ServerEvent) => void): Unsubscribe; + detach(): Promise; + prompt(text: string): Promise; + steer(text: string): Promise; + abort(): Promise; + setModel(model: ModelRef): Promise; + setThinking(thinkingLevel: ThinkingLevel): Promise; +} + +interface SessionHandleCallbacks { + isAttached(): boolean; + getSnapshot(): SessionSnapshot | undefined; + subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe; + onEvent(listener: (event: ServerEvent) => void): Unsubscribe; + request(command: TCommand): Promise>; +} + +class SessionHandle implements PiSessionHandle { + readonly id: string; + readonly #callbacks: SessionHandleCallbacks; + + constructor(id: string, callbacks: SessionHandleCallbacks) { + this.id = id; + this.#callbacks = callbacks; + } + + get attached(): boolean { + return this.#callbacks.isAttached(); + } + + get snapshot(): SessionSnapshot | undefined { + return this.#callbacks.getSnapshot(); + } + + subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe { + return this.#callbacks.subscribe(listener); + } + + onEvent(listener: (event: ServerEvent) => void): Unsubscribe { + return this.#callbacks.onEvent(listener); + } + + async detach(): Promise { + await this.#request({ command: "detach", sessionId: this.id }); + } + + async prompt(text: string): Promise { + return (await this.#request({ command: "prompt", sessionId: this.id, text })).session; + } + + async steer(text: string): Promise { + return (await this.#request({ command: "steer", sessionId: this.id, text })).session; + } + + async abort(): Promise { + return (await this.#request({ command: "abort", sessionId: this.id })).session; + } + + async setModel(model: ModelRef): Promise { + return (await this.#request({ command: "set_model", sessionId: this.id, model })).session; + } + + async setThinking(thinkingLevel: ThinkingLevel): Promise { + return (await this.#request({ command: "set_thinking", sessionId: this.id, thinkingLevel })).session; + } + + #request(command: TCommand): Promise> { + return this.#callbacks.request(command); + } +} + +interface PendingRequest { + command: Command; + resolve(result: CommandResult): void; + reject(error: Error): void; +} export class PiClient { readonly #options: PiClientOptions; - readonly #maxFrameLength: number; - readonly #stateStore: ClientStateStore; - readonly #pendingRequests = new PendingRequests(); - readonly #sessionHandles = new Map(); - readonly #stateListeners = new Set<(change: ConnectionStateChange) => void>(); - #connection: Connection = { state: "disconnected" }; - #connectionSequence = 0; + readonly #connection: Connection; + readonly #state: ClientState; + readonly #pendingRequests = new Map(); + readonly #sessions = new Map(); + readonly #connectionStateListeners = new Set<(change: ConnectionStateChange) => void>(); + #requestSequence = 0; constructor(options: PiClientOptions) { this.#options = options; - this.#maxFrameLength = options.maxFrameLength ?? DEFAULT_MAX_FRAME_LENGTH; - if ( - !Number.isSafeInteger(this.#maxFrameLength) || - this.#maxFrameLength <= 0 || - this.#maxFrameLength > MAX_UINT32 - ) { - throw new TypeError(`PiClient maxFrameLength must be an integer between 1 and ${MAX_UINT32}`); - } - this.#stateStore = new ClientStateStore(options.onListenerError); + this.#state = new ClientState(options.onListenerError); + this.#connection = new Connection({ + token: options.token, + transportFactory: options.transportFactory, + maxFrameLength: options.maxFrameLength, + onHandshake: (snapshot) => this.#state.applyServerSnapshot(snapshot), + onMessage: (message) => this.#handleMessage(message), + onStateChange: (change) => this.#handleConnectionStateChange(change), + }); } get connectionState(): ConnectionState { @@ -70,38 +139,12 @@ export class PiClient { } get snapshot(): ServerSnapshot | undefined { - return this.#stateStore.snapshot; - } - - get sessions(): readonly SessionSummary[] { - return this.#stateStore.sessions; + return this.#state.snapshot; } connect(): Promise { - if (this.#connection.state !== "disconnected") { - return Promise.reject(new PiDisconnectedError(`PiClient is already ${this.#connection.state}`)); - } - this.#stateStore.reset(); - const id = ++this.#connectionSequence; - const deferred = Promise.withResolvers(); - this.#connection = { - state: "connecting", - id, - decoder: new ServerMessageDecoder({ maxFrameLength: this.#maxFrameLength }), - deferred, - }; - this.#notifyConnectionState(); - const handlers = { - onData: (chunk) => this.#handleTransportData(id, chunk), - onClose: () => { - if (this.#isCurrentConnection(id)) this.#handleTransportClose(); - }, - onError: (error) => { - if (this.#isCurrentConnection(id)) this.#handleTransportError(error); - }, - } satisfies ByteTransportHandlers; - void this.#openTransport(id, handlers); - return deferred.promise; + if (this.#connection.state === "disconnected") this.#state.reset(); + return this.#connection.connect(); } reconnect(): Promise { @@ -109,190 +152,91 @@ export class PiClient { } disconnect(reason = "Client disconnected"): void { - if (this.#connection.state === "disconnected") return; - const transport = this.#connection.transport; - this.#failConnection(new PiDisconnectedError(reason)); - transport?.close(); + this.#connection.disconnect(reason); } subscribe(listener: (snapshot: ServerSnapshot) => void): Unsubscribe { - return this.#stateStore.subscribe(listener); + return this.#state.subscribe(listener); } onEvent(listener: (event: ServerEvent) => void): Unsubscribe { - return this.#stateStore.onEvent(listener); + return this.#state.onEvent(listener); } onConnectionStateChange(listener: (change: ConnectionStateChange) => void): Unsubscribe { - this.#stateListeners.add(listener); - return () => this.#stateListeners.delete(listener); - } - - getSession(sessionId: string): PiSessionClient { - let handle = this.#sessionHandles.get(sessionId); - if (handle) return handle; - const operations: SessionClientOperations = { - isAttached: () => this.#stateStore.isSessionAttached(sessionId), - getSnapshot: () => this.#stateStore.getSessionSnapshot(sessionId), - subscribe: (listener) => this.#stateStore.subscribeSession(sessionId, listener), - onEvent: (listener) => this.#stateStore.onSessionEvent(sessionId, listener), - request: (command) => { - this.#assertSessionAttached(sessionId); - return this.request(command); - }, - }; - handle = new PiSessionClient(sessionId, operations); - this.#sessionHandles.set(sessionId, handle); - return handle; + this.#connectionStateListeners.add(listener); + return () => this.#connectionStateListeners.delete(listener); } async listSessions(): Promise { - return (await this.request({ command: "list" })).sessions; + return (await this.#request({ command: "list" })).sessions; } - async createSession(options: CreateSessionOptions = {}): Promise { - const result = await this.request({ command: "create", ...options }); - return this.getSession(result.session.id); + async createSession(options: CreateSessionOptions = {}): Promise { + const result = await this.#request({ command: "create", ...options }); + return this.#getOrCreateSessionHandle(result.session.id); } - async attachSession(sessionId: string): Promise { - const previous = this.#stateStore.forgetSessionSnapshot(sessionId); + async attachSession(sessionId: string): Promise { + const previous = this.#state.forgetSessionSnapshot(sessionId); try { - await this.request({ command: "attach", sessionId }); - return this.getSession(sessionId); + await this.#request({ command: "attach", sessionId }); + return this.#getOrCreateSessionHandle(sessionId); } catch (error) { - if (previous) this.#stateStore.restoreSessionSnapshot(previous); + if (previous) this.#state.restoreSessionSnapshot(previous); throw error; } } - request(command: TCommand): Promise> { - const connection = this.#connection; - if (connection.state !== "connected") return Promise.reject(new PiDisconnectedError()); + #request(command: TCommand): Promise> { + if (!this.connected) return Promise.reject(new PiDisconnectedError()); + const id = `request-${++this.#requestSequence}`; + const { promise, resolve, reject } = Promise.withResolvers(); + this.#pendingRequests.set(id, { command, resolve, reject }); let frame: Uint8Array; - const pending = this.#pendingRequests.create(command); try { frame = encodeClientMessage( - { type: "request", id: pending.id, request: command }, - { maxFrameLength: this.#maxFrameLength }, - ); - } catch (error) { - this.#pendingRequests.take(pending.id)?.reject(toError(error)); - return pending.promise; - } - this.#sendFrame(connection.transport, frame); - return pending.promise; - } - - async #openTransport(connectionId: number, handlers: ByteTransportHandlers): Promise { - let transport: ByteTransport; - try { - transport = await this.#options.transportFactory(handlers); - } catch (error) { - if (this.#isCurrentConnection(connectionId)) this.#failConnection(toDisconnectedError(error)); - return; - } - const connection = this.#connection; - if (connection.state !== "connecting" || connection.id !== connectionId) { - transport.close(); - return; - } - this.#connection = { ...connection, transport }; - try { - await transport.send( - encodeClientMessage( - { type: "hello", version: PROTOCOL_VERSION, token: this.#options.token }, - { maxFrameLength: this.#maxFrameLength }, - ), + { type: "request", id, request: command }, + { maxFrameLength: this.#connection.maxFrameLength }, ); } catch (error) { - if (this.#isCurrentConnection(connectionId)) { - this.#failConnection(toDisconnectedError(error)); - transport.close(); - } + this.#takePendingRequest(id)?.reject(toError(error)); + return promise as Promise>; } + this.#connection.send(frame); + return promise as Promise>; } - #sendFrame(transport: ByteTransport, frame: Uint8Array): void { - let sending: Promise; - try { - sending = transport.send(frame); - } catch (error) { - this.#handleTransportError(toError(error)); - return; - } - void sending.catch((error: unknown) => { - const connection = this.#connection; - if (connection.state !== "disconnected" && connection.transport === transport) { - this.#handleTransportError(toError(error)); - } - }); - } - - #handleTransportData(connectionId: number, chunk: Uint8Array): void { - const connection = this.#connection; - if (connection.state === "disconnected" || connection.id !== connectionId) return; - if (connection.state === "connecting" && !connection.transport) { - this.#protocolFailure(new ProtocolValidationError("Received server data before the client hello was sent")); - return; - } - let messages: ServerMessage[]; - try { - messages = connection.decoder.push(chunk); - } catch (error) { - this.#protocolFailure(toError(error)); - return; - } - for (const message of messages) { - if (this.#connection.state === "disconnected") return; - this.#handleMessage(message); - } + #getOrCreateSessionHandle(sessionId: string): PiSessionHandle { + let handle = this.#sessions.get(sessionId); + if (handle) return handle; + const callbacks: SessionHandleCallbacks = { + isAttached: () => this.#state.isSessionAttached(sessionId), + getSnapshot: () => this.#state.getSessionSnapshot(sessionId), + subscribe: (listener) => this.#state.subscribeSession(sessionId, listener), + onEvent: (listener) => this.#state.onSessionEvent(sessionId, listener), + request: (command) => { + this.#assertSessionAttached(sessionId); + return this.#request(command); + }, + }; + handle = new SessionHandle(sessionId, callbacks); + this.#sessions.set(sessionId, handle); + return handle; } - #handleMessage(message: ServerMessage): void { - const connection = this.#connection; - if (connection.state === "connecting") { - if (message.type === "hello_error") { - this.#failAndClose(new PiError(message.error)); - return; - } - if (message.type !== "hello") { - this.#protocolFailure(new ProtocolValidationError("Expected server hello as first message")); - return; - } - if (!connection.transport) { - this.#protocolFailure( - new ProtocolValidationError("Received server hello before the client hello was sent"), - ); - return; - } - this.#connection = { - state: "connected", - id: connection.id, - decoder: connection.decoder, - transport: connection.transport, - }; - this.#stateStore.applyServerSnapshot(message.snapshot); - this.#notifyConnectionState(); - connection.deferred.resolve(message.snapshot); - return; - } - if (connection.state !== "connected") return; - if (message.type === "hello" || message.type === "hello_error") { - this.#protocolFailure(new ProtocolValidationError("Unexpected handshake message")); - return; - } + #handleMessage(message: ResponseEnvelope | EventEnvelope): void { if (message.type === "event") { - this.#stateStore.applyEvent(message.event); + this.#state.applyEvent(message.event); return; } - const pending = this.#pendingRequests.take(message.id); + const pending = this.#takePendingRequest(message.id); if (!pending) { - this.#protocolFailure(new ProtocolValidationError("Response has no matching request")); + this.#connection.fail(new ProtocolValidationError("Response has no matching request")); return; } if (!message.ok) { - pending.reject(new PiError(message.error)); + pending.reject(new PiServerError(message.error)); return; } if (message.result.command !== pending.command.command) { @@ -300,61 +244,54 @@ export class PiClient { `Response command ${message.result.command} does not match ${pending.command.command}`, ); pending.reject(error); - this.#protocolFailure(error); + this.#connection.fail(error); return; } - this.#stateStore.applyResult(message.result); + this.#state.applyResult(message.result); pending.resolve(message.result); } - #handleTransportClose(): void { - const connection = this.#connection; - if (connection.state === "disconnected") return; - let error: Error = new PiDisconnectedError("Byte transport closed"); - try { - connection.decoder.end(); - } catch (decoderError) { - error = toError(decoderError); + #handleConnectionStateChange(change: ConnectionStateChange): void { + if (change.state === "disconnected") { + this.#state.clearAttachments(); + this.#rejectPendingRequests(change.error ?? new PiDisconnectedError()); } - this.#failConnection(error); - } - - #handleTransportError(error: Error): void { - this.#failAndClose(toDisconnectedError(error)); + this.#notifyConnectionStateListeners(change); } - #protocolFailure(error: Error): void { - this.#failAndClose(error); + #takePendingRequest(id: string): PendingRequest | undefined { + const request = this.#pendingRequests.get(id); + if (request) this.#pendingRequests.delete(id); + return request; } - #failAndClose(error: Error): void { - const connection = this.#connection; - const transport = connection.state === "disconnected" ? undefined : connection.transport; - this.#failConnection(error); - transport?.close(); - } - - #failConnection(error: Error): void { - const connection = this.#connection; - if (connection.state === "disconnected") return; - this.#connection = { state: "disconnected" }; - this.#stateStore.clearAttachments(); - if (connection.state === "connecting") connection.deferred.reject(error); - this.#pendingRequests.rejectAll(error); - this.#notifyConnectionState(error); + #rejectPendingRequests(error: Error): void { + const requests = [...this.#pendingRequests.values()]; + this.#pendingRequests.clear(); + for (const request of requests) request.reject(error); } #assertSessionAttached(sessionId: string): void { - if (this.#connection.state !== "connected") throw new PiDisconnectedError(); - if (!this.#stateStore.isSessionAttached(sessionId)) throw new PiSessionDetachedError(sessionId); + if (!this.connected) throw new PiDisconnectedError(); + if (!this.#state.isSessionAttached(sessionId)) throw new PiSessionDetachedError(sessionId); } - #isCurrentConnection(connectionId: number): boolean { - return this.#connection.state !== "disconnected" && this.#connection.id === connectionId; + #notifyConnectionStateListeners(change: ConnectionStateChange): void { + for (const listener of this.#connectionStateListeners) { + try { + listener(change); + } catch (error) { + this.#reportListenerError(error); + } + } } - #notifyConnectionState(error?: Error): void { - const state = this.#connection.state; - notifyListeners(this.#stateListeners, error ? { state, error } : { state }, this.#options.onListenerError); + #reportListenerError(error: unknown): void { + if (!this.#options.onListenerError) return; + try { + this.#options.onListenerError(toError(error)); + } catch { + // Diagnostics cannot affect protocol or transport state. + } } } diff --git a/packages/client/src/connection.ts b/packages/client/src/connection.ts new file mode 100644 index 00000000000..8a6d7b904fc --- /dev/null +++ b/packages/client/src/connection.ts @@ -0,0 +1,239 @@ +import { + DEFAULT_MAX_FRAME_LENGTH, + encodeClientMessage, + PROTOCOL_VERSION, + ProtocolValidationError, + type ServerMessage, + ServerMessageDecoder, + type ServerSnapshot, +} from "@earendil-works/pi-protocol"; +import { PiDisconnectedError, PiServerError, toDisconnectedError, toError } from "./errors.ts"; +import type { ByteTransport, ByteTransportFactory, ByteTransportHandlers } from "./transport.ts"; +import type { ConnectionState, ConnectionStateChange } from "./types.ts"; + +const MAX_UINT32 = 0xffff_ffff; + +type ActiveConnection = { + id: number; + decoder: ServerMessageDecoder; + transport?: ByteTransport; +}; + +type ConnectionLifecycle = + | { state: "disconnected" } + | ({ state: "connecting"; handshake: PromiseWithResolvers } & ActiveConnection) + | ({ + state: "connected"; + transport: ByteTransport; + handshake: PromiseWithResolvers | undefined; + } & ActiveConnection); + +interface ConnectionOptions { + token: string; + transportFactory: ByteTransportFactory; + maxFrameLength?: number; + onHandshake(snapshot: ServerSnapshot): void; + onMessage(message: Exclude): void; + onStateChange(change: ConnectionStateChange): void; +} + +export class Connection { + readonly #options: ConnectionOptions; + readonly #maxFrameLength: number; + #lifecycle: ConnectionLifecycle = { state: "disconnected" }; + #sequence = 0; + + constructor(options: ConnectionOptions) { + this.#options = options; + this.#maxFrameLength = options.maxFrameLength ?? DEFAULT_MAX_FRAME_LENGTH; + if ( + !Number.isSafeInteger(this.#maxFrameLength) || + this.#maxFrameLength <= 0 || + this.#maxFrameLength > MAX_UINT32 + ) { + throw new TypeError(`PiClient maxFrameLength must be an integer between 1 and ${MAX_UINT32}`); + } + } + + get state(): ConnectionState { + return this.#lifecycle.state; + } + + get maxFrameLength(): number { + return this.#maxFrameLength; + } + + connect(): Promise { + if (this.#lifecycle.state !== "disconnected") { + return Promise.reject(new PiDisconnectedError(`PiClient is already ${this.#lifecycle.state}`)); + } + const id = ++this.#sequence; + const handshake = Promise.withResolvers(); + this.#lifecycle = { + state: "connecting", + id, + decoder: new ServerMessageDecoder({ maxFrameLength: this.#maxFrameLength }), + handshake, + }; + this.#options.onStateChange({ state: "connecting" }); + const handlers = { + onData: (chunk) => this.#handleData(id, chunk), + onClose: () => { + if (this.#isCurrent(id)) this.#handleClose(); + }, + onError: (error) => { + if (this.#isCurrent(id)) this.#failAndClose(toDisconnectedError(error)); + }, + } satisfies ByteTransportHandlers; + void this.#openTransport(id, handlers); + return handshake.promise; + } + + disconnect(reason = "Client disconnected"): void { + if (this.#lifecycle.state === "disconnected") return; + this.#failAndClose(new PiDisconnectedError(reason)); + } + + fail(error: Error): void { + this.#failAndClose(error); + } + + send(frame: Uint8Array): void { + const lifecycle = this.#lifecycle; + if (lifecycle.state !== "connected") throw new PiDisconnectedError(); + let sending: Promise; + try { + sending = lifecycle.transport.send(frame); + } catch (error) { + this.#failAndClose(toDisconnectedError(error)); + return; + } + void sending.catch((error: unknown) => { + const current = this.#lifecycle; + if (current.state !== "disconnected" && current.transport === lifecycle.transport) { + this.#failAndClose(toDisconnectedError(error)); + } + }); + } + + async #openTransport(id: number, handlers: ByteTransportHandlers): Promise { + let transport: ByteTransport; + try { + transport = await this.#options.transportFactory(handlers); + } catch (error) { + if (this.#isCurrent(id)) this.#fail(toDisconnectedError(error)); + return; + } + const lifecycle = this.#lifecycle; + if (lifecycle.state !== "connecting" || lifecycle.id !== id) { + transport.close(); + return; + } + this.#lifecycle = { ...lifecycle, transport }; + try { + await transport.send( + encodeClientMessage( + { type: "hello", version: PROTOCOL_VERSION, token: this.#options.token }, + { maxFrameLength: this.#maxFrameLength }, + ), + ); + } catch (error) { + if (this.#isCurrent(id)) this.#failAndClose(toDisconnectedError(error)); + } + } + + #handleData(id: number, chunk: Uint8Array): void { + const lifecycle = this.#lifecycle; + if (lifecycle.state === "disconnected" || lifecycle.id !== id) return; + if (lifecycle.state === "connecting" && !lifecycle.transport) { + this.#failAndClose(new ProtocolValidationError("Received server data before the client hello was sent")); + return; + } + let messages: ServerMessage[]; + try { + messages = lifecycle.decoder.push(chunk); + } catch (error) { + this.#failAndClose(toError(error)); + return; + } + for (const message of messages) { + if (this.#lifecycle.state === "disconnected") return; + this.#handleMessage(message); + } + } + + #handleMessage(message: ServerMessage): void { + const lifecycle = this.#lifecycle; + if (lifecycle.state === "connecting") { + if (message.type === "hello_error") { + this.#failAndClose(new PiServerError(message.error)); + return; + } + if (message.type !== "hello") { + this.#failAndClose(new ProtocolValidationError("Expected server hello as first message")); + return; + } + if (!lifecycle.transport) { + this.#failAndClose(new ProtocolValidationError("Received server hello before the client hello was sent")); + return; + } + const connected = { + state: "connected", + id: lifecycle.id, + decoder: lifecycle.decoder, + transport: lifecycle.transport, + handshake: lifecycle.handshake, + } satisfies Extract; + this.#lifecycle = connected; + try { + this.#options.onHandshake(message.snapshot); + } catch (error) { + if (this.#lifecycle === connected) this.#failAndClose(toError(error)); + return; + } + if (this.#lifecycle !== connected) return; + this.#options.onStateChange({ state: "connected" }); + if (this.#lifecycle !== connected) return; + this.#lifecycle = { ...connected, handshake: undefined }; + lifecycle.handshake.resolve(message.snapshot); + return; + } + if (lifecycle.state !== "connected") return; + if (message.type === "hello" || message.type === "hello_error") { + this.#failAndClose(new ProtocolValidationError("Unexpected handshake message")); + return; + } + this.#options.onMessage(message); + } + + #handleClose(): void { + const lifecycle = this.#lifecycle; + if (lifecycle.state === "disconnected") return; + let error: Error = new PiDisconnectedError("Byte transport closed"); + try { + lifecycle.decoder.end(); + } catch (decoderError) { + error = toError(decoderError); + } + this.#fail(error); + } + + #failAndClose(error: Error): void { + const lifecycle = this.#lifecycle; + const transport = lifecycle.state === "disconnected" ? undefined : lifecycle.transport; + this.#fail(error); + transport?.close(); + } + + #fail(error: Error): void { + const lifecycle = this.#lifecycle; + if (lifecycle.state === "disconnected") return; + this.#lifecycle = { state: "disconnected" }; + lifecycle.handshake?.reject(error); + this.#options.onStateChange({ state: "disconnected", error }); + } + + #isCurrent(id: number): boolean { + return this.#lifecycle.state !== "disconnected" && this.#lifecycle.id === id; + } +} diff --git a/packages/client/src/errors.ts b/packages/client/src/errors.ts index 54b6d16be1e..60b0fc7aade 100644 --- a/packages/client/src/errors.ts +++ b/packages/client/src/errors.ts @@ -1,12 +1,12 @@ import type { JsonValue, ProtocolError, ProtocolErrorCode } from "@earendil-works/pi-protocol"; -export class PiError extends Error { +export class PiServerError extends Error { readonly code: ProtocolErrorCode; readonly details: JsonValue | undefined; constructor(error: ProtocolError) { super(error.message); - this.name = "PiError"; + this.name = "PiServerError"; this.code = error.code; this.details = error.details; } diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts index 0b1235b6020..b2c9e628394 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -1,6 +1,6 @@ +export type { PiSessionHandle } from "./client.ts"; export { PiClient } from "./client.ts"; -export { PiDisconnectedError, PiError, PiSessionDetachedError } from "./errors.ts"; -export { PiSessionClient } from "./session-client.ts"; +export { PiDisconnectedError, PiServerError, PiSessionDetachedError } from "./errors.ts"; export type { ByteTransport, ByteTransportFactory, ByteTransportHandlers } from "./transport.ts"; export type { ConnectionState, diff --git a/packages/client/src/listeners.ts b/packages/client/src/listeners.ts deleted file mode 100644 index 9a850b34eeb..00000000000 --- a/packages/client/src/listeners.ts +++ /dev/null @@ -1,21 +0,0 @@ -import { toError } from "./errors.ts"; -import type { ListenerErrorHandler } from "./types.ts"; - -export function notifyListeners( - listeners: Iterable<(value: T) => void> | undefined, - value: T, - onError?: ListenerErrorHandler, -): void { - for (const listener of listeners ?? []) { - try { - listener(value); - } catch (error) { - if (!onError) continue; - try { - onError(toError(error)); - } catch { - // Diagnostics cannot affect protocol or transport state. - } - } - } -} diff --git a/packages/client/src/pending-requests.ts b/packages/client/src/pending-requests.ts deleted file mode 100644 index 6fbe1fc0292..00000000000 --- a/packages/client/src/pending-requests.ts +++ /dev/null @@ -1,33 +0,0 @@ -import type { Command, CommandResult, ResultForCommand } from "@earendil-works/pi-protocol"; - -interface PendingRequest { - command: Command; - resolve(result: CommandResult): void; - reject(error: Error): void; -} - -export class PendingRequests { - readonly #requests = new Map(); - #sequence = 0; - - create( - command: TCommand, - ): { id: string; promise: Promise> } { - const id = `request-${++this.#sequence}`; - const { promise, resolve, reject } = Promise.withResolvers(); - this.#requests.set(id, { command, resolve, reject }); - return { id, promise: promise as Promise> }; - } - - take(id: string): PendingRequest | undefined { - const request = this.#requests.get(id); - if (request) this.#requests.delete(id); - return request; - } - - rejectAll(error: Error): void { - const requests = [...this.#requests.values()]; - this.#requests.clear(); - for (const request of requests) request.reject(error); - } -} diff --git a/packages/client/src/session-client.ts b/packages/client/src/session-client.ts deleted file mode 100644 index a5c1dba920c..00000000000 --- a/packages/client/src/session-client.ts +++ /dev/null @@ -1,74 +0,0 @@ -import type { - Command, - ModelRef, - ResultForCommand, - ServerEvent, - SessionSnapshot, - ThinkingLevel, -} from "@earendil-works/pi-protocol"; -import type { Unsubscribe } from "./types.ts"; - -type SessionCommand = Extract; - -export interface SessionClientOperations { - isAttached(): boolean; - getSnapshot(): SessionSnapshot | undefined; - subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe; - onEvent(listener: (event: ServerEvent) => void): Unsubscribe; - request(command: TCommand): Promise>; -} - -export class PiSessionClient { - readonly id: string; - readonly #operations: SessionClientOperations; - - /** @internal Construct session handles through `PiClient`. */ - constructor(id: string, operations: SessionClientOperations) { - this.id = id; - this.#operations = operations; - } - - get attached(): boolean { - return this.#operations.isAttached(); - } - - get snapshot(): SessionSnapshot | undefined { - return this.#operations.getSnapshot(); - } - - subscribe(listener: (snapshot: SessionSnapshot) => void): Unsubscribe { - return this.#operations.subscribe(listener); - } - - onEvent(listener: (event: ServerEvent) => void): Unsubscribe { - return this.#operations.onEvent(listener); - } - - async detach(): Promise { - await this.#request({ command: "detach", sessionId: this.id }); - } - - async prompt(text: string): Promise { - return (await this.#request({ command: "prompt", sessionId: this.id, text })).session; - } - - async steer(text: string): Promise { - return (await this.#request({ command: "steer", sessionId: this.id, text })).session; - } - - async abort(): Promise { - return (await this.#request({ command: "abort", sessionId: this.id })).session; - } - - async setModel(model: ModelRef): Promise { - return (await this.#request({ command: "set_model", sessionId: this.id, model })).session; - } - - async setThinking(thinkingLevel: ThinkingLevel): Promise { - return (await this.#request({ command: "set_thinking", sessionId: this.id, thinkingLevel })).session; - } - - #request(command: TCommand): Promise> { - return this.#operations.request(command); - } -} diff --git a/packages/client/src/state-store.ts b/packages/client/src/state.ts similarity index 83% rename from packages/client/src/state-store.ts rename to packages/client/src/state.ts index 097dece20f6..fff3e2c2e9c 100644 --- a/packages/client/src/state-store.ts +++ b/packages/client/src/state.ts @@ -1,14 +1,8 @@ -import type { - CommandResult, - ServerEvent, - ServerSnapshot, - SessionSnapshot, - SessionSummary, -} from "@earendil-works/pi-protocol"; -import { notifyListeners } from "./listeners.ts"; +import type { CommandResult, ServerEvent, ServerSnapshot, SessionSnapshot } from "@earendil-works/pi-protocol"; +import { toError } from "./errors.ts"; import type { ListenerErrorHandler, Unsubscribe } from "./types.ts"; -export class ClientStateStore { +export class ClientState { readonly #sessionSnapshots = new Map(); readonly #attachedSessionIds = new Set(); readonly #snapshotListeners = new Set<(snapshot: ServerSnapshot) => void>(); @@ -26,10 +20,6 @@ export class ClientStateStore { return this.#snapshot; } - get sessions(): readonly SessionSummary[] { - return this.#snapshot?.sessions ?? []; - } - reset(): void { this.#snapshot = undefined; this.#sessionSnapshots.clear(); @@ -94,9 +84,9 @@ export class ClientStateStore { this.#sessionSnapshots.delete(event.sessionId); this.#attachedSessionIds.delete(event.sessionId); } - notifyListeners(this.#eventListeners, event, this.#onListenerError); + this.#notify(this.#eventListeners, event); const sessionId = getEventSessionId(event); - if (sessionId) notifyListeners(this.#sessionEventListeners.get(sessionId), event, this.#onListenerError); + if (sessionId) this.#notify(this.#sessionEventListeners.get(sessionId), event); } applyServerSnapshot(snapshot: ServerSnapshot): void { @@ -104,7 +94,7 @@ export class ClientStateStore { this.#snapshot = snapshot; this.#attachedSessionIds.clear(); for (const session of snapshot.sessions) if (session.attached) this.#attachedSessionIds.add(session.id); - notifyListeners(this.#snapshotListeners, snapshot, this.#onListenerError); + this.#notify(this.#snapshotListeners, snapshot); } #applySessionSnapshot(snapshot: SessionSnapshot, force = false): void { @@ -113,7 +103,26 @@ export class ClientStateStore { this.#sessionSnapshots.set(snapshot.id, snapshot); if (snapshot.attached) this.#attachedSessionIds.add(snapshot.id); else this.#attachedSessionIds.delete(snapshot.id); - notifyListeners(this.#sessionSnapshotListeners.get(snapshot.id), snapshot, this.#onListenerError); + this.#notify(this.#sessionSnapshotListeners.get(snapshot.id), snapshot); + } + + #notify(listeners: Iterable<(value: T) => void> | undefined, value: T): void { + for (const listener of listeners ?? []) { + try { + listener(value); + } catch (error) { + this.#reportListenerError(error); + } + } + } + + #reportListenerError(error: unknown): void { + if (!this.#onListenerError) return; + try { + this.#onListenerError(toError(error)); + } catch { + // Diagnostics cannot affect client state. + } } } diff --git a/packages/client/test/client-state.test.ts b/packages/client/test/client-state.test.ts deleted file mode 100644 index 0f894cf2bb1..00000000000 --- a/packages/client/test/client-state.test.ts +++ /dev/null @@ -1,200 +0,0 @@ -import { encodeCbor, encodeFrame, PROTOCOL_VERSION, ProtocolValidationError } from "@earendil-works/pi-protocol"; -import { describe, expect, test } from "vitest"; -import { PiClient } from "../src/index.ts"; -import { baseServerSnapshot, collectRequests, connectClient, MemoryByteServer, sessionSnapshot } from "./support.ts"; - -describe("PiClient", () => { - test("enforces the configured frame limit for outbound and inbound messages", async () => { - const server = new MemoryByteServer(); - server.onMessage((message) => { - if (message.type === "hello") { - server.send({ - type: "hello", - version: PROTOCOL_VERSION, - connectionId: "connection-1", - snapshot: baseServerSnapshot, - }); - } - }); - const client = new PiClient({ - token: "bearer-secret", - maxFrameLength: 512, - transportFactory: (handlers) => server.connect(handlers), - }); - await client.connect(); - const sentBefore = server.sentByClient.length; - await expect( - client.request({ command: "prompt", sessionId: "session-1", text: "x".repeat(1_000) }), - ).rejects.toBeInstanceOf(ProtocolValidationError); - expect(server.sentByClient).toHaveLength(sentBefore); - - server.sendRaw(new Uint8Array([0, 0, 2, 1])); - expect(client.connectionState).toBe("disconnected"); - }); - - test("disconnects on invalid protocol data", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - server.sendRaw(encodeFrame(encodeCbor({ type: "event", event: { type: "session_removed", sessionId: 1 } }))); - expect(client.connectionState).toBe("disconnected"); - }); - - test("reports truncated framing when the transport closes", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - const pending = client.listSessions(); - server.sendRaw(new Uint8Array([0, 0, 0, 2, 1])); - server.close(); - - await expect(pending).rejects.toMatchObject({ - name: "ProtocolValidationError", - message: expect.stringMatching(/truncated/i), - }); - expect(client.connectionState).toBe("disconnected"); - }); - - test("rejects a mismatched response instead of leaving its request pending", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - const requests = collectRequests(server); - const listed = client.listSessions(); - expect(requests).toMatchObject([{ request: { command: "list" } }]); - server.send({ - type: "response", - id: requests[0]!.id, - ok: true, - result: { command: "attach", session: sessionSnapshot("session-1") }, - }); - - await expect(listed).rejects.toMatchObject({ - name: "ProtocolValidationError", - message: "Response command attach does not match list", - }); - expect(client.connectionState).toBe("disconnected"); - }); - - test("does not let a delayed command response replace a newer event snapshot", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - const initial = sessionSnapshot("session-1", { revision: 1, thinkingLevel: "off" }); - server.send({ type: "event", event: { type: "session_snapshot", snapshot: initial } }); - const handle = client.getSession("session-1"); - const requests = collectRequests(server); - const changing = handle.setThinking("high"); - const request = requests.find((candidate) => candidate.request.command === "set_thinking"); - if (!request) throw new Error("Missing set_thinking request"); - server.send({ - type: "event", - event: { - type: "session_snapshot", - snapshot: sessionSnapshot("session-1", { revision: 3, thinkingLevel: "high" }), - }, - }); - server.send({ - type: "response", - id: request.id, - ok: true, - result: { - command: "set_thinking", - session: sessionSnapshot("session-1", { revision: 2, thinkingLevel: "medium" }), - }, - }); - - await changing; - expect(handle.snapshot).toMatchObject({ revision: 3, thinkingLevel: "high" }); - }); - - test("does not let an attach response replace a newer snapshot from the reacquired runtime", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - server.send({ - type: "event", - event: { - type: "session_snapshot", - snapshot: sessionSnapshot("session-1", { revision: 10, attached: false }), - }, - }); - server.onMessage((message) => { - if (message.type !== "request" || message.request.command !== "attach") return; - server.send({ - type: "event", - event: { - type: "session_snapshot", - snapshot: sessionSnapshot("session-1", { revision: 3, thinkingLevel: "high" }), - }, - }); - server.send({ - type: "response", - id: message.id, - ok: true, - result: { - command: "attach", - session: sessionSnapshot("session-1", { revision: 2, thinkingLevel: "medium" }), - }, - }); - }); - - const handle = await client.attachSession("session-1"); - expect(handle.snapshot).toMatchObject({ revision: 3, thinkingLevel: "high" }); - }); - - test("accepts a lower revision after detaching and reacquiring the same session", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - let attachCount = 0; - server.onMessage((message) => { - if (message.type !== "request") return; - if (message.request.command === "attach") { - server.send({ - type: "response", - id: message.id, - ok: true, - result: { - command: "attach", - session: sessionSnapshot("session-1", { revision: attachCount++ === 0 ? 10 : 0 }), - }, - }); - } - if (message.request.command === "detach") { - server.send({ - type: "response", - id: message.id, - ok: true, - result: { command: "detach", sessionId: "session-1" }, - }); - } - }); - - const first = await client.attachSession("session-1"); - expect(first.snapshot?.revision).toBe(10); - await first.detach(); - const reopened = await client.attachSession("session-1"); - expect(reopened.snapshot?.revision).toBe(0); - }); - - test("rejects frame limits outside the unsigned 32-bit range", () => { - const server = new MemoryByteServer(); - expect( - () => - new PiClient({ - token: "secret", - maxFrameLength: 0x1_0000_0000, - transportFactory: (handlers) => server.connect(handlers), - }), - ).toThrow(/maxFrameLength/); - }); - - test("surfaces typed request errors", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - const requests = collectRequests(server); - const attaching = client.attachSession("locked"); - server.send({ - type: "response", - id: requests[0]?.id ?? "missing", - ok: false, - error: { code: "session_locked", message: "Already attached" }, - }); - await expect(attaching).rejects.toMatchObject({ name: "PiError", code: "session_locked" }); - }); -}); diff --git a/packages/client/test/client-connection.test.ts b/packages/client/test/connection.test.ts similarity index 65% rename from packages/client/test/client-connection.test.ts rename to packages/client/test/connection.test.ts index 9c842740c2b..55cb07954a7 100644 --- a/packages/client/test/client-connection.test.ts +++ b/packages/client/test/connection.test.ts @@ -1,14 +1,17 @@ import { type ClientMessage, + encodeCbor, + encodeFrame, encodeServerMessage, PROTOCOL_VERSION, + ProtocolValidationError, type ServerSnapshot, } from "@earendil-works/pi-protocol"; import { describe, expect, test } from "vitest"; -import { type ByteTransportFactory, PiClient, PiDisconnectedError, PiSessionDetachedError } from "../src/index.ts"; +import { type ByteTransportFactory, PiClient, PiDisconnectedError } from "../src/index.ts"; import { + attachSession, baseServerSnapshot, - collectRequests, connectClient, createClient, MemoryByteServer, @@ -123,139 +126,77 @@ describe("PiClient", () => { expect(client.connectionState).toBe("connected"); }); - test("rejects a typed handshake authentication error", async () => { + test("does not restore a connection after a snapshot listener disconnects during handshake", async () => { const server = new MemoryByteServer(); - server.onMessage(() => { + server.onMessage((message) => { + if (message.type !== "hello") return; server.send({ - type: "hello_error", - error: { code: "auth", message: "Invalid token" }, + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, }); }); - const client = createClient(server, "wrong"); + const client = createClient(server); + client.subscribe(() => client.disconnect()); - await expect(client.connect()).rejects.toMatchObject({ - name: "PiError", - code: "auth", - message: "Invalid token", - }); + await expect(client.connect()).rejects.toBeInstanceOf(PiDisconnectedError); expect(client.connectionState).toBe("disconnected"); expect(server.clientCloseCount).toBe(1); }); - test("correlates coalesced out-of-order responses", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - const requests = collectRequests(server); - const listed = client.listSessions(); - const attached = client.attachSession("session-1"); - expect(requests).toHaveLength(2); - - const attachRequest = requests.find((request) => request.request.command === "attach"); - const listRequest = requests.find((request) => request.request.command === "list"); - if (!attachRequest || !listRequest) throw new Error("Missing requests"); - server.sendTogether([ - { - type: "response", - id: attachRequest.id, - ok: true, - result: { command: "attach", session: sessionSnapshot("session-1") }, - }, - { - type: "response", - id: listRequest.id, - ok: true, - result: { command: "list", sessions: [] }, - }, - ]); - - await expect(listed).resolves.toEqual([]); - await expect(attached).resolves.toMatchObject({ id: "session-1", attached: true }); - }); - - test("reduces only authoritative snapshots and supports unsubscribe", async () => { - const server = new MemoryByteServer(); - const client = await connectClient(server); - const requests = collectRequests(server); - const initial = sessionSnapshot("session-1", { revision: 1, phase: "idle" }); - server.send({ type: "event", event: { type: "session_snapshot", snapshot: initial } }); - const handle = client.getSession("session-1"); - const observed: number[] = []; - const progressTypes: string[] = []; - const unsubscribe = handle.subscribe((snapshot) => observed.push(snapshot.revision)); - const unsubscribeEvents = handle.onEvent((event) => progressTypes.push(event.type)); - server.send({ - type: "event", - event: { - type: "session_progress", - sessionId: "session-1", - progress: { - type: "assistant_delta", - messageId: "assistant-1", - contentIndex: 0, - kind: "text", - delta: "hi", - }, - }, + test("does not restore a stale connection when a snapshot listener reconnects during handshake", async () => { + const first = new MemoryByteServer(); + const second = new MemoryByteServer(); + let connection = 0; + for (const server of [first, second]) { + server.onMessage((message) => { + if (message.type !== "hello") return; + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: `connection-${connection}`, + snapshot: { ...baseServerSnapshot, revision: connection }, + }); + }); + } + const client = new PiClient({ + token: "bearer-secret", + transportFactory: (handlers) => (connection++ === 0 ? first : second).connect(handlers), }); - expect(progressTypes).toEqual(["session_progress"]); - expect(handle.snapshot).toEqual(initial); - - const prompting = handle.prompt("hello"); - expect(handle.snapshot).toEqual(initial); - const promptRequest = requests.find((request) => request.request.command === "prompt"); - if (!promptRequest) throw new Error("Missing prompt request"); - const updated = sessionSnapshot("session-1", { revision: 2, phase: "turn" }); - server.send({ - type: "response", - id: promptRequest.id, - ok: true, - result: { command: "prompt", session: updated }, + let reconnect: Promise | undefined; + let reconnectRequested = false; + client.subscribe(() => { + if (reconnectRequested) return; + reconnectRequested = true; + client.disconnect(); + reconnect = client.reconnect(); }); - await expect(prompting).resolves.toEqual(updated); - expect(handle.snapshot).toEqual(updated); - expect(observed).toEqual([2]); - unsubscribe(); - unsubscribeEvents(); - server.send({ - type: "event", - event: { type: "session_snapshot", snapshot: sessionSnapshot("session-1", { revision: 3 }) }, - }); - expect(observed).toEqual([2]); + await expect(client.connect()).rejects.toBeInstanceOf(PiDisconnectedError); + expect(reconnect).toBeDefined(); + await expect(reconnect).resolves.toMatchObject({ revision: 2 }); + expect(client.connectionState).toBe("connected"); + expect(first.clientCloseCount).toBe(1); }); - test("keeps multiple session handles independent and enforces detach", async () => { + test("rejects a typed handshake authentication error", async () => { const server = new MemoryByteServer(); - const client = await connectClient(server); - server.onMessage((message) => { - if (message.type !== "request") return; - const request = message.request; - if (request.command === "attach") { - server.send({ - type: "response", - id: message.id, - ok: true, - result: { command: "attach", session: sessionSnapshot(request.sessionId) }, - }); - } - if (request.command === "detach") { - server.send({ - type: "response", - id: message.id, - ok: true, - result: { command: "detach", sessionId: request.sessionId }, - }); - } + server.onMessage(() => { + server.send({ + type: "hello_error", + error: { code: "auth", message: "Invalid token" }, + }); }); + const client = createClient(server, "wrong"); - const first = await client.attachSession("session-1"); - const second = await client.attachSession("session-2"); - expect(first.attached).toBe(true); - expect(second.attached).toBe(true); - await first.detach(); - expect(first.attached).toBe(false); - expect(second.attached).toBe(true); - await expect(first.abort()).rejects.toBeInstanceOf(PiSessionDetachedError); + await expect(client.connect()).rejects.toMatchObject({ + name: "PiServerError", + code: "auth", + message: "Invalid token", + }); + expect(client.connectionState).toBe("disconnected"); + expect(server.clientCloseCount).toBe(1); }); test("rejects pending requests on close and reconnects through a fresh factory result", async () => { @@ -330,4 +271,64 @@ describe("PiClient", () => { await expect(pending).rejects.toMatchObject({ name: "PiDisconnectedError", message: "read failed" }); expect(client.connectionState).toBe("disconnected"); }); + + test("enforces the configured frame limit for outbound and inbound messages", async () => { + const server = new MemoryByteServer(); + server.onMessage((message) => { + if (message.type === "hello") { + server.send({ + type: "hello", + version: PROTOCOL_VERSION, + connectionId: "connection-1", + snapshot: baseServerSnapshot, + }); + } + }); + const client = new PiClient({ + token: "bearer-secret", + maxFrameLength: 512, + transportFactory: (handlers) => server.connect(handlers), + }); + await client.connect(); + const handle = await attachSession(client, server, sessionSnapshot("session-1")); + const sentBefore = server.sentByClient.length; + await expect(handle.prompt("x".repeat(1_000))).rejects.toBeInstanceOf(ProtocolValidationError); + expect(server.sentByClient).toHaveLength(sentBefore); + + server.sendRaw(new Uint8Array([0, 0, 2, 1])); + expect(client.connectionState).toBe("disconnected"); + }); + + test("disconnects on invalid protocol data", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + server.sendRaw(encodeFrame(encodeCbor({ type: "event", event: { type: "session_removed", sessionId: 1 } }))); + expect(client.connectionState).toBe("disconnected"); + }); + + test("reports truncated framing when the transport closes", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const pending = client.listSessions(); + server.sendRaw(new Uint8Array([0, 0, 0, 2, 1])); + server.close(); + + await expect(pending).rejects.toMatchObject({ + name: "ProtocolValidationError", + message: expect.stringMatching(/truncated/i), + }); + expect(client.connectionState).toBe("disconnected"); + }); + + test("rejects frame limits outside the unsigned 32-bit range", () => { + const server = new MemoryByteServer(); + expect( + () => + new PiClient({ + token: "secret", + maxFrameLength: 0x1_0000_0000, + transportFactory: (handlers) => server.connect(handlers), + }), + ).toThrow(/maxFrameLength/); + }); }); diff --git a/packages/client/test/requests.test.ts b/packages/client/test/requests.test.ts new file mode 100644 index 00000000000..3ddfea3b363 --- /dev/null +++ b/packages/client/test/requests.test.ts @@ -0,0 +1,68 @@ +import { describe, expect, test } from "vitest"; +import { collectRequests, connectClient, MemoryByteServer, sessionSnapshot } from "./support.ts"; + +describe("PiClient", () => { + test("correlates coalesced out-of-order responses", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const listed = client.listSessions(); + const attached = client.attachSession("session-1"); + expect(requests).toHaveLength(2); + + const attachRequest = requests.find((request) => request.request.command === "attach"); + const listRequest = requests.find((request) => request.request.command === "list"); + if (!attachRequest || !listRequest) throw new Error("Missing requests"); + server.sendTogether([ + { + type: "response", + id: attachRequest.id, + ok: true, + result: { command: "attach", session: sessionSnapshot("session-1") }, + }, + { + type: "response", + id: listRequest.id, + ok: true, + result: { command: "list", sessions: [] }, + }, + ]); + + await expect(listed).resolves.toEqual([]); + await expect(attached).resolves.toMatchObject({ id: "session-1", attached: true }); + }); + + test("rejects a mismatched response instead of leaving its request pending", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const listed = client.listSessions(); + expect(requests).toMatchObject([{ request: { command: "list" } }]); + server.send({ + type: "response", + id: requests[0]!.id, + ok: true, + result: { command: "attach", session: sessionSnapshot("session-1") }, + }); + + await expect(listed).rejects.toMatchObject({ + name: "ProtocolValidationError", + message: "Response command attach does not match list", + }); + expect(client.connectionState).toBe("disconnected"); + }); + + test("surfaces typed request errors", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const attaching = client.attachSession("locked"); + server.send({ + type: "response", + id: requests[0]?.id ?? "missing", + ok: false, + error: { code: "session_locked", message: "Already attached" }, + }); + await expect(attaching).rejects.toMatchObject({ name: "PiServerError", code: "session_locked" }); + }); +}); diff --git a/packages/client/test/sessions.test.ts b/packages/client/test/sessions.test.ts new file mode 100644 index 00000000000..098111e918b --- /dev/null +++ b/packages/client/test/sessions.test.ts @@ -0,0 +1,74 @@ +import { describe, expect, test } from "vitest"; +import { PiSessionDetachedError } from "../src/index.ts"; +import { connectClient, MemoryByteServer, sessionSnapshot } from "./support.ts"; + +describe("PiClient", () => { + test("keeps multiple session handles independent and enforces detach", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + server.onMessage((message) => { + if (message.type !== "request") return; + const request = message.request; + if (request.command === "attach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { command: "attach", session: sessionSnapshot(request.sessionId) }, + }); + } + if (request.command === "detach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { command: "detach", sessionId: request.sessionId }, + }); + } + }); + + const first = await client.attachSession("session-1"); + const second = await client.attachSession("session-2"); + expect(first.attached).toBe(true); + expect(second.attached).toBe(true); + await first.detach(); + expect(first.attached).toBe(false); + expect(second.attached).toBe(true); + await expect(first.abort()).rejects.toBeInstanceOf(PiSessionDetachedError); + }); + + test("accepts a lower revision after detaching and reacquiring the same session", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + let attachCount = 0; + server.onMessage((message) => { + if (message.type !== "request") return; + if (message.request.command === "attach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { + command: "attach", + session: sessionSnapshot("session-1", { revision: attachCount++ === 0 ? 10 : 0 }), + }, + }); + } + if (message.request.command === "detach") { + server.send({ + type: "response", + id: message.id, + ok: true, + result: { command: "detach", sessionId: "session-1" }, + }); + } + }); + + const first = await client.attachSession("session-1"); + expect(first.snapshot?.revision).toBe(10); + await first.detach(); + const reopened = await client.attachSession("session-1"); + expect(reopened).toBe(first); + expect(reopened.snapshot?.revision).toBe(0); + }); +}); diff --git a/packages/client/test/state.test.ts b/packages/client/test/state.test.ts new file mode 100644 index 00000000000..1e32117e967 --- /dev/null +++ b/packages/client/test/state.test.ts @@ -0,0 +1,119 @@ +import { describe, expect, test } from "vitest"; +import { attachSession, collectRequests, connectClient, MemoryByteServer, sessionSnapshot } from "./support.ts"; + +describe("PiClient", () => { + test("reduces only authoritative snapshots and supports unsubscribe", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const requests = collectRequests(server); + const initial = sessionSnapshot("session-1", { revision: 1, phase: "idle" }); + const handle = await attachSession(client, server, initial); + const observed: number[] = []; + const progressTypes: string[] = []; + const unsubscribe = handle.subscribe((snapshot) => observed.push(snapshot.revision)); + const unsubscribeEvents = handle.onEvent((event) => progressTypes.push(event.type)); + server.send({ + type: "event", + event: { + type: "session_progress", + sessionId: "session-1", + progress: { + type: "assistant_delta", + messageId: "assistant-1", + contentIndex: 0, + kind: "text", + delta: "hi", + }, + }, + }); + expect(progressTypes).toEqual(["session_progress"]); + expect(handle.snapshot).toEqual(initial); + + const prompting = handle.prompt("hello"); + expect(handle.snapshot).toEqual(initial); + const promptRequest = requests.find((request) => request.request.command === "prompt"); + if (!promptRequest) throw new Error("Missing prompt request"); + const updated = sessionSnapshot("session-1", { revision: 2, phase: "turn" }); + server.send({ + type: "response", + id: promptRequest.id, + ok: true, + result: { command: "prompt", session: updated }, + }); + await expect(prompting).resolves.toEqual(updated); + expect(handle.snapshot).toEqual(updated); + expect(observed).toEqual([2]); + + unsubscribe(); + unsubscribeEvents(); + server.send({ + type: "event", + event: { type: "session_snapshot", snapshot: sessionSnapshot("session-1", { revision: 3 }) }, + }); + expect(observed).toEqual([2]); + }); + + test("does not let a delayed command response replace a newer event snapshot", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + const initial = sessionSnapshot("session-1", { revision: 1, thinkingLevel: "off" }); + const handle = await attachSession(client, server, initial); + const requests = collectRequests(server); + const changing = handle.setThinking("high"); + const request = requests.find((candidate) => candidate.request.command === "set_thinking"); + if (!request) throw new Error("Missing set_thinking request"); + server.send({ + type: "event", + event: { + type: "session_snapshot", + snapshot: sessionSnapshot("session-1", { revision: 3, thinkingLevel: "high" }), + }, + }); + server.send({ + type: "response", + id: request.id, + ok: true, + result: { + command: "set_thinking", + session: sessionSnapshot("session-1", { revision: 2, thinkingLevel: "medium" }), + }, + }); + + await changing; + expect(handle.snapshot).toMatchObject({ revision: 3, thinkingLevel: "high" }); + }); + + test("does not let an attach response replace a newer snapshot from the reacquired runtime", async () => { + const server = new MemoryByteServer(); + const client = await connectClient(server); + server.send({ + type: "event", + event: { + type: "session_snapshot", + snapshot: sessionSnapshot("session-1", { revision: 10, attached: false }), + }, + }); + server.onMessage((message) => { + if (message.type !== "request" || message.request.command !== "attach") return; + server.send({ + type: "event", + event: { + type: "session_snapshot", + snapshot: sessionSnapshot("session-1", { revision: 3, thinkingLevel: "high" }), + }, + }); + server.send({ + type: "response", + id: message.id, + ok: true, + result: { + command: "attach", + session: sessionSnapshot("session-1", { revision: 2, thinkingLevel: "medium" }), + }, + }); + }); + + const handle = await client.attachSession("session-1"); + expect(handle.snapshot).toMatchObject({ revision: 3, thinkingLevel: "high" }); + }); +}); diff --git a/packages/client/test/support.ts b/packages/client/test/support.ts index 03818670913..fdd3e8a2928 100644 --- a/packages/client/test/support.ts +++ b/packages/client/test/support.ts @@ -8,7 +8,7 @@ import { type ServerSnapshot, type SessionSnapshot, } from "@earendil-works/pi-protocol"; -import type { ByteTransport, ByteTransportHandlers } from "../src/index.ts"; +import type { ByteTransport, ByteTransportHandlers, PiSessionHandle } from "../src/index.ts"; import { PiClient } from "../src/index.ts"; export class MemoryByteServer { @@ -134,3 +134,21 @@ export function collectRequests(server: MemoryByteServer): RequestEnvelope[] { }); return requests; } + +export async function attachSession( + client: PiClient, + server: MemoryByteServer, + snapshot: SessionSnapshot, +): Promise { + const requests = collectRequests(server); + const attaching = client.attachSession(snapshot.id); + const request = requests.find((candidate) => candidate.request.command === "attach"); + if (!request) throw new Error("Missing attach request"); + server.send({ + type: "response", + id: request.id, + ok: true, + result: { command: "attach", session: snapshot }, + }); + return attaching; +} From 6e16f771610359b959501acd4c42ab73a1c914e8 Mon Sep 17 00:00:00 2001 From: Christian Klotz Date: Fri, 31 Jul 2026 11:38:43 +0300 Subject: [PATCH 4/5] refactor(client): retain ES2022 compatibility --- packages/client/src/client.ts | 3 ++- packages/client/src/connection.ts | 7 ++++--- packages/client/src/promise.ts | 15 +++++++++++++++ tsconfig.base.json | 2 +- 4 files changed, 22 insertions(+), 5 deletions(-) create mode 100644 packages/client/src/promise.ts diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index 7318d662896..896a00559ab 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -15,6 +15,7 @@ import { } from "@earendil-works/pi-protocol"; import { Connection } from "./connection.ts"; import { PiDisconnectedError, PiServerError, PiSessionDetachedError, toError } from "./errors.ts"; +import { createPromiseResolvers } from "./promise.ts"; import { ClientState } from "./state.ts"; import type { ConnectionState, @@ -191,7 +192,7 @@ export class PiClient { #request(command: TCommand): Promise> { if (!this.connected) return Promise.reject(new PiDisconnectedError()); const id = `request-${++this.#requestSequence}`; - const { promise, resolve, reject } = Promise.withResolvers(); + const { promise, resolve, reject } = createPromiseResolvers(); this.#pendingRequests.set(id, { command, resolve, reject }); let frame: Uint8Array; try { diff --git a/packages/client/src/connection.ts b/packages/client/src/connection.ts index 8a6d7b904fc..db4cd62e800 100644 --- a/packages/client/src/connection.ts +++ b/packages/client/src/connection.ts @@ -8,6 +8,7 @@ import { type ServerSnapshot, } from "@earendil-works/pi-protocol"; import { PiDisconnectedError, PiServerError, toDisconnectedError, toError } from "./errors.ts"; +import { createPromiseResolvers, type PromiseResolvers } from "./promise.ts"; import type { ByteTransport, ByteTransportFactory, ByteTransportHandlers } from "./transport.ts"; import type { ConnectionState, ConnectionStateChange } from "./types.ts"; @@ -21,11 +22,11 @@ type ActiveConnection = { type ConnectionLifecycle = | { state: "disconnected" } - | ({ state: "connecting"; handshake: PromiseWithResolvers } & ActiveConnection) + | ({ state: "connecting"; handshake: PromiseResolvers } & ActiveConnection) | ({ state: "connected"; transport: ByteTransport; - handshake: PromiseWithResolvers | undefined; + handshake: PromiseResolvers | undefined; } & ActiveConnection); interface ConnectionOptions { @@ -68,7 +69,7 @@ export class Connection { return Promise.reject(new PiDisconnectedError(`PiClient is already ${this.#lifecycle.state}`)); } const id = ++this.#sequence; - const handshake = Promise.withResolvers(); + const handshake = createPromiseResolvers(); this.#lifecycle = { state: "connecting", id, diff --git a/packages/client/src/promise.ts b/packages/client/src/promise.ts new file mode 100644 index 00000000000..7de26806011 --- /dev/null +++ b/packages/client/src/promise.ts @@ -0,0 +1,15 @@ +export interface PromiseResolvers { + promise: Promise; + resolve(value: T | PromiseLike): void; + reject(reason?: unknown): void; +} + +export function createPromiseResolvers(): PromiseResolvers { + let resolve!: PromiseResolvers["resolve"]; + let reject!: PromiseResolvers["reject"]; + const promise = new Promise((resolvePromise, rejectPromise) => { + resolve = resolvePromise; + reject = rejectPromise; + }); + return { promise, resolve, reject }; +} diff --git a/tsconfig.base.json b/tsconfig.base.json index 2f338f05805..57e97d6e361 100644 --- a/tsconfig.base.json +++ b/tsconfig.base.json @@ -2,7 +2,7 @@ "compilerOptions": { "target": "ES2022", "module": "Node16", - "lib": ["ES2024"], + "lib": ["ES2022"], "strict": true, "erasableSyntaxOnly": true, "esModuleInterop": true, From f9b06345348f13799cdc5c0c8acec74e13e01276 Mon Sep 17 00:00:00 2001 From: Christian Klotz Date: Fri, 31 Jul 2026 11:43:52 +0300 Subject: [PATCH 5/5] docs(client): note promise helper removal --- packages/client/src/promise.ts | 1 + 1 file changed, 1 insertion(+) diff --git a/packages/client/src/promise.ts b/packages/client/src/promise.ts index 7de26806011..f5a65fc3677 100644 --- a/packages/client/src/promise.ts +++ b/packages/client/src/promise.ts @@ -4,6 +4,7 @@ export interface PromiseResolvers { reject(reason?: unknown): void; } +/** Remove in favor of `Promise.withResolvers()` when the repository's TypeScript lib baseline moves to ES2024. */ export function createPromiseResolvers(): PromiseResolvers { let resolve!: PromiseResolvers["resolve"]; let reject!: PromiseResolvers["reject"];