From 33bc0a7b83cde2a3cd2b568ed3107d9bdf7a768a Mon Sep 17 00:00:00 2001 From: Christian Klotz Date: Fri, 31 Jul 2026 00:41:53 +0300 Subject: [PATCH] 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 e55c44da2..1149569d9 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 107d54bce..d63205796 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 2c9938ba3..dc8244e44 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 9b1ec714b..d59d06244 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 000000000..b94aff2dd --- /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 000000000..e251429bb --- /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 000000000..5ca572253 --- /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 000000000..d79d9741c --- /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 000000000..54b6d16be --- /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 000000000..d7ddefdb5 --- /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 000000000..a94680c4c --- /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 000000000..dccaccb97 --- /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 000000000..71b7489c3 --- /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 000000000..aba7b4966 --- /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 000000000..d76d12953 --- /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 000000000..0f894cf2b --- /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 000000000..038186709 --- /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 000000000..0e72ae51a --- /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 000000000..b5e61b3d6 --- /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 000000000..7e6abe3f9 --- /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 64927b6e5..bfb811901 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 3b275fc88..dd046fada 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 967f61373..0053f6b1c 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 4d36db2f7..a409cdf11 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"],