From dd9e6ecd037c223594aea2ade82965548d31ea30 Mon Sep 17 00:00:00 2001 From: Zachary Davison Date: Thu, 3 Sep 2026 17:33:31 +0200 Subject: [PATCH 1/4] feat: pluggable client transport with a plain WebSocket option UseAIClient now takes a UseAITransport instead of building its own Socket.IO socket. Two transports ship: SocketIOTransport keeps today's behaviour and stays the default, WebSocketTransport carries JSON frames over a plain WebSocket for servers that do not speak Socket.IO. The server serves both on one port. webSocketPath (default '/ws') adds a plain listener beside Socket.IO in both runtime adapters. ClientSession.socket narrows to ClientConnection, which a Socket.IO socket satisfies, so plugins are unchanged. UseAIProvider accepts a transport prop, read once so an inline object does not churn the connection. new UseAIClient(url) and a provider given only serverUrl behave exactly as before. The UseAIClient suite now runs table-driven over both transports. An integration test drives a real WebSocketTransport through a full run on both the Bun and Node adapters. Co-Authored-By: Claude Opus 5 (1M context) --- CLAUDE.md | 2 + README.md | 42 + bun.lock | 5 + docs/websocket-protocol.md | 138 +++ packages/client/src/client.test.ts | 809 +++++++++--------- packages/client/src/client.ts | 146 ++-- packages/client/src/index.ts | 8 + .../client/src/providers/useAIProvider.tsx | 30 +- .../src/transport/SocketIOTransport.test.ts | 183 ++++ .../client/src/transport/SocketIOTransport.ts | 115 +++ .../src/transport/WebSocketTransport.test.ts | 286 +++++++ .../src/transport/WebSocketTransport.ts | 224 +++++ .../client/src/transport/handlerRegistry.ts | 32 + packages/client/src/transport/index.ts | 5 + packages/client/src/transport/types.ts | 42 + packages/client/src/types.ts | 5 +- packages/server/package.json | 3 + packages/server/src/agents/types.ts | 20 +- packages/server/src/index.ts | 2 +- .../src/runtime/bun/BunRuntimeAdapter.ts | 42 +- .../server/src/runtime/bun/rawWebSocket.ts | 59 ++ packages/server/src/runtime/index.ts | 2 + .../src/runtime/node/NodeRuntimeAdapter.ts | 22 + .../server/src/runtime/node/rawWebSocket.ts | 31 + packages/server/src/runtime/types.ts | 37 + packages/server/src/server.ts | 222 +++-- packages/server/src/types.ts | 10 + packages/server/src/webSocketConnection.ts | 22 + .../websocket-transport.integration.test.ts | 198 +++++ 29 files changed, 2158 insertions(+), 584 deletions(-) create mode 100644 docs/websocket-protocol.md create mode 100644 packages/client/src/transport/SocketIOTransport.test.ts create mode 100644 packages/client/src/transport/SocketIOTransport.ts create mode 100644 packages/client/src/transport/WebSocketTransport.test.ts create mode 100644 packages/client/src/transport/WebSocketTransport.ts create mode 100644 packages/client/src/transport/handlerRegistry.ts create mode 100644 packages/client/src/transport/index.ts create mode 100644 packages/client/src/transport/types.ts create mode 100644 packages/server/src/runtime/bun/rawWebSocket.ts create mode 100644 packages/server/src/runtime/node/rawWebSocket.ts create mode 100644 packages/server/src/webSocketConnection.ts create mode 100644 packages/server/src/websocket-transport.integration.test.ts diff --git a/CLAUDE.md b/CLAUDE.md index e51248f5..3e60fb0d 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -61,6 +61,8 @@ bun run kill # Kill processes on ports 3000, 3002, 8081 Use `wss://your-domain.com` for secure WebSocket connections. For local development without SSL, use `ws://localhost:8081`. +The client reaches the server through a `UseAITransport`. `SocketIOTransport` is the default. `WebSocketTransport` carries JSON frames over a plain WebSocket, which the server serves at `webSocketPath` (default `/ws`) on the same port. The framing is documented in `docs/websocket-protocol.md`. + ## Core Architecture ### Data Flow diff --git a/README.md b/README.md index 84ce4d36..f5f112ba 100644 --- a/README.md +++ b/README.md @@ -25,6 +25,7 @@ A React client/framework for easily enabling AI to control your users frontend. - [Features](#features) - [General](#general) - [AG-UI Protocol](#ag-ui-protocol) + - [Transports](#transports) - [Client](#client) - [`useAI` hook](#useai-hook) - [`UseAIProvider`](#useaiprovider) @@ -297,6 +298,44 @@ There are some minor extensions to the protocol: **Message Types**: - `run_workflow`: Trigger headless workflow (use-ai extension) [see `@meetsmore-oss/use-ai-plugin-workflows`] +### Transports + +The client reaches the server through a `UseAITransport`. Two transports ship with the library. + +| Transport | Wire | Server endpoint | +| -------------------- | ---------------------------------------- | ------------------------------ | +| `SocketIOTransport` | Socket.IO, over polling and WebSocket | `/socket.io/` (the default) | +| `WebSocketTransport` | JSON text frames, over a plain WebSocket | `webSocketPath`, default `/ws` | + +`UseAIProvider` builds a `SocketIOTransport` from `serverUrl` when you do not pass one, so +nothing changes if you use the bundled server. + +Pass `WebSocketTransport` to reach a server that does not serve Socket.IO. Such a server +does not have to be Node. It must accept a WebSocket connection. It must then exchange +the documented frames. + +```tsx +import { UseAIProvider, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; + +root.render( + + + +); +``` + +The bundled server serves both listeners on one port. Set `webSocketPath: null` to serve +Socket.IO only. + +To carry the same messages over something else, implement `UseAITransport` yourself. The +interface has five members: `connect`, `disconnect`, `send`, `on` and `connected`. + +See [docs/websocket-protocol.md](docs/websocket-protocol.md) for the frames, the turn +sequence, and the reconnection behaviour. + ## Client ### `useAI` hook @@ -353,6 +392,8 @@ root.render( ); ``` +Pass `transport` to reach a server over something other than Socket.IO. See [Transports](#transports). + ### Component State via `prompt` When you call `useAI`, you can provide a prompt that is used to tell the LLM the state of the component in a text-friendly way. @@ -1048,6 +1089,7 @@ const server = new UseAIServer({ }) }, defaultAgent: 'claude', + webSocketPath: '/ws', // plain WebSocket listener, see 'Transports'. null to disable. rateLimitMaxRequests: 1_000, rateLimitWindowMs: 60_000, plugins: [ // see 'Plugins' diff --git a/bun.lock b/bun.lock index 5178962a..09a1ccb3 100644 --- a/bun.lock +++ b/bun.lock @@ -137,15 +137,18 @@ "picomatch": "^4.0.3", "socket.io": "^4.8.1", "uuid": "^11.1.0", + "ws": "^8.18.3", "zod": "^3.24.1", "zod-to-json-schema": "^3.25.1", }, "devDependencies": { + "@meetsmore-oss/use-ai-client": "workspace:^", "@types/bun": "^1.3.13", "@types/cors": "^2.8.17", "@types/json-schema": "^7.0.15", "@types/picomatch": "^4.0.2", "@types/uuid": "^10.0.0", + "@types/ws": "^8.18.1", "socket.io-client": "^4.8.1", "testcontainers": "^11.8.1", "typescript": "^5.9.3", @@ -545,6 +548,8 @@ "@types/uuid": ["@types/uuid@10.0.0", "", {}, "sha512-7gqG38EyHgyP1S+7+xomFtL+ZNHcKv6DwNaCZmJmo1vgMugyF3TCnXVg4t1uk89mLNwnLtnY3TpOpCOyp1/xHQ=="], + "@types/ws": ["@types/ws@8.18.1", "", { "dependencies": { "@types/node": "*" } }, "sha512-ThVF6DCVhA8kUGy+aazFQ4kXQ7E1Ty7A3ypFOe0IcJV8O/M511G99AW24irKrW56Wt44yG9+ij8FaqoBGkuBXg=="], + "@vercel/oidc": ["@vercel/oidc@3.0.5", "", {}, "sha512-fnYhv671l+eTTp48gB4zEsTW/YtRgRPnkI2nT7x6qw5rkI1Lq2hTmQIpHPgyThI0znLK+vX2n9XxKdXZ7BUbbw=="], "abort-controller": ["abort-controller@3.0.0", "", { "dependencies": { "event-target-shim": "^5.0.0" } }, "sha512-h8lQ8tacZYnR3vNQTgibj+tODHI5/+l06Au2Pcriv/Gmet0eaj4TwWH41sO9wnHDiQsEj19q0drzdWdeAHtweg=="], diff --git a/docs/websocket-protocol.md b/docs/websocket-protocol.md new file mode 100644 index 00000000..5c2b37a0 --- /dev/null +++ b/docs/websocket-protocol.md @@ -0,0 +1,138 @@ +# Plain WebSocket protocol + +The bundled server serves Socket.IO by default. It also serves a plain WebSocket +listener on the same port, at `/ws`. A client reaches that listener with +`WebSocketTransport` instead of the default `SocketIOTransport`. + +Use this protocol to connect the `use-ai` chat UI and hooks to your own server. +Your server does not have to be Node. It does not have to implement Socket.IO. +Your server must accept a WebSocket connection. It must then exchange the JSON text +frames below. + +## Client setup + +```tsx +import { UseAIProvider, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; + +root.render( + + + +); +``` + +The provider reads `transport` once, on the first render. An inline object therefore +does not reconnect the client on every render. To change transports, remount the provider. + +When you pass a transport, the provider does not use `serverUrl`. The prop stays +required. The provider reports `serverUrl` on the context for application code that +reads it. + +## Server setup + +The bundled server enables the listener by default: + +```typescript +const server = new UseAIServer({ + agents: { claude }, + defaultAgent: 'claude', + webSocketPath: '/agent', // default: '/ws'. Pass null to serve Socket.IO only. +}); +``` + +## Upstream frames + +The client sends the `UseAIClientMessage` object, serialized, with nothing wrapped +around it. Each message is one text frame. + +```json +{ "type": "run_agent", "data": { "threadId": "...", "runId": "...", "messages": [], "tools": [], "state": null, "forwardedProps": {} } } +``` + +The message types are `run_agent`, `tool_result`, `tool_approval_response`, +`abort_run` and `message_feedback`. Plugins add more. See `UseAIClientMessage` in +`@meetsmore-oss/use-ai-core` for each payload. + +## Downstream frames + +A plain WebSocket has no event names of its own, so the server wraps each payload in +a named envelope. Each envelope is one text frame. + +```json +{ "name": "event", "data": { "type": "TEXT_MESSAGE_CONTENT", "messageId": "...", "delta": "Hello" } } +{ "name": "agents", "data": { "agents": [{ "id": "claude", "name": "Claude" }], "defaultAgent": "claude" } } +{ "name": "config", "data": { "langfuseEnabled": true } } +``` + +| Name | Payload | When | +| -------- | -------------------------------------------------- | ------------------------------------------------------ | +| `agents` | The agent list and the default agent id | Once, after connect | +| `config` | Server capability flags, such as `langfuseEnabled` | Once, after connect, if the server has flags to report | +| `event` | One AG-UI event | Throughout a run | + +The client reads `name` and finds the handlers for that name. It then passes `data` +to each handler. + +The client ignores a frame with an unknown `name`. An unknown name is not an error. +The client does not close the connection. A server can therefore send a name that an +older client does not know. + +The `event` payload is an AG-UI event. The event types the client handles are +`RUN_STARTED`, `RUN_FINISHED`, `RUN_ERROR`, `STEP_STARTED`, `STEP_FINISHED`, +`TEXT_MESSAGE_START`, `TEXT_MESSAGE_CONTENT`, `TEXT_MESSAGE_END`, `TOOL_CALL_START`, +`TOOL_CALL_ARGS`, `TOOL_CALL_END`, `TOOL_CALL_RESULT`, the `REASONING_*` events, and +the `TOOL_APPROVAL_REQUEST` extension. See the +[AG-UI protocol](https://docs.ag-ui.com/introduction) for each event, and +`packages/core/src/types.ts` for the types this library uses. + +## One turn, step by step + +1. The client opens the connection. +2. The server sends `agents`. It then sends `config`. +3. The client sends `run_agent` with the prompt, the tool definitions and the app state. +4. The server sends `event` frames for `RUN_STARTED`, then the model output. +5. For a client-side tool, the server sends `TOOL_CALL_START`, `TOOL_CALL_ARGS` and `TOOL_CALL_END`. +6. The client runs the tool. It then sends `tool_result` with the output. +7. The server resumes the model. It then sends the remaining `event` frames. +8. The server sends `RUN_FINISHED`. + +## Reconnection + +A plain WebSocket has no reconnection. `WebSocketTransport` therefore runs its own +retry loop. It retries indefinitely, with exponential backoff capped at ten seconds. +The cap matches the Socket.IO settings. A mobile app in the background, or a device +in airplane mode, thus recovers without frequent retries. + +Set both delays in the options: + +```typescript +new WebSocketTransport('wss://your-server.com/ws', { + reconnectionDelay: 1000, // first retry, in milliseconds + reconnectionDelayMax: 10000, // cap on the backoff, in milliseconds +}); +``` + +The server destroys the session when the connection closes. A reconnected client +therefore starts a new session. The client sends its conversation history with the +next `run_agent`. + +## Writing your own transport + +`UseAITransport` has five members: + +- `connect` +- `disconnect` +- `send` +- `on` +- `connected` + +Implement `UseAITransport` to carry the same messages over something else. + +```typescript +import { UseAIClient, type UseAITransport } from '@meetsmore-oss/use-ai-client'; + +const client = new UseAIClient(myTransport); +``` diff --git a/packages/client/src/client.test.ts b/packages/client/src/client.test.ts index 5036feee..d7274c80 100644 --- a/packages/client/src/client.test.ts +++ b/packages/client/src/client.test.ts @@ -1,16 +1,26 @@ import { describe, test, expect, mock, beforeEach, afterEach, spyOn } from 'bun:test'; import type { Socket } from 'socket.io-client'; +import type { UseAIClientMessage } from './types'; +import type { WebSocketLike } from './transport/WebSocketTransport'; -// Store event handlers registered via socket.on() -let eventHandlers: Record = {}; +/** + * The UseAIClient suite, run over every bundled transport. + * + * Each harness drives its transport at that transport's own wire level, so the + * two are held to one behaviour. Wiring specific to a single transport is tested + * next to it, in transport/*.test.ts. + */ + +// ── Socket.IO harness ─────────────────────────────────────────────────────── + +let socketHandlers: Record = {}; let mockSocket: Partial & { connected: boolean }; function createMockSocket() { - eventHandlers = {}; + socketHandlers = {}; mockSocket = { on: mock((event: string, handler: Function) => { - if (!eventHandlers[event]) eventHandlers[event] = []; - eventHandlers[event].push(handler); + (socketHandlers[event] ??= []).push(handler); return mockSocket as Socket; }), emit: mock(() => mockSocket as Socket), @@ -21,245 +31,235 @@ function createMockSocket() { transport: { name: 'polling' }, on: mock(), }, - } as any, + } as never, }; return mockSocket as Socket; } -// Helper to emit socket events in tests -function emitSocketEvent(event: string, ...args: any[]) { - eventHandlers[event]?.forEach(handler => handler(...args)); -} - -// Mock socket.io-client module mock.module('socket.io-client', () => ({ io: () => createMockSocket(), })); -// Import after mocking +// ── Plain WebSocket harness ───────────────────────────────────────────────── + +class FakeWebSocket implements WebSocketLike { + sent: string[] = []; + onopen: ((event: unknown) => void) | null = null; + onmessage: ((event: { data: unknown }) => void) | null = null; + onclose: ((event: unknown) => void) | null = null; + onerror: ((event: unknown) => void) | null = null; + + send(data: string): void { + this.sent.push(data); + } + + close(): void {} +} + +// Imported after the module mock so SocketIOTransport picks it up. const { UseAIClient } = await import('./client'); +const { SocketIOTransport } = await import('./transport/SocketIOTransport'); +const { WebSocketTransport } = await import('./transport/WebSocketTransport'); + +type Client = InstanceType; + +/** Controls over the wire beneath a connected client. */ +interface Harness { + client: Client; + /** + * Simulates the connection being accepted. After a drop, waits for the + * transport's own reconnection first, so the two transports read the same + * in a test even though only one of them owns its retry loop. + */ + open(): Promise; + /** Simulates the connection dropping. */ + close(reason: string): void; + /** Simulates a named payload arriving from the server. */ + deliver(name: 'event' | 'agents' | 'config', data: unknown): void; + /** Messages the client has sent upstream, oldest first. */ + sent(): UseAIClientMessage[]; +} + +async function waitUntil(condition: () => boolean, timeoutMs = 500): Promise { + const deadline = Date.now() + timeoutMs; + while (!condition()) { + if (Date.now() > deadline) throw new Error('Timed out waiting for the transport to reconnect'); + await new Promise(resolve => setTimeout(resolve, 1)); + } +} + +const HARNESSES: Array<[string, () => Harness]> = [ + [ + 'SocketIOTransport', + () => { + const client = new UseAIClient(new SocketIOTransport('http://localhost:8081')); + client.connect(); + const socket = mockSocket; + const fire = (event: string, ...args: unknown[]) => + socketHandlers[event]?.forEach(handler => handler(...args)); + + return { + client, + async open() { + socket.connected = true; + fire('connect'); + }, + close(reason: string) { + socket.connected = false; + fire('disconnect', reason); + }, + deliver(name, data) { + fire(name, data); + }, + sent() { + const emit = socket.emit as unknown as { mock: { calls: unknown[][] } }; + return emit.mock.calls + .filter(call => call[0] === 'message') + .map(call => call[1] as UseAIClientMessage); + }, + }; + }, + ], + [ + 'WebSocketTransport', + () => { + let socket!: FakeWebSocket; + const client = new UseAIClient( + new WebSocketTransport('wss://localhost:8081/ws', { + reconnectionDelay: 1, + reconnectionDelayMax: 1, + createWebSocket: () => (socket = new FakeWebSocket()), + }), + ); + client.connect(); + + return { + client, + async open() { + const dropped = socket; + // A closed socket is detached; the replacement arrives on the backoff timer. + if (dropped.onopen === null) { + await waitUntil(() => socket !== dropped); + } + socket.onopen?.({}); + }, + close(reason: string) { + // The reason a plain WebSocket reports comes from the close frame, not + // from the caller; the client only logs it. + void reason; + socket.onclose?.({}); + }, + deliver(name, data) { + socket.onmessage?.({ data: JSON.stringify({ name, data }) }); + }, + sent() { + return socket.sent.map(frame => JSON.parse(frame) as UseAIClientMessage); + }, + }; + }, + ], +]; -describe('UseAIClient', () => { +describe.each(HARNESSES)('UseAIClient over %s', (_name, createHarness) => { let consoleLogSpy: ReturnType; let consoleWarnSpy: ReturnType; + let harness: Harness; + let client: Client; + + /** The last message the client sent upstream. */ + const lastSent = () => harness.sent()[harness.sent().length - 1]; + const emitEvent = (event: Record) => harness.deliver('event', event); beforeEach(() => { consoleLogSpy = spyOn(console, 'log').mockImplementation(() => {}); consoleWarnSpy = spyOn(console, 'warn').mockImplementation(() => {}); + // Never reconnects during a test: the harness owns connection lifecycle. + harness = createHarness(); + client = harness.client; }); afterEach(() => { + client.disconnect(); consoleLogSpy.mockRestore(); consoleWarnSpy.mockRestore(); }); - describe('connect()', () => { - test('notifies connected state on successful connection', () => { - const client = new UseAIClient('http://localhost:8081'); + describe('connection state', () => { + test('notifies connected state on successful connection', async () => { const stateChanges: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges.push(connected)); - client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); - - client.connect(); - - // Simulate successful connection - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); // Initial state (false) + connected (true) expect(stateChanges).toEqual([false, true]); }); - test('notifies disconnected state on disconnect', () => { - const client = new UseAIClient('http://localhost:8081'); + test('notifies disconnected state on disconnect', async () => { const stateChanges: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges.push(connected)); - client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); - - client.connect(); - - // Connect first - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Then disconnect - mockSocket.connected = false; - emitSocketEvent('disconnect', 'transport close'); + await harness.open(); + harness.close('transport close'); expect(stateChanges).toEqual([false, true, false]); }); - test('logs warning on connection error without throwing', () => { - const client = new UseAIClient('http://localhost:8081'); - - client.connect(); - - // Simulate connection error - emitSocketEvent('connect_error', new Error('Connection refused')); - - // Should use console.warn, not console.error - expect(consoleWarnSpy).toHaveBeenCalledWith( - '[UseAI] Connection error:', - 'Connection refused' - ); - }); - }); - - describe('reconnection scenarios', () => { - test('reconnects successfully after 2 failed attempts', () => { - const client = new UseAIClient('http://localhost:8081'); + test('reconnects after disconnect', async () => { const stateChanges: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges.push(connected)); - client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); - - client.connect(); - - // 1st attempt: connection error - emitSocketEvent('connect_error', new Error('Attempt 1 failed')); - - // 2nd attempt: connection error - emitSocketEvent('connect_error', new Error('Attempt 2 failed')); - - // 3rd attempt: success - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Initial (false) + connected (true) - // Connection errors don't change state, only connect/disconnect events do - expect(stateChanges).toEqual([false, true]); - expect(consoleWarnSpy).toHaveBeenCalledTimes(2); - }); - - test('reconnects after disconnect', () => { - const client = new UseAIClient('http://localhost:8081'); - const stateChanges: boolean[] = []; - - client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); - - client.connect(); - - // Initial connection - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Disconnect - mockSocket.connected = false; - emitSocketEvent('disconnect', 'transport close'); - - // Reconnect after 1 failed attempt - emitSocketEvent('connect_error', new Error('Reconnect attempt 1 failed')); - - // Successful reconnection - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); + harness.close('transport close'); + await harness.open(); // false (initial) -> true (connect) -> false (disconnect) -> true (reconnect) expect(stateChanges).toEqual([false, true, false, true]); }); - test('handles multiple disconnect/reconnect cycles', () => { - const client = new UseAIClient('http://localhost:8081'); + test('handles multiple disconnect/reconnect cycles', async () => { const stateChanges: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges.push(connected)); - client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); - - client.connect(); - - // Cycle 1: connect - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Cycle 1: disconnect - mockSocket.connected = false; - emitSocketEvent('disconnect', 'ping timeout'); - - // Cycle 2: reconnect - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Cycle 2: disconnect - mockSocket.connected = false; - emitSocketEvent('disconnect', 'transport error'); - - // Cycle 3: reconnect - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); + harness.close('ping timeout'); + await harness.open(); + harness.close('transport error'); + await harness.open(); expect(stateChanges).toEqual([false, true, false, true, false, true]); }); }); describe('onConnectionStateChange()', () => { - test('immediately notifies current state on subscribe', () => { - const client = new UseAIClient('http://localhost:8081'); + test('immediately notifies current state on subscribe', async () => { const stateChanges: boolean[] = []; - client.connect(); - - // Connect first - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Subscribe after connection - should immediately get current state - client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); + await harness.open(); + client.onConnectionStateChange(connected => stateChanges.push(connected)); expect(stateChanges).toEqual([true]); }); - test('unsubscribe stops notifications', () => { - const client = new UseAIClient('http://localhost:8081'); + test('unsubscribe stops notifications', async () => { const stateChanges: boolean[] = []; + const unsubscribe = client.onConnectionStateChange(connected => stateChanges.push(connected)); - const unsubscribe = client.onConnectionStateChange((connected) => { - stateChanges.push(connected); - }); - - client.connect(); - - // Connect - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Unsubscribe + await harness.open(); unsubscribe(); - - // Disconnect - should not be notified - mockSocket.connected = false; - emitSocketEvent('disconnect', 'transport close'); + harness.close('transport close'); // Only initial (false) + connect (true), no disconnect notification expect(stateChanges).toEqual([false, true]); }); - test('supports multiple subscribers', () => { - const client = new UseAIClient('http://localhost:8081'); + test('supports multiple subscribers', async () => { const stateChanges1: boolean[] = []; const stateChanges2: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges1.push(connected)); + client.onConnectionStateChange(connected => stateChanges2.push(connected)); - client.onConnectionStateChange((connected) => { - stateChanges1.push(connected); - }); - - client.onConnectionStateChange((connected) => { - stateChanges2.push(connected); - }); - - client.connect(); - - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); expect(stateChanges1).toEqual([false, true]); expect(stateChanges2).toEqual([false, true]); @@ -267,141 +267,146 @@ describe('UseAIClient', () => { }); describe('isConnected()', () => { - test('returns false before connect', () => { - const client = new UseAIClient('http://localhost:8081'); + test('returns false before the connection is accepted', () => { expect(client.isConnected()).toBe(false); }); - test('returns true when connected', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); + test('returns true when connected', async () => { + await harness.open(); + expect(client.isConnected()).toBe(true); + }); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('returns false after disconnect', async () => { + await harness.open(); + harness.close('transport close'); + expect(client.isConnected()).toBe(false); + }); + }); - expect(client.isConnected()).toBe(true); + describe('server payloads', () => { + test('agents payload updates the available agents', async () => { + const received: Array<[unknown, unknown]> = []; + client.onAgentsChange((agents, defaultAgent) => received.push([agents, defaultAgent])); + + await harness.open(); + harness.deliver('agents', { + agents: [{ id: 'claude', name: 'Claude' }], + defaultAgent: 'claude', + }); + + expect(client.availableAgents).toEqual([{ id: 'claude', name: 'Claude' }]); + expect(client.defaultAgent).toBe('claude'); + expect(received[received.length - 1]).toEqual([[{ id: 'claude', name: 'Claude' }], 'claude']); }); - test('returns false after disconnect', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); + test('config payload updates the Langfuse flag', async () => { + const received: boolean[] = []; + client.onLangfuseConfigChange(enabled => received.push(enabled)); + + await harness.open(); + harness.deliver('config', { langfuseEnabled: true }); - mockSocket.connected = true; - emitSocketEvent('connect'); + expect(received).toEqual([false, true]); + }); - mockSocket.connected = false; - emitSocketEvent('disconnect', 'transport close'); + test('submitFeedback sends feedback once Langfuse is enabled', async () => { + await harness.open(); + harness.deliver('config', { langfuseEnabled: true }); - expect(client.isConnected()).toBe(false); + client.submitFeedback('msg-1', 'trace-1', 'upvote'); + + expect(lastSent()).toEqual({ + type: 'message_feedback', + data: { messageId: 'msg-1', traceId: 'trace-1', feedback: 'upvote' }, + }); + }); + + test('submitFeedback is a no-op while disconnected', () => { + client.submitFeedback('msg-1', 'trace-1', 'upvote'); + + expect(harness.sent()).toEqual([]); + expect(consoleWarnSpy).toHaveBeenCalledWith('[UseAI] Cannot submit feedback: not connected'); }); }); - describe('sendPrompt()', () => { + describe('sendPrompt()', async () => { test('sends message without forwardedProps when not provided', async () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); await client.sendPrompt('Hello'); - expect(mockSocket.emit).toHaveBeenCalledWith('message', expect.objectContaining({ - type: 'run_agent', - data: expect.objectContaining({ - forwardedProps: {}, - }), - })); + const message = lastSent(); + expect(message.type).toBe('run_agent'); + expect((message.data as { forwardedProps: unknown }).forwardedProps).toEqual({}); }); test('sends message with forwardedProps when provided', async () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); await client.sendPrompt('Hello', undefined, { mcpHeaders: { - 'https://api.example.com': { headers: { 'Authorization': 'Bearer token' } }, + 'https://api.example.com': { headers: { Authorization: 'Bearer token' } }, }, telemetryMetadata: { userId: 'user-123', evaluationId: 'eval-456' }, - }); - expect(mockSocket.emit).toHaveBeenCalledWith('message', expect.objectContaining({ - type: 'run_agent', - data: expect.objectContaining({ - forwardedProps: { - mcpHeaders: { - 'https://api.example.com': { headers: { 'Authorization': 'Bearer token' } }, - }, - telemetryMetadata: { userId: 'user-123', evaluationId: 'eval-456' }, - }, - }), - })); + const message = lastSent(); + expect(message.type).toBe('run_agent'); + expect((message.data as { forwardedProps: unknown }).forwardedProps).toEqual({ + mcpHeaders: { + 'https://api.example.com': { headers: { Authorization: 'Bearer token' } }, + }, + telemetryMetadata: { userId: 'user-123', evaluationId: 'eval-456' }, + }); }); test('merges forwardedProps with selected agent', async () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - - mockSocket.connected = true; - emitSocketEvent('connect'); - - // Set selected agent + await harness.open(); client.setAgent('claude-opus'); await client.sendPrompt('Hello', undefined, { telemetryMetadata: { userId: 'user-123' }, }); - expect(mockSocket.emit).toHaveBeenCalledWith('message', expect.objectContaining({ - type: 'run_agent', - data: expect.objectContaining({ - forwardedProps: { - agent: 'claude-opus', - telemetryMetadata: { userId: 'user-123' }, - }, - }), - })); + const message = lastSent(); + expect(message.type).toBe('run_agent'); + expect((message.data as { forwardedProps: unknown }).forwardedProps).toEqual({ + agent: 'claude-opus', + telemetryMetadata: { userId: 'user-123' }, + }); }); }); - describe('message ordering after tool call turn', () => { - function simulateToolCallTurn(client: InstanceType) { + describe('message ordering after tool call turn', async () => { + function simulateToolCallTurn() { // Simulate sending a user message client.sendPrompt('Add a todo: buy groceries'); // Simulate server events for a tool call turn - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'toolu_123', toolCallName: 'addTodo' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'toolu_123', delta: '{"text":"buy groceries"}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'toolu_123' }); + emitEvent({ type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'toolu_123', toolCallName: 'addTodo' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'toolu_123', delta: '{"text":"buy groceries"}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'toolu_123' }); // Client executes tool and sends result client.sendToolResponse('toolu_123', { success: true, message: 'Todo added' }); // Server emits STEP_FINISHED after the tool-call step (the real AISDKAgent // emits one per step), which flushes assistant(toolCalls) + tool_result. - emitSocketEvent('event', { type: 'STEP_FINISHED' }); + emitEvent({ type: 'STEP_FINISHED' }); // Server sends final text response - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: "I've added the todo!" }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: "I've added the todo!" }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); - // RUN_FINISHED - emitSocketEvent('event', { type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); } - test('messages are in correct API order: user → assistant(toolCalls) → tool → assistant(text)', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('messages are in correct API order: user → assistant(toolCalls) → tool → assistant(text)', async () => { + await harness.open(); - simulateToolCallTurn(client); + simulateToolCallTurn(); const messages = client.messages; @@ -424,23 +429,19 @@ describe('UseAIClient', () => { expect((assistantText as Record).toolCalls).toBeUndefined(); }); - test('TOOL_CALL_RESULT stores server-side tool result in conversation history', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('TOOL_CALL_RESULT stores server-side tool result in conversation history', async () => { + await harness.open(); - // Simulate sending a user message client.sendPrompt('What is the weather in Tokyo?'); // Server-side tool call (MCP tool) — client does NOT call sendToolResponse - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'toolu_mcp_1', toolCallName: 'mcp_get_weather' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'toolu_mcp_1', delta: '{"location":"Tokyo"}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'toolu_mcp_1' }); + emitEvent({ type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'toolu_mcp_1', toolCallName: 'mcp_get_weather' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'toolu_mcp_1', delta: '{"location":"Tokyo"}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'toolu_mcp_1' }); // Server sends the actual MCP tool result via TOOL_CALL_RESULT - emitSocketEvent('event', { + emitEvent({ type: 'TOOL_CALL_RESULT', messageId: 'msg-result-1', toolCallId: 'toolu_mcp_1', @@ -449,13 +450,13 @@ describe('UseAIClient', () => { }); // STEP_FINISHED flushes the tool-call step (assistant(toolCalls) + result). - emitSocketEvent('event', { type: 'STEP_FINISHED' }); + emitEvent({ type: 'STEP_FINISHED' }); // Server sends final text response - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: 'It is 15°C and cloudy in Tokyo.' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: 'It is 15°C and cloudy in Tokyo.' }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); + emitEvent({ type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); const messages = client.messages; expect(messages).toHaveLength(4); // user, assistant(toolCalls), tool, assistant(text) @@ -466,24 +467,21 @@ describe('UseAIClient', () => { expect((toolResult as Record).toolCallId).toBe('toolu_mcp_1'); }); - test('TOOL_CALL_RESULT does not duplicate result for client-side tools', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('TOOL_CALL_RESULT does not duplicate result for client-side tools', async () => { + await harness.open(); client.sendPrompt('Add a todo: buy groceries'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'toolu_client_1', toolCallName: 'addTodo' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'toolu_client_1', delta: '{"text":"buy groceries"}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'toolu_client_1' }); + emitEvent({ type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'toolu_client_1', toolCallName: 'addTodo' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'toolu_client_1', delta: '{"text":"buy groceries"}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'toolu_client_1' }); // Client executes tool and sends result (this pushes to _pendingToolResults) client.sendToolResponse('toolu_client_1', { success: true }); // Server also sends TOOL_CALL_RESULT for the same toolCallId (should be deduplicated) - emitSocketEvent('event', { + emitEvent({ type: 'TOOL_CALL_RESULT', messageId: 'msg-result-dup', toolCallId: 'toolu_client_1', @@ -492,102 +490,77 @@ describe('UseAIClient', () => { }); // STEP_FINISHED flushes the tool-call step (assistant(toolCalls) + result). - emitSocketEvent('event', { type: 'STEP_FINISHED' }); + emitEvent({ type: 'STEP_FINISHED' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: 'Done!' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: 'Done!' }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); + emitEvent({ type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); const messages = client.messages; const toolResults = messages.filter(m => m.role === 'tool'); expect(toolResults).toHaveLength(1); // Only one, not duplicated }); - test('abortRun emits abort_run with the in-flight runId from sendPrompt', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('abortRun sends abort_run with the in-flight runId from sendPrompt', async () => { + await harness.open(); client.sendPrompt('Hello'); const runId = client.currentRunId; expect(typeof runId).toBe('string'); - // Clear prior emit calls (run_agent) for a clean assertion. - const emitMock = mockSocket.emit as ReturnType; - emitMock.mockClear(); - client.abortRun(); - expect(emitMock).toHaveBeenCalledWith('message', { - type: 'abort_run', - data: { runId }, - }); + expect(lastSent()).toEqual({ type: 'abort_run', data: { runId } }); }); - test('abortRun is a no-op when no run is in flight', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); - - const emitMock = mockSocket.emit as ReturnType; - emitMock.mockClear(); + test('abortRun is a no-op when no run is in flight', async () => { + await harness.open(); client.abortRun(); - expect(emitMock).not.toHaveBeenCalled(); + expect(harness.sent()).toEqual([]); }); - test('currentRunId is cleared after RUN_FINISHED', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('currentRunId is cleared after RUN_FINISHED', async () => { + await harness.open(); client.sendPrompt('Hello'); expect(client.currentRunId).not.toBeNull(); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'm' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'm' }); - emitSocketEvent('event', { type: 'RUN_FINISHED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'm' }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'm' }); + emitEvent({ type: 'RUN_FINISHED', threadId: 't', runId: 'r' }); expect(client.currentRunId).toBeNull(); }); - test('currentRunId is cleared after RUN_ERROR', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('currentRunId is cleared after RUN_ERROR', async () => { + await harness.open(); client.sendPrompt('Hello'); expect(client.currentRunId).not.toBeNull(); - emitSocketEvent('event', { type: 'RUN_ERROR', message: 'ABORTED' }); + emitEvent({ type: 'RUN_ERROR', message: 'ABORTED' }); expect(client.currentRunId).toBeNull(); }); - test('stopping while a tool is still running backfills results for the unanswered tool calls', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('stopping while a tool is still running backfills results for the unanswered tool calls', async () => { + await harness.open(); client.sendPrompt('Add two todos'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r' }); // Two tool_use blocks streamed, but the client never responds for the // second one (aborted mid-execution). - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'addTodo' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{"text":"a"}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'tc1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'tc2', toolCallName: 'addTodo' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'tc2', delta: '{"text":"b"}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'tc2' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'addTodo' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{"text":"a"}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'tc1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'tc2', toolCallName: 'addTodo' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'tc2', delta: '{"text":"b"}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'tc2' }); // Only the first tool got a result before abort. client.sendToolResponse('tc1', { ok: 1 }); @@ -611,18 +584,15 @@ describe('UseAIClient', () => { expect(JSON.parse(synthetic!.content)).toMatchObject({ aborted: true }); }); - test('stopping mid-TOOL_CALL_ARGS (before TOOL_CALL_END) leaves no orphaned tool_use in history', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('stopping mid-TOOL_CALL_ARGS (before TOOL_CALL_END) leaves no orphaned tool_use in history', async () => { + await harness.open(); client.sendPrompt('Add a todo'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'addTodo' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'addTodo' }); // Partial args delta — TOOL_CALL_END never arrives before abort. - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{"text":"buy' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{"text":"buy' }); client.finalizeRun({ aborted: true }); @@ -643,28 +613,24 @@ describe('UseAIClient', () => { // finalizeAbortedRun clears in-progress reasoning blocks. expect(client.currentReasoningBlocks).toEqual([]); - }); - test('does not duplicate the step text in _messages after a STEP_FINISHED → ABORT sequence', () => { + test('does not duplicate the step text in _messages after a STEP_FINISHED → ABORT sequence', async () => { // Regression: aborting between STEP_FINISHED and the next TEXT_MESSAGE_START used to save the step text twice. - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); client.sendPrompt('test prompt'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r' }); - emitSocketEvent('event', { type: 'STEP_STARTED', stepName: 'step-0' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'm1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'm1', delta: 'step text' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'm1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'testTool' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'tc1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_RESULT', messageId: 'tr1', toolCallId: 'tc1', content: '[]', role: 'tool' }); - emitSocketEvent('event', { type: 'STEP_FINISHED', stepName: 'step-0' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'STEP_STARTED', stepName: 'step-0' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'm1' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'm1', delta: 'step text' }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'm1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'testTool' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'tc1' }); + emitEvent({ type: 'TOOL_CALL_RESULT', messageId: 'tr1', toolCallId: 'tc1', content: '[]', role: 'tool' }); + emitEvent({ type: 'STEP_FINISHED', stepName: 'step-0' }); client.finalizeRun({ aborted: true }); @@ -676,29 +642,26 @@ describe('UseAIClient', () => { expect((textMatches[0].toolCalls as Array<{ id: string }>)[0].id).toBe('tc1'); }); - test('a tool-only aborted run does not leak the previous run\'s text', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test("a tool-only aborted run does not leak the previous run's text", async () => { + await harness.open(); // Run 1: ends with a final text answer. client.sendPrompt('list the tools'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'm1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'm1', delta: 'Here are the tools' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'm1' }); - emitSocketEvent('event', { type: 'RUN_FINISHED', threadId: 't', runId: 'r1' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r1' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'm1' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'm1', delta: 'Here are the tools' }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'm1' }); + emitEvent({ type: 'RUN_FINISHED', threadId: 't', runId: 'r1' }); // Run 2: tool-only step (no TEXT_MESSAGE_START), aborted mid-execution. // RUN_STARTED must reset the leftover text so it is not persisted again. client.sendPrompt('wait 5 seconds'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r2' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r2' }); expect(client.currentMessageContent).toBe(''); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'wait' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{"seconds":5}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'tc1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'wait' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{"seconds":5}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'tc1' }); client.finalizeRun({ aborted: true }); // No assistant message carries the run-1 text after run 2's abort. @@ -708,17 +671,14 @@ describe('UseAIClient', () => { expect(leaked).toHaveLength(1); // only the legitimate run-1 message }); - test('stopping while the assistant is streaming text keeps the partial text', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('stopping while the assistant is streaming text keeps the partial text', async () => { + await harness.open(); client.sendPrompt('Tell me a story'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'm' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'm', delta: 'Once upon a' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'm' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'm', delta: 'Once upon a' }); // Note: no TEXT_MESSAGE_END — mid-stream abort. client.finalizeRun({ aborted: true }); @@ -729,33 +689,30 @@ describe('UseAIClient', () => { expect(msgs[1].content).toBe('Once upon a'); }); - test('stopping while the assistant is reasoning drops the unfinished thinking but keeps earlier steps\' reasoning', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test("stopping while the assistant is reasoning drops the unfinished thinking but keeps earlier steps' reasoning", async () => { + await harness.open(); client.sendPrompt('Do two things'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 't', runId: 'r' }); + emitEvent({ type: 'RUN_STARTED', threadId: 't', runId: 'r' }); // Step 1: complete reasoning + tool_use + STEP_FINISHED. Reasoning gets // attached to the step-1 assistant message and survives the abort. - emitSocketEvent('event', { type: 'REASONING_MESSAGE_START', messageId: 'rm1' }); - emitSocketEvent('event', { type: 'REASONING_MESSAGE_CONTENT', messageId: 'rm1', delta: 'think 1' }); - emitSocketEvent('event', { type: 'REASONING_MESSAGE_END', messageId: 'rm1' }); - emitSocketEvent('event', { type: 'REASONING_ENCRYPTED_VALUE', subtype: 'message', encryptedValue: 'sig1' }); - emitSocketEvent('event', { type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'doThing' }); - emitSocketEvent('event', { type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{}' }); - emitSocketEvent('event', { type: 'TOOL_CALL_END', toolCallId: 'tc1' }); + emitEvent({ type: 'REASONING_MESSAGE_START', messageId: 'rm1' }); + emitEvent({ type: 'REASONING_MESSAGE_CONTENT', messageId: 'rm1', delta: 'think 1' }); + emitEvent({ type: 'REASONING_MESSAGE_END', messageId: 'rm1' }); + emitEvent({ type: 'REASONING_ENCRYPTED_VALUE', subtype: 'message', encryptedValue: 'sig1' }); + emitEvent({ type: 'TOOL_CALL_START', toolCallId: 'tc1', toolCallName: 'doThing' }); + emitEvent({ type: 'TOOL_CALL_ARGS', toolCallId: 'tc1', delta: '{}' }); + emitEvent({ type: 'TOOL_CALL_END', toolCallId: 'tc1' }); client.sendToolResponse('tc1', { ok: 1 }); - emitSocketEvent('event', { type: 'STEP_FINISHED' }); + emitEvent({ type: 'STEP_FINISHED' }); // Step 2: reasoning streamed but END/encrypted not received before abort. - emitSocketEvent('event', { type: 'REASONING_MESSAGE_START', messageId: 'rm2' }); - emitSocketEvent('event', { type: 'REASONING_MESSAGE_CONTENT', messageId: 'rm2', delta: 'think 2 partial' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'm2' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'm2', delta: 'About to' }); + emitEvent({ type: 'REASONING_MESSAGE_START', messageId: 'rm2' }); + emitEvent({ type: 'REASONING_MESSAGE_CONTENT', messageId: 'rm2', delta: 'think 2 partial' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'm2' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'm2', delta: 'About to' }); client.finalizeRun({ aborted: true }); @@ -776,19 +733,16 @@ describe('UseAIClient', () => { expect(client.currentReasoningBlocks).toEqual([]); }); - test('simple text response (no tool calls) still works', () => { - const client = new UseAIClient('http://localhost:8081'); - client.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + test('simple text response (no tool calls) still works', async () => { + await harness.open(); client.sendPrompt('Hello'); - emitSocketEvent('event', { type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: 'Hi there!' }); - emitSocketEvent('event', { type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); - emitSocketEvent('event', { type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'RUN_STARTED', threadId: 'thread-1', runId: 'run-1' }); + emitEvent({ type: 'TEXT_MESSAGE_START', messageId: 'msg-1' }); + emitEvent({ type: 'TEXT_MESSAGE_CONTENT', messageId: 'msg-1', delta: 'Hi there!' }); + emitEvent({ type: 'TEXT_MESSAGE_END', messageId: 'msg-1' }); + emitEvent({ type: 'RUN_FINISHED', threadId: 'thread-1', runId: 'run-1' }); const messages = client.messages; expect(messages).toHaveLength(2); @@ -798,3 +752,38 @@ describe('UseAIClient', () => { }); }); }); + +describe('UseAIClient construction', () => { + test('a server URL builds a Socket.IO transport', () => { + const client = new UseAIClient('http://localhost:8081'); + client.connect(); + + expect(mockSocket).toBeDefined(); + expect(client.isConnected()).toBe(false); + + mockSocket.connected = true; + expect(client.isConnected()).toBe(true); + + client.disconnect(); + }); + + test('disconnect() unsubscribes from the transport', () => { + let socket!: FakeWebSocket; + const client = new UseAIClient( + new WebSocketTransport('wss://localhost:8081/ws', { + createWebSocket: () => (socket = new FakeWebSocket()), + }), + ); + const stateChanges: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges.push(connected)); + client.connect(); + socket.onopen?.({}); + + client.disconnect(); + // The transport is closed, but a late frame from the old socket must not + // reach a client that has stopped listening. + socket.onclose?.({}); + + expect(stateChanges).toEqual([false, true]); + }); +}); diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index c8f1fa00..893cd270 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -1,4 +1,3 @@ -import { io, Socket } from 'socket.io-client'; import { EventType } from '@meetsmore-oss/use-ai-core'; import type { ToolDefinition, @@ -23,6 +22,8 @@ import type { ReasoningEncryptedValueEvent, ReasoningPart, } from './types'; +import { SocketIOTransport } from './transport/SocketIOTransport'; +import type { UseAITransport } from './transport/types'; import { v4 as uuidv4 } from 'uuid'; /** @@ -45,11 +46,11 @@ export type ToolCallHandler = ( ) => void; /** - * Socket.IO client for communicating with the UseAI server. + * Client for communicating with the UseAI server. * Uses the AG-UI protocol (https://docs.ag-ui.com/), so will be compatible with other AG-UI compliant servers. * * Handles: - * - Connection management and automatic reconnection + * - Connection management, via a {@link UseAITransport} * - Sending RunAgentInput messages to server * - Parsing AG-UI event streams from server * - Tool execution coordination @@ -57,14 +58,9 @@ export type ToolCallHandler = ( * You probably don't need to use this directly, instead use {@link UseAIProvider}. */ export class UseAIClient { - private socket: Socket | null = null; + private transport: UseAITransport; + private transportUnsubscribes: Array<() => void> = []; private eventHandlers: Map = new Map(); - // Reconnect indefinitely so clients recover after extended outages (mobile - // app backgrounded long enough for server pingTimeout, airplane mode, etc.). - // Socket.IO applies exponential backoff capped at reconnectionDelayMax, - // so steady-state retry frequency is ~one attempt per 10s. - private reconnectDelay = 1000; - private reconnectDelayMax = 10_000; // Session state private _threadId: string | null = null; @@ -111,81 +107,62 @@ export class UseAIClient { /** * Creates a new UseAI client instance. * - * @param serverUrl - The WebSocket URL of the UseAI server + * @param target - The URL of the UseAI server, which is reached over Socket.IO, or a + * {@link UseAITransport} to reach it over something else. + * @example + * ```typescript + * new UseAIClient('wss://your-server.com'); + * new UseAIClient(new WebSocketTransport('wss://your-server.com/ws')); + * ``` */ - constructor(private serverUrl: string) {} + constructor(target: string | UseAITransport) { + this.transport = typeof target === 'string' ? new SocketIOTransport(target) : target; + } /** - * Establishes a Socket.IO connection to the server. + * Opens the transport's connection to the server. * Connection state changes are notified via onConnectionStateChange(). - * Socket.IO handles reconnection automatically. + * Reconnection is the transport's responsibility and is automatic for both bundled transports. */ connect(): void { - this.socket = io(this.serverUrl, { - transports: ['polling', 'websocket'], - reconnection: true, - reconnectionAttempts: Infinity, - reconnectionDelay: this.reconnectDelay, - reconnectionDelayMax: this.reconnectDelayMax, - withCredentials: true, - }); - - this.socket.on('connect', () => { - console.log('[UseAI] Connected to server'); - console.log('[UseAI] Transport:', this.socket?.io?.engine?.transport?.name); - - // Listen for transport upgrades (only if engine is available) - const engine = this.socket?.io?.engine; - if (engine) { - engine.on('upgrade', (transport: { name: string }) => { - console.log('[UseAI] Upgraded to transport:', transport.name); - }); - - engine.on('upgradeError', (err: { message: string }) => { - console.warn('[UseAI] Upgrade error:', err.message); - }); - } - - // Notify connection state handlers - this.connectionStateHandlers.forEach(handler => handler(true)); - }); + this.transportUnsubscribes.push( + this.transport.on('connect', () => { + console.log('[UseAI] Connected to server'); + this.connectionStateHandlers.forEach(handler => handler(true)); + }), + + this.transport.on('event', (data) => { + const aguiEvent = data as AGUIEvent; + try { + console.log('[Client] Received event:', aguiEvent.type); + this.handleEvent(aguiEvent); + } catch (error) { + console.error('[UseAI] Error handling event:', error); + } + }), - this.socket.on('event', (aguiEvent: AGUIEvent) => { - try { - console.log('[Client] Received event:', aguiEvent.type); - this.handleEvent(aguiEvent); - } catch (error) { - console.error('[UseAI] Error handling event:', error); - } - }); + this.transport.on('agents', (data) => { + const { agents, defaultAgent } = data as { agents: AgentInfo[]; defaultAgent: string }; + console.log('[Client] Received available agents:', data); + this._availableAgents = agents; + this._defaultAgent = defaultAgent; + this.agentsChangeHandlers.forEach(handler => handler(agents, defaultAgent)); + }), - // Listen for available agents from server - this.socket.on('agents', (data: { agents: AgentInfo[]; defaultAgent: string }) => { - console.log('[Client] Received available agents:', data); - this._availableAgents = data.agents; - this._defaultAgent = data.defaultAgent; - // Notify listeners - this.agentsChangeHandlers.forEach(handler => handler(data.agents, data.defaultAgent)); - }); + this.transport.on('config', (data) => { + const { langfuseEnabled } = data as { langfuseEnabled?: boolean }; + console.log('[Client] Received server config:', data); + this._langfuseEnabled = langfuseEnabled ?? false; + this.langfuseConfigHandlers.forEach(handler => handler(this._langfuseEnabled)); + }), - // Listen for server config (including Langfuse enabled status) - this.socket.on('config', (data: { langfuseEnabled?: boolean }) => { - console.log('[Client] Received server config:', data); - this._langfuseEnabled = data.langfuseEnabled ?? false; - // Notify listeners - this.langfuseConfigHandlers.forEach(handler => handler(this._langfuseEnabled)); - }); + this.transport.on('disconnect', (reason) => { + console.log('[UseAI] Disconnected:', reason); + this.connectionStateHandlers.forEach(handler => handler(false)); + }), + ); - this.socket.on('connect_error', (error) => { - // Use warn instead of error to avoid triggering Next.js error overlay - console.warn('[UseAI] Connection error:', error.message); - }); - - this.socket.on('disconnect', (reason) => { - console.log('[UseAI] Disconnected:', reason); - // Notify connection state handlers - this.connectionStateHandlers.forEach(handler => handler(false)); - }); + this.transport.connect(); } @@ -869,21 +846,20 @@ export class UseAIClient { } send(message: UseAIClientMessage) { - if (this.socket && this.socket.connected) { - this.socket.emit('message', message); + if (this.transport.connected) { + this.transport.send(message); } else { - console.error('Socket.IO is not connected'); + console.error('[UseAI] Not connected to server'); } } /** - * Closes the Socket.IO connection to the server. + * Closes the connection to the server and unsubscribes from the transport. */ disconnect() { - if (this.socket) { - this.socket.disconnect(); - this.socket = null; - } + this.transportUnsubscribes.forEach(unsubscribe => unsubscribe()); + this.transportUnsubscribes = []; + this.transport.disconnect(); } /** @@ -892,7 +868,7 @@ export class UseAIClient { * @returns true if connected, false otherwise */ isConnected(): boolean { - return this.socket !== null && this.socket.connected; + return this.transport.connected; } /** @@ -919,7 +895,7 @@ export class UseAIClient { * @param feedback - 'upvote' for positive, 'downvote' for negative, null to remove */ submitFeedback(messageId: string, traceId: string, feedback: FeedbackValue): void { - if (!this.socket?.connected) { + if (!this.transport.connected) { console.warn('[UseAI] Cannot submit feedback: not connected'); return; } diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts index 6caed550..e21f1ada 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -2,6 +2,14 @@ export { useAI } from './useAI'; export { useAIWorkflow } from './useAIWorkflow'; export { UseAIProvider, useAIContext } from './providers/useAIProvider'; export { UseAIClient } from './client'; +export { SocketIOTransport, WebSocketTransport } from './transport'; +export type { + UseAITransport, + UseAITransportEventName, + SocketIOTransportOptions, + WebSocketTransportOptions, + WebSocketLike, +} from './transport'; export { defineTool, executeDefinedTool, convertToolsToDefinitions } from './defineTool'; /** @hidden */ export { z } from 'zod'; diff --git a/packages/client/src/providers/useAIProvider.tsx b/packages/client/src/providers/useAIProvider.tsx index b7b74739..2aa7637b 100644 --- a/packages/client/src/providers/useAIProvider.tsx +++ b/packages/client/src/providers/useAIProvider.tsx @@ -5,6 +5,7 @@ import { UseAIChatPanel } from '../components/UseAIChatPanel'; import { UseAIFloatingChatWrapper, CloseButton } from '../components/UseAIFloatingChatWrapper'; import { __UseAIChatContext, type ChatUIContextValue } from '../components/UseAIChat'; import { UseAIClient } from '../client'; +import type { UseAITransport } from '../transport/types'; import { convertToolsToDefinitions, type ToolsDefinition } from '../defineTool'; import type { ChatRepository, Chat, ChatMetadata, CreateChatOptions, PersistedMessage, PersistedMessageContent } from './chatRepository/types'; import { LocalStorageChatRepository } from './chatRepository/LocalStorageChatRepository'; @@ -263,6 +264,24 @@ export interface ChatPanelProps { export interface UseAIProviderProps extends UseAIConfig { children: ReactNode; systemPrompt?: string; + /** + * Transport used to reach the server. Defaults to a {@link SocketIOTransport} built + * from `serverUrl`, which is what the bundled server serves. Supply one to reach a + * server that speaks something else — {@link WebSocketTransport} carries documented + * JSON frames over a plain WebSocket. + * + * Only the value from the first render is read, so an inline object does not churn + * the connection. Change transports by remounting the provider. + * + * @example + * ```tsx + * + * ``` + */ + transport?: UseAITransport; CustomButton?: React.ComponentType | null; CustomChat?: React.ComponentType | null; /** Default component overrides for every built-in chat rendered by this provider. */ @@ -471,6 +490,7 @@ const DEFAULT_FILE_UPLOAD_CONFIG: FileUploadConfig = { */ export function UseAIProvider({ serverUrl, + transport, children, systemPrompt, CustomButton, @@ -573,16 +593,20 @@ export function UseAIProvider({ const handleDisconnectRef = useRef(serverEvents.handleDisconnect); handleDisconnectRef.current = serverEvents.handleDisconnect; + // Read once: an inline `transport` object would otherwise re-create the client + // on every render. + const transportRef = useRef(transport); + useEffect(() => { console.log('[UseAIProvider] Initializing client with serverUrl:', serverUrl); - const client = new UseAIClient(serverUrl); + const client = new UseAIClient(transportRef.current ?? serverUrl); const unsubscribeConnection = client.onConnectionStateChange((isConnected) => { console.log('[UseAIProvider] Connection state changed:', isConnected); setConnected(isConnected); if (!isConnected) { - // The server destroys its session on disconnect (keyed by socket.id), - // so any in-flight run is unrecoverable even after Socket.IO reconnects. + // The server destroys its session on disconnect (keyed by connection id), + // so any in-flight run is unrecoverable even after the transport reconnects. // Reset UI state so the user can send a new message instead of being // stuck in a permanent "loading" state. handleDisconnectRef.current(); diff --git a/packages/client/src/transport/SocketIOTransport.test.ts b/packages/client/src/transport/SocketIOTransport.test.ts new file mode 100644 index 00000000..5d75bc4d --- /dev/null +++ b/packages/client/src/transport/SocketIOTransport.test.ts @@ -0,0 +1,183 @@ +import { describe, test, expect, mock, beforeEach, afterEach, spyOn } from 'bun:test'; +import type { Socket } from 'socket.io-client'; + +// Socket.IO builds its own socket, so the module is the only seam for testing +// this transport's wiring. Everything above it is tested through a transport +// instead: see client.test.ts. +let handlers: Record = {}; +let ioOptions: Record | undefined; +let mockSocket: Partial & { connected: boolean }; + +function createMockSocket() { + handlers = {}; + mockSocket = { + on: mock((event: string, handler: Function) => { + (handlers[event] ??= []).push(handler); + return mockSocket as Socket; + }), + emit: mock(() => mockSocket as Socket), + connected: false, + disconnect: mock(() => mockSocket as Socket), + io: { + engine: { + transport: { name: 'polling' }, + on: mock((event: string, handler: Function) => { + (handlers[`engine:${event}`] ??= []).push(handler); + }), + }, + } as never, + }; + return mockSocket as Socket; +} + +function emitSocketEvent(event: string, ...args: unknown[]) { + handlers[event]?.forEach(handler => handler(...args)); +} + +mock.module('socket.io-client', () => ({ + io: (_url: string, options: Record) => { + ioOptions = options; + return createMockSocket(); + }, +})); + +const { SocketIOTransport } = await import('./SocketIOTransport'); + +describe('SocketIOTransport', () => { + let consoleLogSpy: ReturnType; + let consoleWarnSpy: ReturnType; + + beforeEach(() => { + ioOptions = undefined; + consoleLogSpy = spyOn(console, 'log').mockImplementation(() => {}); + consoleWarnSpy = spyOn(console, 'warn').mockImplementation(() => {}); + }); + + afterEach(() => { + consoleLogSpy.mockRestore(); + consoleWarnSpy.mockRestore(); + }); + + test('reconnects indefinitely with backoff capped at ten seconds', () => { + new SocketIOTransport('http://localhost:8081').connect(); + + expect(ioOptions).toMatchObject({ + transports: ['polling', 'websocket'], + reconnection: true, + reconnectionAttempts: Infinity, + reconnectionDelay: 1000, + reconnectionDelayMax: 10_000, + withCredentials: true, + }); + }); + + test('reconnection delays are configurable', () => { + new SocketIOTransport('http://localhost:8081', { + reconnectionDelay: 250, + reconnectionDelayMax: 2000, + }).connect(); + + expect(ioOptions).toMatchObject({ reconnectionDelay: 250, reconnectionDelayMax: 2000 }); + }); + + test('dispatches connect, disconnect and named payloads', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + const received: Array<[string, unknown]> = []; + for (const name of ['connect', 'disconnect', 'event', 'agents', 'config'] as const) { + transport.on(name, data => received.push([name, data])); + } + + transport.connect(); + + emitSocketEvent('connect'); + emitSocketEvent('event', { type: 'RUN_STARTED' }); + emitSocketEvent('agents', { agents: [], defaultAgent: 'claude' }); + emitSocketEvent('config', { langfuseEnabled: true }); + emitSocketEvent('disconnect', 'transport close'); + + expect(received).toEqual([ + ['connect', undefined], + ['event', { type: 'RUN_STARTED' }], + ['agents', { agents: [], defaultAgent: 'claude' }], + ['config', { langfuseEnabled: true }], + ['disconnect', 'transport close'], + ]); + }); + + test('logs a warning on connection error without throwing', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + transport.connect(); + + emitSocketEvent('connect_error', new Error('Connection refused')); + + // Warn, not error: console.error triggers the Next.js error overlay. + expect(consoleWarnSpy).toHaveBeenCalledWith('[UseAI] Connection error:', 'Connection refused'); + }); + + test('repeated connection errors do not dispatch a connection state change', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + const states: string[] = []; + transport.on('connect', () => states.push('connect')); + transport.on('disconnect', () => states.push('disconnect')); + + transport.connect(); + + emitSocketEvent('connect_error', new Error('Attempt 1 failed')); + emitSocketEvent('connect_error', new Error('Attempt 2 failed')); + emitSocketEvent('connect'); + + expect(states).toEqual(['connect']); + expect(consoleWarnSpy).toHaveBeenCalledTimes(2); + }); + + test('logs transport upgrades', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + transport.connect(); + emitSocketEvent('connect'); + + expect(consoleLogSpy).toHaveBeenCalledWith('[UseAI] Transport:', 'polling'); + + emitSocketEvent('engine:upgrade', { name: 'websocket' }); + expect(consoleLogSpy).toHaveBeenCalledWith('[UseAI] Upgraded to transport:', 'websocket'); + + emitSocketEvent('engine:upgradeError', { message: 'upgrade failed' }); + expect(consoleWarnSpy).toHaveBeenCalledWith('[UseAI] Upgrade error:', 'upgrade failed'); + }); + + test('send emits the message on the message channel', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + transport.connect(); + mockSocket.connected = true; + emitSocketEvent('connect'); + + transport.send({ type: 'abort_run', data: { runId: 'run-1' } }); + + expect(mockSocket.emit).toHaveBeenCalledWith('message', { + type: 'abort_run', + data: { runId: 'run-1' }, + }); + }); + + test('connected follows the socket', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + expect(transport.connected).toBe(false); + + transport.connect(); + mockSocket.connected = true; + expect(transport.connected).toBe(true); + + mockSocket.connected = false; + expect(transport.connected).toBe(false); + }); + + test('disconnect closes the socket and reports disconnected', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + transport.connect(); + mockSocket.connected = true; + + transport.disconnect(); + + expect(mockSocket.disconnect).toHaveBeenCalled(); + expect(transport.connected).toBe(false); + }); +}); diff --git a/packages/client/src/transport/SocketIOTransport.ts b/packages/client/src/transport/SocketIOTransport.ts new file mode 100644 index 00000000..99f0650e --- /dev/null +++ b/packages/client/src/transport/SocketIOTransport.ts @@ -0,0 +1,115 @@ +import { io, Socket } from 'socket.io-client'; +import type { UseAIClientMessage } from '../types'; +import { TransportHandlerRegistry } from './handlerRegistry'; +import type { UseAITransport, UseAITransportEventName } from './types'; + +/** + * Options for {@link SocketIOTransport}. + */ +export interface SocketIOTransportOptions { + /** + * Delay before the first reconnection attempt, in milliseconds. + * + * @default 1000 + */ + reconnectionDelay?: number; + /** + * Upper bound on the exponential backoff between reconnection attempts, in milliseconds. + * + * @default 10000 + */ + reconnectionDelayMax?: number; +} + +/** + * Transport over Socket.IO. This is what {@link UseAIProvider} uses when given only a `serverUrl`, + * and what the bundled `@meetsmore-oss/use-ai-server` serves. + * + * @example + * ```typescript + * const transport = new SocketIOTransport('wss://your-server.com'); + * ``` + */ +export class SocketIOTransport implements UseAITransport { + private socket: Socket | null = null; + private registry = new TransportHandlerRegistry(); + // Reconnect indefinitely so clients recover after extended outages (mobile + // app backgrounded long enough for server pingTimeout, airplane mode, etc.). + // Socket.IO applies exponential backoff capped at reconnectionDelayMax, + // so steady-state retry frequency is ~one attempt per 10s. + private reconnectionDelay: number; + private reconnectionDelayMax: number; + + /** + * @param serverUrl - The URL of the UseAI server + * @example + * ```typescript + * new SocketIOTransport('ws://localhost:8081'); + * ``` + */ + constructor(private serverUrl: string, options: SocketIOTransportOptions = {}) { + this.reconnectionDelay = options.reconnectionDelay ?? 1000; + this.reconnectionDelayMax = options.reconnectionDelayMax ?? 10_000; + } + + get connected(): boolean { + return this.socket !== null && this.socket.connected; + } + + connect(): void { + const socket = io(this.serverUrl, { + transports: ['polling', 'websocket'], + reconnection: true, + reconnectionAttempts: Infinity, + reconnectionDelay: this.reconnectionDelay, + reconnectionDelayMax: this.reconnectionDelayMax, + withCredentials: true, + }); + this.socket = socket; + + socket.on('connect', () => { + console.log('[UseAI] Transport:', socket.io?.engine?.transport?.name); + + const engine = socket.io?.engine; + if (engine) { + engine.on('upgrade', (transport: { name: string }) => { + console.log('[UseAI] Upgraded to transport:', transport.name); + }); + + engine.on('upgradeError', (err: { message: string }) => { + console.warn('[UseAI] Upgrade error:', err.message); + }); + } + + this.registry.dispatch('connect', undefined); + }); + + socket.on('event', (event: unknown) => this.registry.dispatch('event', event)); + socket.on('agents', (data: unknown) => this.registry.dispatch('agents', data)); + socket.on('config', (data: unknown) => this.registry.dispatch('config', data)); + + socket.on('connect_error', (error: Error) => { + // Use warn instead of error to avoid triggering Next.js error overlay + console.warn('[UseAI] Connection error:', error.message); + }); + + socket.on('disconnect', (reason: string) => { + this.registry.dispatch('disconnect', reason); + }); + } + + disconnect(): void { + if (this.socket) { + this.socket.disconnect(); + this.socket = null; + } + } + + send(message: UseAIClientMessage): void { + this.socket?.emit('message', message); + } + + on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void { + return this.registry.on(name, handler); + } +} diff --git a/packages/client/src/transport/WebSocketTransport.test.ts b/packages/client/src/transport/WebSocketTransport.test.ts new file mode 100644 index 00000000..f14dfa8a --- /dev/null +++ b/packages/client/src/transport/WebSocketTransport.test.ts @@ -0,0 +1,286 @@ +import { describe, test, expect, beforeEach, afterEach, spyOn } from 'bun:test'; +import { WebSocketTransport, type WebSocketLike } from './WebSocketTransport'; + +class FakeWebSocket implements WebSocketLike { + static instances: FakeWebSocket[] = []; + + sent: string[] = []; + closeCalls = 0; + onopen: ((event: unknown) => void) | null = null; + onmessage: ((event: { data: unknown }) => void) | null = null; + onclose: ((event: unknown) => void) | null = null; + onerror: ((event: unknown) => void) | null = null; + + constructor(readonly url: string) { + FakeWebSocket.instances.push(this); + } + + static get latest(): FakeWebSocket { + return FakeWebSocket.instances[FakeWebSocket.instances.length - 1]; + } + + send(data: string): void { + this.sent.push(data); + } + + close(): void { + this.closeCalls++; + } + + /** Simulates the server accepting the connection. */ + serverOpen(): void { + this.onopen?.({}); + } + + /** Simulates a frame arriving from the server. */ + serverSend(data: unknown): void { + this.onmessage?.({ data }); + } + + /** Simulates the connection dropping. */ + serverClose(): void { + this.onclose?.({}); + } +} + +function makeTransport(options: { reconnectionDelay?: number; reconnectionDelayMax?: number } = {}) { + return new WebSocketTransport('wss://server.example/ws', { + ...options, + createWebSocket: (url) => new FakeWebSocket(url), + }); +} + +const tick = (ms: number) => new Promise(resolve => setTimeout(resolve, ms)); + +describe('WebSocketTransport', () => { + let consoleLogSpy: ReturnType; + let consoleWarnSpy: ReturnType; + + beforeEach(() => { + FakeWebSocket.instances = []; + consoleLogSpy = spyOn(console, 'log').mockImplementation(() => {}); + consoleWarnSpy = spyOn(console, 'warn').mockImplementation(() => {}); + }); + + afterEach(() => { + consoleLogSpy.mockRestore(); + consoleWarnSpy.mockRestore(); + }); + + test('opens a socket at the configured url', () => { + const transport = makeTransport(); + transport.connect(); + + expect(FakeWebSocket.instances).toHaveLength(1); + expect(FakeWebSocket.latest.url).toBe('wss://server.example/ws'); + + transport.disconnect(); + }); + + test('an incoming frame calls the handlers for its name', () => { + const transport = makeTransport(); + const events: unknown[] = []; + const configs: unknown[] = []; + transport.on('event', data => events.push(data)); + transport.on('config', data => configs.push(data)); + + transport.connect(); + FakeWebSocket.latest.serverOpen(); + + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: { type: 'RUN_STARTED' } })); + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'config', data: { langfuseEnabled: true } })); + + expect(events).toEqual([{ type: 'RUN_STARTED' }]); + expect(configs).toEqual([{ langfuseEnabled: true }]); + + transport.disconnect(); + }); + + test('delivers a frame to every subscriber of its name', () => { + const transport = makeTransport(); + const first: unknown[] = []; + const second: unknown[] = []; + transport.on('agents', data => first.push(data)); + transport.on('agents', data => second.push(data)); + + transport.connect(); + FakeWebSocket.latest.serverOpen(); + FakeWebSocket.latest.serverSend( + JSON.stringify({ name: 'agents', data: { agents: [], defaultAgent: 'claude' } }), + ); + + expect(first).toEqual([{ agents: [], defaultAgent: 'claude' }]); + expect(second).toEqual([{ agents: [], defaultAgent: 'claude' }]); + + transport.disconnect(); + }); + + test('unsubscribing stops delivery', () => { + const transport = makeTransport(); + const events: unknown[] = []; + const unsubscribe = transport.on('event', data => events.push(data)); + + transport.connect(); + FakeWebSocket.latest.serverOpen(); + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 1 })); + unsubscribe(); + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 2 })); + + expect(events).toEqual([1]); + + transport.disconnect(); + }); + + test('a frame with an unknown name is ignored, not an error', () => { + const transport = makeTransport(); + const events: unknown[] = []; + transport.on('event', data => events.push(data)); + + transport.connect(); + FakeWebSocket.latest.serverOpen(); + + expect(() => { + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'a_name_from_a_later_version', data: {} })); + }).not.toThrow(); + + // The connection survives, so the next known frame still arrives. + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 'still here' })); + expect(events).toEqual(['still here']); + expect(transport.connected).toBe(true); + + transport.disconnect(); + }); + + test('a malformed frame is ignored, not an error', () => { + const transport = makeTransport(); + const events: unknown[] = []; + transport.on('event', data => events.push(data)); + + transport.connect(); + FakeWebSocket.latest.serverOpen(); + + expect(() => FakeWebSocket.latest.serverSend('not json')).not.toThrow(); + expect(() => FakeWebSocket.latest.serverSend(JSON.stringify({ noNameHere: true }))).not.toThrow(); + + FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 'still here' })); + expect(events).toEqual(['still here']); + + transport.disconnect(); + }); + + test('send serializes the message with nothing wrapped around it', () => { + const transport = makeTransport(); + transport.connect(); + FakeWebSocket.latest.serverOpen(); + + transport.send({ type: 'abort_run', data: { runId: 'run-1' } }); + + expect(FakeWebSocket.latest.sent).toEqual(['{"type":"abort_run","data":{"runId":"run-1"}}']); + + transport.disconnect(); + }); + + test('connected follows the socket opening and closing', () => { + const transport = makeTransport({ reconnectionDelay: 10_000 }); + expect(transport.connected).toBe(false); + + transport.connect(); + expect(transport.connected).toBe(false); + + FakeWebSocket.latest.serverOpen(); + expect(transport.connected).toBe(true); + + FakeWebSocket.latest.serverClose(); + expect(transport.connected).toBe(false); + + transport.disconnect(); + }); + + test('a close after opening dispatches disconnect', () => { + const transport = makeTransport({ reconnectionDelay: 10_000 }); + const states: string[] = []; + transport.on('connect', () => states.push('connect')); + transport.on('disconnect', () => states.push('disconnect')); + + transport.connect(); + FakeWebSocket.latest.serverOpen(); + FakeWebSocket.latest.serverClose(); + + expect(states).toEqual(['connect', 'disconnect']); + + transport.disconnect(); + }); + + test('a failed connection attempt does not dispatch disconnect', () => { + const transport = makeTransport({ reconnectionDelay: 10_000 }); + const states: string[] = []; + transport.on('disconnect', () => states.push('disconnect')); + + transport.connect(); + // Never opened: the socket closes straight from the connecting state. + FakeWebSocket.latest.serverClose(); + + expect(states).toEqual([]); + + transport.disconnect(); + }); + + test('backoff reconnects after the socket drops', async () => { + const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); + transport.connect(); + FakeWebSocket.latest.serverOpen(); + expect(FakeWebSocket.instances).toHaveLength(1); + + FakeWebSocket.latest.serverClose(); + await tick(20); + + expect(FakeWebSocket.instances.length).toBeGreaterThan(1); + + // The reconnected socket is live: opening it restores connected. + FakeWebSocket.latest.serverOpen(); + expect(transport.connected).toBe(true); + + transport.disconnect(); + }); + + test('backoff keeps retrying while attempts fail', async () => { + const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); + transport.connect(); + + for (let i = 0; i < 3; i++) { + FakeWebSocket.latest.serverClose(); + await tick(10); + } + + expect(FakeWebSocket.instances.length).toBeGreaterThanOrEqual(4); + + transport.disconnect(); + }); + + test('disconnect() stops the backoff', async () => { + const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); + transport.connect(); + FakeWebSocket.latest.serverOpen(); + + FakeWebSocket.latest.serverClose(); + transport.disconnect(); + const openedByNow = FakeWebSocket.instances.length; + + await tick(20); + + expect(FakeWebSocket.instances).toHaveLength(openedByNow); + expect(transport.connected).toBe(false); + }); + + test('disconnect() closes the open socket', () => { + const transport = makeTransport(); + transport.connect(); + FakeWebSocket.latest.serverOpen(); + + const socket = FakeWebSocket.latest; + transport.disconnect(); + + expect(socket.closeCalls).toBe(1); + expect(transport.connected).toBe(false); + }); +}); diff --git a/packages/client/src/transport/WebSocketTransport.ts b/packages/client/src/transport/WebSocketTransport.ts new file mode 100644 index 00000000..521a2d91 --- /dev/null +++ b/packages/client/src/transport/WebSocketTransport.ts @@ -0,0 +1,224 @@ +import type { UseAIClientMessage } from '../types'; +import { TransportHandlerRegistry } from './handlerRegistry'; +import type { UseAITransport, UseAITransportEventName } from './types'; + +/** + * The subset of the WHATWG `WebSocket` API that {@link WebSocketTransport} uses. + * Declared structurally so a test double, or a polyfill on a runtime without a + * global `WebSocket`, can stand in for the real thing. + */ +export interface WebSocketLike { + send(data: string): void; + close(): void; + onopen: ((event: unknown) => void) | null; + onmessage: ((event: { data: unknown }) => void) | null; + onclose: ((event: unknown) => void) | null; + onerror: ((event: unknown) => void) | null; +} + +/** + * A downstream frame, as sent by the server. + * + * @example + * ```json + * { "name": "config", "data": { "langfuseEnabled": true } } + * ``` + */ +interface DownstreamFrame { + name: string; + data: unknown; +} + +/** + * Options for {@link WebSocketTransport}. + */ +export interface WebSocketTransportOptions { + /** + * Delay before the first reconnection attempt, in milliseconds. + * Subsequent attempts double this, up to {@link reconnectionDelayMax}. + * + * @default 1000 + */ + reconnectionDelay?: number; + /** + * Upper bound on the exponential backoff between reconnection attempts, in milliseconds. + * + * @default 10000 + */ + reconnectionDelayMax?: number; + /** + * Opens the underlying socket. + * + * @default (url) => new WebSocket(url) + */ + createWebSocket?: (url: string) => WebSocketLike; +} + +/** + * Transport over a plain WebSocket carrying JSON text frames. + * + * Use this to reach a server that does not speak Socket.IO. The framing is: + * + * - **Upstream** — the `UseAIClientMessage`, serialized, with nothing wrapped around it: + * `{"type":"run_agent","data":{...}}` + * - **Downstream** — a named envelope, because a plain WebSocket has no event names of its own: + * `{"name":"event","data":{...}}`. The names are `event`, `agents` and `config`. + * A frame with any other name is ignored, so a server may add names without breaking + * older clients. + * + * A server should send `agents` and `config` once, after the connection opens. + * + * @example + * ```typescript + * + * ``` + */ +export class WebSocketTransport implements UseAITransport { + private socket: WebSocketLike | null = null; + private registry = new TransportHandlerRegistry(); + private _connected = false; + // A plain WebSocket has no reconnection of its own, so this transport matches + // the Socket.IO settings: retry indefinitely with exponential backoff capped at + // reconnectionDelayMax, so a client recovers after an extended outage (mobile app + // backgrounded, airplane mode) without hammering the server in the meantime. + private reconnectionDelay: number; + private reconnectionDelayMax: number; + private reconnectAttempts = 0; + private reconnectTimer: ReturnType | null = null; + private reconnecting = false; + private createWebSocket: (url: string) => WebSocketLike; + + /** + * @param url - WebSocket URL of the server + * @example + * ```typescript + * new WebSocketTransport('wss://your-server.com/ws'); + * ``` + */ + constructor(private url: string, options: WebSocketTransportOptions = {}) { + this.reconnectionDelay = options.reconnectionDelay ?? 1000; + this.reconnectionDelayMax = options.reconnectionDelayMax ?? 10_000; + this.createWebSocket = + options.createWebSocket ?? ((url: string) => new WebSocket(url) as unknown as WebSocketLike); + } + + get connected(): boolean { + return this._connected; + } + + connect(): void { + this.reconnecting = true; + this.open(); + } + + disconnect(): void { + this.reconnecting = false; + if (this.reconnectTimer !== null) { + clearTimeout(this.reconnectTimer); + this.reconnectTimer = null; + } + + const socket = this.socket; + this.socket = null; + this._connected = false; + if (socket) { + this.detach(socket); + socket.close(); + } + } + + send(message: UseAIClientMessage): void { + this.socket?.send(JSON.stringify(message)); + } + + on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void { + return this.registry.on(name, handler); + } + + private open(): void { + let socket: WebSocketLike; + try { + socket = this.createWebSocket(this.url); + } catch (error) { + // Use warn instead of error to avoid triggering Next.js error overlay + console.warn('[UseAI] Connection error:', error instanceof Error ? error.message : error); + this.scheduleReconnect(); + return; + } + this.socket = socket; + + socket.onopen = () => { + this._connected = true; + this.reconnectAttempts = 0; + this.registry.dispatch('connect', undefined); + }; + + socket.onmessage = (event) => { + const frame = this.parseFrame(event.data); + if (!frame) return; + this.registry.dispatch(frame.name, frame.data); + }; + + socket.onerror = () => { + // onclose always follows, and that is where reconnection is scheduled. + console.warn('[UseAI] Connection error:', this.url); + }; + + socket.onclose = () => { + this.detach(socket); + if (this.socket !== socket) return; + this.socket = null; + + const wasConnected = this._connected; + this._connected = false; + // A socket that never opened reports only a failed attempt, not a disconnection. + if (wasConnected) { + this.registry.dispatch('disconnect', 'transport close'); + } + this.scheduleReconnect(); + }; + } + + private parseFrame(data: unknown): DownstreamFrame | null { + if (typeof data !== 'string') { + console.warn('[UseAI] Ignoring non-text frame'); + return null; + } + let parsed: unknown; + try { + parsed = JSON.parse(data); + } catch { + console.warn('[UseAI] Ignoring malformed frame'); + return null; + } + if (typeof parsed !== 'object' || parsed === null) return null; + const frame = parsed as Partial; + if (typeof frame.name !== 'string') return null; + return { name: frame.name, data: frame.data }; + } + + private detach(socket: WebSocketLike): void { + socket.onopen = null; + socket.onmessage = null; + socket.onerror = null; + socket.onclose = null; + } + + private scheduleReconnect(): void { + if (!this.reconnecting || this.reconnectTimer !== null) return; + + const delay = Math.min( + this.reconnectionDelay * 2 ** this.reconnectAttempts, + this.reconnectionDelayMax, + ); + this.reconnectAttempts++; + + this.reconnectTimer = setTimeout(() => { + this.reconnectTimer = null; + if (this.reconnecting) this.open(); + }, delay); + } +} diff --git a/packages/client/src/transport/handlerRegistry.ts b/packages/client/src/transport/handlerRegistry.ts new file mode 100644 index 00000000..2938237a --- /dev/null +++ b/packages/client/src/transport/handlerRegistry.ts @@ -0,0 +1,32 @@ +import type { UseAITransportEventName } from './types'; + +type Handler = (data: unknown) => void; + +/** + * The subscribe/dispatch bookkeeping shared by the bundled transports. + * A name with no subscribers dispatches to nobody, which is what makes an + * unrecognised downstream frame a no-op rather than an error. + */ +export class TransportHandlerRegistry { + private handlers: Map> = new Map(); + + on(name: UseAITransportEventName, handler: Handler): () => void { + let set = this.handlers.get(name); + if (!set) { + set = new Set(); + this.handlers.set(name, set); + } + set.add(handler); + return () => { + set.delete(handler); + }; + } + + dispatch(name: string, data: unknown): void { + const set = this.handlers.get(name); + if (!set) return; + for (const handler of [...set]) { + handler(data); + } + } +} diff --git a/packages/client/src/transport/index.ts b/packages/client/src/transport/index.ts new file mode 100644 index 00000000..7dd56cc1 --- /dev/null +++ b/packages/client/src/transport/index.ts @@ -0,0 +1,5 @@ +export type { UseAITransport, UseAITransportEventName } from './types'; +export { SocketIOTransport } from './SocketIOTransport'; +export type { SocketIOTransportOptions } from './SocketIOTransport'; +export { WebSocketTransport } from './WebSocketTransport'; +export type { WebSocketTransportOptions, WebSocketLike } from './WebSocketTransport'; diff --git a/packages/client/src/transport/types.ts b/packages/client/src/transport/types.ts new file mode 100644 index 00000000..bf85c79c --- /dev/null +++ b/packages/client/src/transport/types.ts @@ -0,0 +1,42 @@ +import type { UseAIClientMessage } from '../types'; + +/** + * Names of the downstream channels a transport delivers to {@link UseAIClient}. + * + * - `connect` / `disconnect` — connection lifecycle. `disconnect` carries a reason string. + * - `event` — an AG-UI event. + * - `agents` — the server's agent list, `{ agents, defaultAgent }`. + * - `config` — server capability flags, `{ langfuseEnabled }`. + */ +export type UseAITransportEventName = 'connect' | 'disconnect' | 'event' | 'agents' | 'config'; + +/** + * The pipe between {@link UseAIClient} and a server. + * + * A transport carries {@link UseAIClientMessage} upstream and named payloads downstream. + * It owns everything protocol-specific: how a connection is opened, how it reconnects, + * and how a named payload is framed on the wire. + * + * Two implementations ship with the library: {@link SocketIOTransport} (the default) + * and {@link WebSocketTransport}. + */ +export interface UseAITransport { + /** Opens the connection. Reconnection until {@link disconnect} is the transport's own responsibility. */ + connect(): void; + + /** Closes the connection and stops reconnecting. */ + disconnect(): void; + + /** Sends a message upstream. Only called while {@link connected} is true. */ + send(message: UseAIClientMessage): void; + + /** + * Subscribes to a downstream channel. + * + * @returns Cleanup function to unsubscribe + */ + on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void; + + /** Whether the connection is currently open. */ + readonly connected: boolean; +} diff --git a/packages/client/src/types.ts b/packages/client/src/types.ts index 71d7c310..35606b18 100644 --- a/packages/client/src/types.ts +++ b/packages/client/src/types.ts @@ -2,7 +2,10 @@ * Configuration for the UseAI client provider. */ export interface UseAIConfig { - /** The WebSocket URL of the UseAI server */ + /** + * The WebSocket URL of the UseAI server. + * Unused when an explicit `transport` is supplied, but still reported on the context. + */ serverUrl: string; } diff --git a/packages/server/package.json b/packages/server/package.json index f2ad6d07..78cd4ce4 100644 --- a/packages/server/package.json +++ b/packages/server/package.json @@ -39,6 +39,7 @@ "picomatch": "^4.0.3", "socket.io": "^4.8.1", "uuid": "^11.1.0", + "ws": "^8.18.3", "zod": "^3.24.1", "zod-to-json-schema": "^3.25.1" }, @@ -50,11 +51,13 @@ "@opentelemetry/sdk-trace-node": "^2.2.0" }, "devDependencies": { + "@meetsmore-oss/use-ai-client": "workspace:^", "@types/bun": "^1.3.13", "@types/cors": "^2.8.17", "@types/json-schema": "^7.0.15", "@types/picomatch": "^4.0.2", "@types/uuid": "^10.0.0", + "@types/ws": "^8.18.1", "socket.io-client": "^4.8.1", "testcontainers": "^11.8.1", "typescript": "^5.9.3" diff --git a/packages/server/src/agents/types.ts b/packages/server/src/agents/types.ts index 79c0adf7..5f4a6ed9 100644 --- a/packages/server/src/agents/types.ts +++ b/packages/server/src/agents/types.ts @@ -1,7 +1,21 @@ -import type { Socket } from 'socket.io'; import type { ModelMessage } from 'ai'; import type { ToolDefinition, AGUIEvent, ToolApprovalRequestEvent } from '../types'; +/** + * A client connection, as much of it as a session needs. + * + * Both connection kinds the server accepts satisfy this: a Socket.IO socket, and a + * plain WebSocket where `emit(name, data)` is written out as `{"name":...,"data":...}`. + */ +export interface ClientConnection { + /** Identifies the connection for the lifetime of the session. */ + readonly id: string; + /** Whether the connection is still open. */ + readonly connected: boolean; + /** Sends a named payload to the client. */ + emit(name: string, data?: unknown): void; +} + /** * Context for a single client session. * Contains all state needed for multi-turn conversations and tool coordination. @@ -11,8 +25,8 @@ export interface ClientSession { clientId: string; /** IP address of the client (used for rate limiting) */ ipAddress: string; - /** Socket.IO socket for bidirectional communication with client */ - socket: Socket; + /** Connection to the client. Emit named payloads on it to reach the client. */ + socket: ClientConnection; /** Unique identifier for the conversation thread */ threadId: string; /** ID of the currently executing run (if any) */ diff --git a/packages/server/src/index.ts b/packages/server/src/index.ts index 314cad83..8dec0d9d 100644 --- a/packages/server/src/index.ts +++ b/packages/server/src/index.ts @@ -1,6 +1,6 @@ export { UseAIServer } from './server'; export type { UseAIServerConfig, McpEndpointConfig, ToolDefinition, CorsOptions } from './types'; -export type { ClientSession } from './server'; +export type { ClientSession, ClientConnection } from './server'; // Attachment ref resolution (host seam + helper) export { resolveAttachmentParts } from './attachmentResolution'; diff --git a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts index 2874db43..8efa6153 100644 --- a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts +++ b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts @@ -2,6 +2,7 @@ import type { Server as SocketIOServer } from 'socket.io'; import { Server as BunEngine } from '@socket.io/bun-engine'; import type { RuntimeAdapter, RuntimeServerConfig, RuntimeServerHandle } from '../types'; import { resolveCorsHeaders, resolvePreflightHeaders } from './cors'; +import { BunRawWebSocket, isRawWebSocket, type RawWebSocketData } from './rawWebSocket'; /** * Runtime adapter for Bun. @@ -33,7 +34,34 @@ export class BunRuntimeAdapter implements RuntimeAdapter { io.bind(this.engine); const handler = this.engine.handleRequest.bind(this.engine); - const websocketHandler = this.engine.handler().websocket; + const engineWebSocket = this.engine.handler().websocket; + + // Bun.serve takes a single websocket handler table for the whole server, so the + // plain listener and the Socket.IO engine share it and dispatch on ws.data. + const rawSockets = new WeakMap(); + type EngineWebSocketHandler = typeof engineWebSocket; + type EngineWebSocket = Parameters[0]; + const websocketHandler: EngineWebSocketHandler = { + ...engineWebSocket, + open: (ws: EngineWebSocket) => { + if (!isRawWebSocket(ws)) return engineWebSocket.open(ws); + const { remoteAddress } = ws.data as unknown as RawWebSocketData; + const connection = new BunRawWebSocket(ws, remoteAddress); + rawSockets.set(ws, connection); + config.websocket?.onConnection(connection); + }, + message: (ws: EngineWebSocket, message: Parameters[1]) => { + if (!isRawWebSocket(ws)) return engineWebSocket.message(ws, message); + rawSockets.get(ws)?.receiveMessage( + typeof message === 'string' ? message : new TextDecoder().decode(message as Uint8Array), + ); + }, + close: (ws: EngineWebSocket, code: number, reason: string) => { + if (!isRawWebSocket(ws)) return engineWebSocket.close(ws, code, reason); + rawSockets.get(ws)?.receiveClose(); + rawSockets.delete(ws); + }, + }; // Start Bun server const bunServer = Bun.serve({ @@ -63,6 +91,18 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }); } + // Plain WebSocket path + if (config.websocket && url.pathname === config.websocket.path) { + const data: RawWebSocketData = { + useAiRawWebSocket: true, + remoteAddress: server.requestIP(req)?.address, + }; + // The engine owns the server's WebSocket data type; the plain listener + // rides along on the same handler table and is told apart by useAiRawWebSocket. + if (server.upgrade(req, { data: data as unknown as EngineWebSocket['data'] })) return undefined; + return new Response('Expected a WebSocket upgrade', { status: 400, headers: corsHeaders }); + } + // Socket.IO path if (url.pathname.startsWith('/socket.io/')) { const response = await handler(req, server); diff --git a/packages/server/src/runtime/bun/rawWebSocket.ts b/packages/server/src/runtime/bun/rawWebSocket.ts new file mode 100644 index 00000000..79f93011 --- /dev/null +++ b/packages/server/src/runtime/bun/rawWebSocket.ts @@ -0,0 +1,59 @@ +import type { RawWebSocket } from '../types'; + +/** Marks a Bun WebSocket as belonging to the plain listener rather than to Socket.IO. */ +export interface RawWebSocketData { + useAiRawWebSocket: true; + remoteAddress?: string; +} + +interface BunWebSocket { + readonly readyState: number; + data: unknown; + send(data: string): unknown; + close(): void; +} + +const OPEN = 1; + +/** + * Adapts a Bun WebSocket, which delivers frames through the server-level handler + * table rather than per-socket callbacks, to {@link RawWebSocket}. + */ +export class BunRawWebSocket implements RawWebSocket { + private messageHandler: ((data: string) => void) | null = null; + private closeHandler: (() => void) | null = null; + + constructor(private ws: BunWebSocket, readonly remoteAddress?: string) {} + + get open(): boolean { + return this.ws.readyState === OPEN; + } + + send(data: string): void { + this.ws.send(data); + } + + close(): void { + this.ws.close(); + } + + onMessage(handler: (data: string) => void): void { + this.messageHandler = handler; + } + + onClose(handler: () => void): void { + this.closeHandler = handler; + } + + receiveMessage(data: string): void { + this.messageHandler?.(data); + } + + receiveClose(): void { + this.closeHandler?.(); + } +} + +export function isRawWebSocket(ws: { data?: unknown }): boolean { + return (ws.data as Partial | undefined)?.useAiRawWebSocket === true; +} diff --git a/packages/server/src/runtime/index.ts b/packages/server/src/runtime/index.ts index 26bf7afa..029f0ce5 100644 --- a/packages/server/src/runtime/index.ts +++ b/packages/server/src/runtime/index.ts @@ -7,6 +7,8 @@ export type { RuntimeType, RuntimeServerConfig, RuntimeServerHandle, + RawWebSocket, + RawWebSocketListener, } from './types'; export { detectRuntime } from './detection'; export { createClientIpTracker, type ClientIpTracker, type ClientIpConnection } from './clientIp'; diff --git a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts index 7e452984..1c55b104 100644 --- a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts +++ b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts @@ -1,6 +1,8 @@ import { createServer, type Server as HttpServer } from 'http'; import type { Server as SocketIOServer } from 'socket.io'; +import { WebSocketServer } from 'ws'; import type { RuntimeAdapter, RuntimeServerConfig, RuntimeServerHandle } from '../types'; +import { NodeRawWebSocket } from './rawWebSocket'; /** * Runtime adapter for Node.js. @@ -34,6 +36,25 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { // Socket.IO will handle the request via its internal listeners }); + // The plain listener claims its path before Socket.IO attaches, so engine.io's + // own upgrade handler sees a handshake already written and leaves the socket alone. + const websocketConfig = config.websocket; + let wss: WebSocketServer | null = null; + if (websocketConfig) { + wss = new WebSocketServer({ noServer: true, maxPayload: config.maxHttpBufferSize }); + httpServer.on('upgrade', (req, socket, head) => { + const url = new URL(req.url || '/', `http://localhost:${config.port}`); + if (url.pathname !== websocketConfig.path) return; + wss!.handleUpgrade(req, socket, head, (ws) => { + const forwardedFor = req.headers['x-forwarded-for']; + const remoteAddress = typeof forwardedFor === 'string' + ? forwardedFor.split(',')[0].trim() + : req.socket.remoteAddress; + websocketConfig.onConnection(new NodeRawWebSocket(ws, remoteAddress)); + }); + }); + } + // Attach Socket.IO to the HTTP server // Socket.IO handles CORS internally for /socket.io/* paths io.attach(httpServer, { @@ -66,6 +87,7 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { return { stop: () => { + wss?.close(); httpServer.close(); }, server: httpServer, diff --git a/packages/server/src/runtime/node/rawWebSocket.ts b/packages/server/src/runtime/node/rawWebSocket.ts new file mode 100644 index 00000000..db26abcc --- /dev/null +++ b/packages/server/src/runtime/node/rawWebSocket.ts @@ -0,0 +1,31 @@ +import type { WebSocket } from 'ws'; +import type { RawWebSocket } from '../types'; + +/** + * Adapts a `ws` WebSocket to {@link RawWebSocket}. + */ +export class NodeRawWebSocket implements RawWebSocket { + constructor(private ws: WebSocket, readonly remoteAddress?: string) {} + + get open(): boolean { + return this.ws.readyState === this.ws.OPEN; + } + + send(data: string): void { + this.ws.send(data); + } + + close(): void { + this.ws.close(); + } + + onMessage(handler: (data: string) => void): void { + this.ws.on('message', (data: unknown, isBinary: boolean) => { + handler(isBinary ? '' : String(data)); + }); + } + + onClose(handler: () => void): void { + this.ws.on('close', handler); + } +} diff --git a/packages/server/src/runtime/types.ts b/packages/server/src/runtime/types.ts index 76f58c81..eeb3baa5 100644 --- a/packages/server/src/runtime/types.ts +++ b/packages/server/src/runtime/types.ts @@ -6,6 +6,38 @@ import type { CorsOptions } from '../types'; */ export type RuntimeType = 'bun' | 'node'; +/** + * A plain WebSocket connection, as the runtime adapters expose it. + * Text frames only; the framing above it is the caller's business. + */ +export interface RawWebSocket { + /** Remote address of the peer, when the runtime exposes one. */ + readonly remoteAddress?: string; + /** Whether the connection is still open. */ + readonly open: boolean; + /** Sends a text frame. */ + send(data: string): void; + /** Closes the connection. */ + close(): void; + /** Registers the handler for text frames arriving from the peer. */ + onMessage(handler: (data: string) => void): void; + /** Registers the handler for the connection closing, for any reason. */ + onClose(handler: () => void): void; +} + +/** + * A plain WebSocket listener, served on the same port and HTTP server as Socket.IO. + */ +export interface RawWebSocketListener { + /** + * Path that upgrades to a plain WebSocket. + * @example '/ws' + */ + path: string; + /** Called once per accepted connection. */ + onConnection(connection: RawWebSocket): void; +} + /** * Configuration for creating a runtime server. */ @@ -27,6 +59,11 @@ export interface RuntimeServerConfig { * Called when a polling transport connection is established. */ onPollingConnection?: (sessionId: string, ip: string) => void; + /** + * Plain WebSocket listener to serve alongside Socket.IO. + * Omit to serve Socket.IO only. + */ + websocket?: RawWebSocketListener; } /** diff --git a/packages/server/src/server.ts b/packages/server/src/server.ts index f9169118..0b2212a8 100644 --- a/packages/server/src/server.ts +++ b/packages/server/src/server.ts @@ -22,7 +22,7 @@ import { logger } from './logger'; import { recordErrorTrace, startTracing } from './instrumentation'; import { v4 as uuidv4 } from 'uuid'; import type { Agent, EventEmitter, AGUIEventExtended } from './agents/types'; -import type { ClientSession } from './agents/types'; +import type { ClientConnection, ClientSession } from './agents/types'; import type { UseAIServerPlugin, MessageHandler } from './plugins/types'; import { FeedbackPlugin } from './plugins/FeedbackPlugin'; import { isRemoteTool, isServerTool } from './utils/toolFilters'; @@ -36,10 +36,12 @@ import { type RuntimeAdapter, type RuntimeServerHandle, type ClientIpTracker, + type RawWebSocket, } from './runtime'; +import { WebSocketClientConnection } from './webSocketConnection'; -// Re-export ClientSession type for external use -export type { ClientSession } from './agents/types'; +// Re-export session types for external use +export type { ClientSession, ClientConnection } from './agents/types'; /** * WebSocket server that coordinates between client applications and AI agents. @@ -96,10 +98,11 @@ export class UseAIServer { private defaultAgentId: string; // ID of the default agent private agents: Record; // Registry of all agents private clients: Map = new Map(); - private config: Required> & { + private config: Required> & { maxHttpBufferSize: number; cors?: CorsOptions; idleTimeout: number; + webSocketPath: string | null; }; private rateLimiter: RateLimiter; private cleanupInterval: NodeJS.Timeout; @@ -131,6 +134,7 @@ export class UseAIServer { maxHttpBufferSize: config.maxHttpBufferSize ?? 20 * 1024 * 1024, // 20MB default cors: config.cors, idleTimeout: config.idleTimeout ?? 30, + webSocketPath: config.webSocketPath === undefined ? '/ws' : config.webSocketPath, }; // Set agents registry @@ -218,6 +222,12 @@ export class UseAIServer { onPollingConnection: (sessionId, ip) => { this.clientIpTracker.trackPollingConnection(sessionId, ip); }, + websocket: this.config.webSocketPath + ? { + path: this.config.webSocketPath, + onConnection: (socket) => this.handleWebSocketConnection(socket), + } + : undefined, }); } @@ -285,8 +295,6 @@ export class UseAIServer { private setupSocketIOServer() { this.io.on('connection', (socket: Socket) => { - const clientId = `client-${++this.clientIdCounter}`; - const threadId = uuidv4(); // Get connection info for IP address resolution const conn = socket.conn as unknown as { id: string; transport: { name: string; socket?: { remoteAddress?: string } } }; // Get IP address for rate limiting: @@ -296,96 +304,150 @@ export class UseAIServer { const ipAddress = this.clientIpTracker.getClientIp(conn) || socket.handshake.address || socket.id; - const transport = conn.transport.name; - logger.info('Client connected', { clientId, threadId, ipAddress, transport }); + + const session = this.createSession(socket, ipAddress); + logger.info('Client connected', { + clientId: session.clientId, + threadId: session.threadId, + ipAddress, + transport: conn.transport.name, + }); // Log transport upgrades socket.conn.on('upgrade', (transport) => { - logger.info('Client upgraded transport', { clientId, transport: transport.name }); + logger.info('Client upgraded transport', { clientId: session.clientId, transport: transport.name }); }); - const session: ClientSession = { - clientId, - ipAddress, - socket, - threadId, - tools: [], - state: null, - pendingToolCalls: new Map(), - pendingToolApprovals: new Map(), - }; + socket.on('message', (message: UseAIClientMessage) => this.receiveClientMessage(session, message)); + + socket.on('disconnect', () => { + logger.info('Client disconnected', { clientId: session.clientId, ipAddress }); + // Clean up polling IP entry + this.clientIpTracker.removePollingConnection(conn.id); + this.destroySession(session); + }); + }); + + logger.info('UseAI server ready', { port: this.config.port }); + } - this.clients.set(socket.id, session); + /** + * Accepts a plain WebSocket connection, the alternative to Socket.IO served on the + * same port. Frames are JSON text: upstream the client message on its own, downstream + * a `{ name, data }` envelope. See docs/websocket-protocol.md. + */ + private handleWebSocketConnection(socket: RawWebSocket) { + const connection = new WebSocketClientConnection(`ws-${uuidv4()}`, socket); + const ipAddress = socket.remoteAddress || connection.id; - // Send available agents to client - const availableAgents = Object.entries(this.agents).map(([id, agent]) => ({ + const session = this.createSession(connection, ipAddress); + logger.info('Client connected', { + clientId: session.clientId, + threadId: session.threadId, + ipAddress, + transport: 'websocket', + }); + + socket.onMessage((data) => { + let message: UseAIClientMessage; + try { + message = JSON.parse(data) as UseAIClientMessage; + } catch { + logger.warn('Discarding malformed frame', { clientId: session.clientId }); + return; + } + void this.receiveClientMessage(session, message); + }); + + socket.onClose(() => { + logger.info('Client disconnected', { clientId: session.clientId, ipAddress }); + this.destroySession(session); + }); + } + + /** + * Creates the session for a newly accepted connection and announces the server's + * agents to it. Both connection kinds land here. + */ + private createSession(connection: ClientConnection, ipAddress: string): ClientSession { + const session: ClientSession = { + clientId: `client-${++this.clientIdCounter}`, + ipAddress, + socket: connection, + threadId: uuidv4(), + tools: [], + state: null, + pendingToolCalls: new Map(), + pendingToolApprovals: new Map(), + }; + + this.clients.set(connection.id, session); + + connection.emit('agents', { + agents: Object.entries(this.agents).map(([id, agent]) => ({ id, name: agent.getName?.() || id, annotation: agent.getAnnotation?.(), - })); - socket.emit('agents', { - agents: availableAgents, - defaultAgent: this.defaultAgentId, - }); - - // Call plugin lifecycle hooks - for (const plugin of this.plugins) { - plugin.onClientConnect?.(session); - } + })), + defaultAgent: this.defaultAgentId, + }); - socket.on('message', async (message: UseAIClientMessage) => { - try { - await this.handleClientMessage(socket, message); - } catch (error) { - logger.error('Error handling message', { - error: error instanceof Error ? error.message : 'Unknown error', - clientId, - }); - if (message.type === 'run_agent') { - const runAgentData = (message as RunAgentMessage).data; - const unhandledForwardedProps = runAgentData?.forwardedProps as UseAIForwardedProps | undefined; - recordErrorTrace({ - runId: runAgentData?.runId || socket.id, - errorCategory: 'unhandled_error', - errorMessage: error instanceof Error ? error.message : 'Unknown error', - sessionId: clientId, - threadId: runAgentData?.threadId, - ipAddress: session?.ipAddress, - metadata: { ...unhandledForwardedProps?.telemetryMetadata }, - }); - } - this.sendEvent(socket, { - type: EventType.RUN_ERROR, - message: error instanceof Error ? error.message : 'Unknown error', - timestamp: Date.now(), - }); - } - }); + for (const plugin of this.plugins) { + plugin.onClientConnect?.(session); + } - socket.on('disconnect', () => { - logger.info('Client disconnected', { clientId, ipAddress }); + return session; + } - // Abort any pending tool calls/approvals for this session - abortRun(session.abortController, new RunAbortedByClientDisconnect()); + /** + * Tears down a session whose connection has gone away. + * Rate limiting persists by IP address across connections, so it is not reset here. + */ + private destroySession(session: ClientSession) { + // Abort any pending tool calls/approvals for this session + abortRun(session.abortController, new RunAbortedByClientDisconnect()); - // Clean up polling IP entry - this.clientIpTracker.removePollingConnection(conn.id); + for (const plugin of this.plugins) { + plugin.onClientDisconnect?.(session); + } - // Call plugin lifecycle hooks - for (const plugin of this.plugins) { - plugin.onClientDisconnect?.(session); - } + this.clients.delete(session.socket.id); + } - // Note: Rate limiting persists by IP address across connections - this.clients.delete(socket.id); + /** + * Runs a client message and reports anything it throws back to that client. + */ + private async receiveClientMessage(session: ClientSession, message: UseAIClientMessage) { + try { + await this.handleClientMessage(session.socket, message); + } catch (error) { + logger.error('Error handling message', { + error: error instanceof Error ? error.message : 'Unknown error', + clientId: session.clientId, }); - }); - - logger.info('UseAI server ready', { port: this.config.port }); + if (message.type === 'run_agent') { + const runAgentData = (message as RunAgentMessage).data; + const unhandledForwardedProps = runAgentData?.forwardedProps as UseAIForwardedProps | undefined; + recordErrorTrace({ + runId: runAgentData?.runId || session.socket.id, + errorCategory: 'unhandled_error', + errorMessage: error instanceof Error ? error.message : 'Unknown error', + sessionId: session.clientId, + threadId: runAgentData?.threadId, + ipAddress: session.ipAddress, + metadata: { ...unhandledForwardedProps?.telemetryMetadata }, + }); + } + this.sendEvent(session.socket, { + type: EventType.RUN_ERROR, + message: error instanceof Error ? error.message : 'Unknown error', + timestamp: Date.now(), + }); + } } - private async handleClientMessage(socket: Socket, message: UseAIClientMessage) { - const session = this.clients.get(socket.id); + private async handleClientMessage(connection: ClientConnection, message: UseAIClientMessage) { + const session = this.clients.get(connection.id); if (!session) return; // Check if a plugin has registered a handler for this message type @@ -930,9 +992,9 @@ export class UseAIServer { logger.info('Run aborted', { clientId: session.clientId, runId }); } - private sendEvent(socket: Socket, event: T) { - if (socket.connected) { - socket.emit('event', event); + private sendEvent(connection: ClientConnection, event: T) { + if (connection.connected) { + connection.emit('event', event); } } diff --git a/packages/server/src/types.ts b/packages/server/src/types.ts index a7bc5f75..e3d2380c 100644 --- a/packages/server/src/types.ts +++ b/packages/server/src/types.ts @@ -134,6 +134,16 @@ export interface UseAIServerConfig { + const client = new UseAIClient(new WebSocketTransport(`ws://localhost:${port}${path}`)); + const events: AGUIEvent[] = []; + client.onEvent('test', event => events.push(event)); + + return new Promise((resolve, reject) => { + const timeout = setTimeout(() => reject(new Error('Timed out connecting')), 5000); + const unsubscribe = client.onConnectionStateChange(connected => { + if (!connected) return; + clearTimeout(timeout); + unsubscribe(); + resolve({ client, events }); + }); + client.connect(); + }); +} + +async function waitFor(condition: () => boolean, message: string, timeoutMs = 5000): Promise { + const deadline = Date.now() + timeoutMs; + while (!condition()) { + if (Date.now() > deadline) throw new Error(`Timed out waiting for ${message}`); + await new Promise(resolve => setTimeout(resolve, 10)); + } +} + +describe.each(RUNTIMES)('WebSocketTransport against a real server: %s runtime', (runtime) => { + const cleanup = new TestCleanupManager(); + const port = runtime === 'bun' ? 9530 : 9540; + let server: UseAIServer; + + beforeAll(() => { + // Step 1 asks for the tool; step 2 answers with text once the result arrives. + const model = createSequentialMockModel([ + { toolCalls: [{ toolCallId: 'tc-1', toolName: 'addTodo', input: { text: 'buy groceries' } }] }, + { text: 'Added it.' }, + ]); + server = new UseAIServer({ + port, + runtime, + agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, + defaultAgent: 'test-agent', + }); + cleanup.trackServer(server); + }); + + afterAll(() => { + cleanup.cleanup(); + }); + + test('receives the agents list on connection', async () => { + const { client } = await connectClient(port); + + await waitFor(() => client.availableAgents.length > 0, 'the agents payload'); + + expect(client.availableAgents.map(a => a.id)).toEqual(['test-agent']); + expect(client.defaultAgent).toBe('test-agent'); + + client.disconnect(); + }); + + test('runs a full turn: prompt → tool call → tool result → RUN_FINISHED', async () => { + const { client, events } = await connectClient(port); + + client.registerTools([ + { + name: 'addTodo', + description: 'Add a todo item', + parameters: { + type: 'object', + properties: { text: { type: 'string' } }, + required: ['text'], + }, + }, + ]); + + await client.sendPrompt('Add a todo: buy groceries'); + + await waitFor( + () => events.some(e => e.type === EventType.TOOL_CALL_END), + 'the tool call', + ); + + const toolCallStart = events.find(e => e.type === EventType.TOOL_CALL_START) as + | { toolCallId: string; toolCallName: string } + | undefined; + expect(toolCallStart?.toolCallName).toBe('addTodo'); + + client.sendToolResponse(toolCallStart!.toolCallId, { success: true }); + + await waitFor( + () => events.some(e => e.type === EventType.RUN_FINISHED), + 'RUN_FINISHED', + ); + + const text = events + .filter(e => e.type === EventType.TEXT_MESSAGE_CONTENT) + .map(e => (e as { delta: string }).delta) + .join(''); + expect(text).toBe('Added it.'); + + // The conversation the client assembled from the stream is the ordinary one. + expect(client.messages.map(m => m.role)).toEqual(['user', 'assistant', 'tool', 'assistant']); + + client.disconnect(); + }); + + test('Socket.IO still serves the same port', async () => { + const socket = await cleanup.createTestClient(port); + expect(socket.connected).toBe(true); + socket.disconnect(); + }); + + test('two plain WebSocket clients get isolated sessions', async () => { + const first = await connectClient(port); + const second = await connectClient(port); + + await waitFor( + () => first.client.availableAgents.length > 0 && second.client.availableAgents.length > 0, + 'both agent payloads', + ); + + // Disconnecting one leaves the other usable. + first.client.disconnect(); + await waitFor(() => !first.client.isConnected(), 'the first client to close'); + + expect(second.client.isConnected()).toBe(true); + + second.client.disconnect(); + }); +}); + +describe('webSocketPath', () => { + const cleanup = new TestCleanupManager(); + + afterAll(() => { + cleanup.cleanup(); + }); + + test('serves the plain listener at a custom path', async () => { + const port = 9550; + const model = createSequentialMockModel([{ text: 'hi' }]); + cleanup.trackServer( + new UseAIServer({ + port, + webSocketPath: '/agent', + agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, + defaultAgent: 'test-agent', + }), + ); + + const { client } = await connectClient(port, '/agent'); + await waitFor(() => client.availableAgents.length > 0, 'the agents payload'); + + expect(client.defaultAgent).toBe('test-agent'); + client.disconnect(); + }); + + test('null serves Socket.IO only', async () => { + const port = 9560; + const model = createSequentialMockModel([{ text: 'hi' }]); + cleanup.trackServer( + new UseAIServer({ + port, + webSocketPath: null, + agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, + defaultAgent: 'test-agent', + }), + ); + + // Socket.IO is unaffected. + const socket = await cleanup.createTestClient(port); + expect(socket.connected).toBe(true); + socket.disconnect(); + + await expect(connectClient(port)).rejects.toThrow('Timed out connecting'); + }, 10_000); +}); From ba8470e9085dd3c22c8eb2a423a9e9db4c336c6e Mon Sep 17 00:00:00 2001 From: Zachary Davison Date: Fri, 4 Sep 2026 11:54:06 +0200 Subject: [PATCH 2/4] Address review: one listener, exclusive props, AG-UI framing, partysocket - UseAIProvider takes exactly one of serverUrl or transport, as a discriminated union on UseAIConfig. Context serverUrl becomes optional. - The server runs one listener: transport: 'socketio' | 'websocket', default 'socketio'. No webSocketPath. The plain listener serves at '/'. The Docker image reads TRANSPORT. - Downstream frames are bare AG-UI events. agents and config travel as AG-UI CUSTOM events, so the wire is AG-UI rather than a bespoke envelope. - WebSocketTransport reconnects through partysocket instead of its own loop. - Protocol doc rewritten. Co-Authored-By: Claude Opus 5 (1M context) --- CLAUDE.md | 2 +- README.md | 33 +-- apps/use-ai-server-app/src/index.ts | 3 + bun.lock | 5 + docs/websocket-protocol.md | 118 ++++----- packages/client/package.json | 1 + packages/client/src/client.test.ts | 77 ++++-- packages/client/src/client.ts | 2 +- packages/client/src/index.ts | 2 +- .../client/src/providers/useAIProvider.tsx | 33 +-- .../src/transport/WebSocketTransport.test.ts | 205 +++++++++------- .../src/transport/WebSocketTransport.ts | 230 +++++++----------- packages/client/src/transport/index.ts | 2 +- packages/client/src/types.ts | 35 ++- .../disabling-at-runtime.integration.test.tsx | 2 +- .../src/runtime/bun/BunRuntimeAdapter.ts | 156 ++++++------ .../server/src/runtime/bun/rawWebSocket.ts | 8 +- packages/server/src/runtime/index.ts | 2 +- .../src/runtime/node/NodeRuntimeAdapter.ts | 94 +++---- packages/server/src/runtime/types.ts | 25 +- packages/server/src/server.ts | 50 ++-- packages/server/src/types.ts | 14 +- packages/server/src/webSocketConnection.ts | 12 +- .../websocket-transport.integration.test.ts | 41 +--- 24 files changed, 559 insertions(+), 593 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 3e60fb0d..0f0b6bdd 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -61,7 +61,7 @@ bun run kill # Kill processes on ports 3000, 3002, 8081 Use `wss://your-domain.com` for secure WebSocket connections. For local development without SSL, use `ws://localhost:8081`. -The client reaches the server through a `UseAITransport`. `SocketIOTransport` is the default. `WebSocketTransport` carries JSON frames over a plain WebSocket, which the server serves at `webSocketPath` (default `/ws`) on the same port. The framing is documented in `docs/websocket-protocol.md`. +The client reaches the server through a `UseAITransport`. `SocketIOTransport` is the default. `WebSocketTransport` carries AG-UI events as JSON frames over a plain WebSocket. The server serves one or the other, chosen by `transport: 'socketio' | 'websocket'` (default `socketio`). The framing is documented in `docs/websocket-protocol.md`. ## Core Architecture diff --git a/README.md b/README.md index f5f112ba..ab2e35d7 100644 --- a/README.md +++ b/README.md @@ -302,39 +302,30 @@ There are some minor extensions to the protocol: The client reaches the server through a `UseAITransport`. Two transports ship with the library. -| Transport | Wire | Server endpoint | -| -------------------- | ---------------------------------------- | ------------------------------ | -| `SocketIOTransport` | Socket.IO, over polling and WebSocket | `/socket.io/` (the default) | -| `WebSocketTransport` | JSON text frames, over a plain WebSocket | `webSocketPath`, default `/ws` | +| Transport | Wire | Bundled server setting | +| -------------------- | ------------------------------------------------- | ------------------------ | +| `SocketIOTransport` | Socket.IO, over polling and WebSocket | `transport: 'socketio'` | +| `WebSocketTransport` | AG-UI events as JSON text, over a plain WebSocket | `transport: 'websocket'` | -`UseAIProvider` builds a `SocketIOTransport` from `serverUrl` when you do not pass one, so -nothing changes if you use the bundled server. +Give `UseAIProvider` either `serverUrl` or `transport`. `serverUrl` builds a `SocketIOTransport`, so nothing changes if you use the bundled server with its defaults. -Pass `WebSocketTransport` to reach a server that does not serve Socket.IO. Such a server -does not have to be Node. It must accept a WebSocket connection. It must then exchange -the documented frames. +Pass `WebSocketTransport` to reach a server that does not serve Socket.IO. Such a server does not have to be Node. It must accept a WebSocket connection. It must then exchange the documented frames. ```tsx import { UseAIProvider, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; root.render( - + ); ``` -The bundled server serves both listeners on one port. Set `webSocketPath: null` to serve -Socket.IO only. +The bundled server serves one transport. Set `transport: 'websocket'` on `UseAIServer`, or `TRANSPORT=websocket` on the Docker image, to serve a plain WebSocket at `/` instead of Socket.IO. -To carry the same messages over something else, implement `UseAITransport` yourself. The -interface has five members: `connect`, `disconnect`, `send`, `on` and `connected`. +To carry the same messages over something else, implement `UseAITransport` yourself. The interface has five members: `connect`, `disconnect`, `send`, `on` and `connected`. -See [docs/websocket-protocol.md](docs/websocket-protocol.md) for the frames, the turn -sequence, and the reconnection behaviour. +See [docs/websocket-protocol.md](docs/websocket-protocol.md) for the frames, the turn sequence, and the reconnection behaviour. ## Client @@ -392,7 +383,7 @@ root.render( ); ``` -Pass `transport` to reach a server over something other than Socket.IO. See [Transports](#transports). +Pass `transport` instead of `serverUrl` to reach a server over something other than Socket.IO. See [Transports](#transports). ### Component State via `prompt` @@ -1089,7 +1080,7 @@ const server = new UseAIServer({ }) }, defaultAgent: 'claude', - webSocketPath: '/ws', // plain WebSocket listener, see 'Transports'. null to disable. + transport: 'socketio', // or 'websocket', see 'Transports' rateLimitMaxRequests: 1_000, rateLimitWindowMs: 60_000, plugins: [ // see 'Plugins' diff --git a/apps/use-ai-server-app/src/index.ts b/apps/use-ai-server-app/src/index.ts index 4e596581..8ee64050 100644 --- a/apps/use-ai-server-app/src/index.ts +++ b/apps/use-ai-server-app/src/index.ts @@ -24,6 +24,7 @@ const maxHttpBufferSize = process.env.MAX_HTTP_BUFFER_SIZE const corsOrigin = process.env.CORS_ORIGIN; // Runtime adapter: 'auto' (default), 'bun', or 'node' const runtime = (process.env.RUNTIME as 'auto' | 'bun' | 'node') || 'auto'; +const transport = (process.env.TRANSPORT as 'socketio' | 'websocket') || 'socketio'; /** * Create agents based on available API keys. @@ -389,6 +390,7 @@ logger.info('Starting UseAI server', { logFormat }); } : undefined, runtime, + transport, }); // Initialize MCP endpoints @@ -401,6 +403,7 @@ logger.info('Starting UseAI server', { logFormat }); console.log(`✓ UseAI server is running on port ${port}`); console.log(` WebSocket URL: ws://localhost:${port}`); console.log(` Runtime: ${runtime} (set RUNTIME=bun or RUNTIME=node to change)`); + console.log(` Transport: ${transport} (set TRANSPORT=websocket for a plain WebSocket)`); console.log(` Log format: ${logFormat} (set LOG_FORMAT=json for structured logs)`); console.log(' Press Ctrl+C to stop'); } diff --git a/bun.lock b/bun.lock index 09a1ccb3..e7498d14 100644 --- a/bun.lock +++ b/bun.lock @@ -73,6 +73,7 @@ "version": "1.18.0", "dependencies": { "@meetsmore-oss/use-ai-core": "workspace:^", + "partysocket": "^1.3.0", "react-markdown": "^8.0.0", "remark-gfm": "3", "socket.io-client": "^4.8.1", @@ -772,6 +773,8 @@ "etag": ["etag@1.8.1", "", {}, "sha512-aIL5Fx7mawVa300al2BnEE4iNvo1qETxLrPI/o05L7z6go7fCw1J6EQmbK4FmJ2AS7kgVF/KEZWufBfdClMcPg=="], + "event-target-polyfill": ["event-target-polyfill@0.0.4", "", {}, "sha512-Gs6RLjzlLRdT8X9ZipJdIZI/Y6/HhRLyq9RdDlCsnpxr/+Nn6bU2EFGuC94GjxqhM+Nmij2Vcq98yoHrU8uNFQ=="], + "event-target-shim": ["event-target-shim@5.0.1", "", {}, "sha512-i/2XbnSz/uxRCU6+NdVJgKWDTM427+MqYbkQzD321DuCQJUqOuJKIA0IM2+W2xtYHdKOmZ4dR6fExsd4SXL+WQ=="], "events": ["events@3.3.0", "", {}, "sha512-mQw+2fkQbALzQ7V0MY0IqdnXNOeTtP4r0lN9z7AAawCXgqea7bDii20AYrIBrFd/Hx0M2Ocz6S111CaFkUcb0Q=="], @@ -1102,6 +1105,8 @@ "parseurl": ["parseurl@1.3.3", "", {}, "sha512-CiyeOxFT/JZyN5m0z9PfXw4SCBJ6Sygz1Dpl0wqjlhDEGGBP1GnsUVEL0p63hoG1fcj3fHynXi9NYO4nWOL+qQ=="], + "partysocket": ["partysocket@1.3.0", "", { "dependencies": { "event-target-polyfill": "^0.0.4" }, "peerDependencies": { "react": ">=17" }, "optionalPeers": ["react"] }, "sha512-1zToNyolZFK/7nuAw/K2bZrNzFqaZyRoCEkS+9vG6WSC5ikrN6qWRe96q6ImU51uptz2r+dAwSkwhJVdQi4LiA=="], + "passport": ["passport@0.7.0", "", { "dependencies": { "passport-strategy": "1.x.x", "pause": "0.0.1", "utils-merge": "^1.0.1" } }, "sha512-cPLl+qZpSc+ireUvt+IzqbED1cHHkDoVYMo30jbJIdOOjQ1MQYZBPiNvmi8UM6lJuOpTPXJGZQk0DtC4y61MYQ=="], "passport-azure-ad-oauth2": ["passport-azure-ad-oauth2@0.0.4", "", { "dependencies": { "passport-oauth": "1.0.x" } }, "sha512-yjwi0qXzGPIrR8yI5mBql2wO6tf/G5+HAFllkwwZ6f2EBCVvRv5z+6CwQeBvlrDbFh8RCXdj/IfB17r8LYDQQQ=="], diff --git a/docs/websocket-protocol.md b/docs/websocket-protocol.md index 5c2b37a0..0050db9e 100644 --- a/docs/websocket-protocol.md +++ b/docs/websocket-protocol.md @@ -1,125 +1,125 @@ # Plain WebSocket protocol -The bundled server serves Socket.IO by default. It also serves a plain WebSocket -listener on the same port, at `/ws`. A client reaches that listener with -`WebSocketTransport` instead of the default `SocketIOTransport`. +The client reaches the server through a `UseAITransport`. `SocketIOTransport` is the default. `WebSocketTransport` is the alternative. It carries AG-UI events as JSON text frames over a plain WebSocket. -Use this protocol to connect the `use-ai` chat UI and hooks to your own server. -Your server does not have to be Node. It does not have to implement Socket.IO. -Your server must accept a WebSocket connection. It must then exchange the JSON text -frames below. +Use `WebSocketTransport` to connect the `use-ai` chat UI and hooks to your own server. Your server does not have to be Node. It does not have to serve Socket.IO. It must accept a WebSocket connection. It must then exchange the frames that this document defines. -## Client setup +## Client ```tsx import { UseAIProvider, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; root.render( - + ); ``` -The provider reads `transport` once, on the first render. An inline object therefore -does not reconnect the client on every render. To change transports, remount the provider. +Give the provider `serverUrl` or `transport`, not both. `serverUrl` connects over Socket.IO. `transport` connects over the transport that you pass. -When you pass a transport, the provider does not use `serverUrl`. The prop stays -required. The provider reports `serverUrl` on the context for application code that -reads it. +The provider reads `transport` once, on the first render. An inline object therefore does not reconnect the client on each render. To change transports, remount the provider. -## Server setup +## Bundled server -The bundled server enables the listener by default: +The bundled server serves one transport. The default is Socket.IO. ```typescript const server = new UseAIServer({ agents: { claude }, defaultAgent: 'claude', - webSocketPath: '/agent', // default: '/ws'. Pass null to serve Socket.IO only. + transport: 'websocket', // default: 'socketio' }); ``` +With `transport: 'websocket'`, the server accepts WebSocket upgrades at `/`. It does not serve Socket.IO. The `/health` endpoint is unchanged. + +The Docker image reads the same setting from the `TRANSPORT` environment variable. + +## Encoding + +Each frame is one JSON object, as a text frame. JSON is the encoding that AG-UI, MCP and the OpenAI Realtime API use on their wires. AG-UI also defines a protobuf encoding for bandwidth. This protocol does not use it. + ## Upstream frames -The client sends the `UseAIClientMessage` object, serialized, with nothing wrapped -around it. Each message is one text frame. +The client sends each `UseAIClientMessage` as one frame, with nothing around it. ```json { "type": "run_agent", "data": { "threadId": "...", "runId": "...", "messages": [], "tools": [], "state": null, "forwardedProps": {} } } ``` -The message types are `run_agent`, `tool_result`, `tool_approval_response`, -`abort_run` and `message_feedback`. Plugins add more. See `UseAIClientMessage` in -`@meetsmore-oss/use-ai-core` for each payload. +The message types are: + +- `run_agent` +- `tool_result` +- `tool_approval_response` +- `abort_run` +- `message_feedback` + +Plugins add more. See `UseAIClientMessage` in `@meetsmore-oss/use-ai-core` for each payload. ## Downstream frames -A plain WebSocket has no event names of its own, so the server wraps each payload in -a named envelope. Each envelope is one text frame. +The server sends one AG-UI event per frame. Each event has a `type` field. ```json -{ "name": "event", "data": { "type": "TEXT_MESSAGE_CONTENT", "messageId": "...", "delta": "Hello" } } -{ "name": "agents", "data": { "agents": [{ "id": "claude", "name": "Claude" }], "defaultAgent": "claude" } } -{ "name": "config", "data": { "langfuseEnabled": true } } +{ "type": "RUN_STARTED", "threadId": "...", "runId": "...", "timestamp": 1700000000000 } +{ "type": "TEXT_MESSAGE_CONTENT", "messageId": "...", "delta": "Hello" } +{ "type": "RUN_FINISHED", "threadId": "...", "runId": "..." } ``` -| Name | Payload | When | -| -------- | -------------------------------------------------- | ------------------------------------------------------ | -| `agents` | The agent list and the default agent id | Once, after connect | -| `config` | Server capability flags, such as `langfuseEnabled` | Once, after connect, if the server has flags to report | -| `event` | One AG-UI event | Throughout a run | +Two payloads are not AG-UI events. The server sends them as AG-UI `CUSTOM` events, once, after the connection opens. + +```json +{ "type": "CUSTOM", "name": "agents", "value": { "agents": [{ "id": "claude", "name": "Claude" }], "defaultAgent": "claude" } } +{ "type": "CUSTOM", "name": "config", "value": { "langfuseEnabled": true } } +``` + +| Name | Value | Required | +| -------- | -------------------------------------------------- | -------- | +| `agents` | The agent list and the default agent id | Yes | +| `config` | Capability flags, such as `langfuseEnabled` | No | + +The client ignores an event type that it does not handle. It also ignores a `CUSTOM` name that it does not know. A server can therefore add events without a change to older clients. -The client reads `name` and finds the handlers for that name. It then passes `data` -to each handler. +The client handles these event types: -The client ignores a frame with an unknown `name`. An unknown name is not an error. -The client does not close the connection. A server can therefore send a name that an -older client does not know. +- `RUN_STARTED`, `RUN_FINISHED`, `RUN_ERROR` +- `STEP_STARTED`, `STEP_FINISHED` +- `TEXT_MESSAGE_START`, `TEXT_MESSAGE_CONTENT`, `TEXT_MESSAGE_END` +- `TOOL_CALL_START`, `TOOL_CALL_ARGS`, `TOOL_CALL_END`, `TOOL_CALL_RESULT` +- `REASONING_MESSAGE_START`, `REASONING_MESSAGE_CONTENT`, `REASONING_MESSAGE_END`, `REASONING_ENCRYPTED_VALUE` +- `TOOL_APPROVAL_REQUEST`, a `use-ai` extension -The `event` payload is an AG-UI event. The event types the client handles are -`RUN_STARTED`, `RUN_FINISHED`, `RUN_ERROR`, `STEP_STARTED`, `STEP_FINISHED`, -`TEXT_MESSAGE_START`, `TEXT_MESSAGE_CONTENT`, `TEXT_MESSAGE_END`, `TOOL_CALL_START`, -`TOOL_CALL_ARGS`, `TOOL_CALL_END`, `TOOL_CALL_RESULT`, the `REASONING_*` events, and -the `TOOL_APPROVAL_REQUEST` extension. See the -[AG-UI protocol](https://docs.ag-ui.com/introduction) for each event, and -`packages/core/src/types.ts` for the types this library uses. +See the [AG-UI protocol](https://docs.ag-ui.com/introduction) for each event. See `packages/core/src/types.ts` for the types that this library uses. ## One turn, step by step 1. The client opens the connection. 2. The server sends `agents`. It then sends `config`. 3. The client sends `run_agent` with the prompt, the tool definitions and the app state. -4. The server sends `event` frames for `RUN_STARTED`, then the model output. +4. The server sends `RUN_STARTED`. It then streams the model output. 5. For a client-side tool, the server sends `TOOL_CALL_START`, `TOOL_CALL_ARGS` and `TOOL_CALL_END`. 6. The client runs the tool. It then sends `tool_result` with the output. -7. The server resumes the model. It then sends the remaining `event` frames. +7. The server resumes the model. It then streams the rest of the output. 8. The server sends `RUN_FINISHED`. ## Reconnection -A plain WebSocket has no reconnection. `WebSocketTransport` therefore runs its own -retry loop. It retries indefinitely, with exponential backoff capped at ten seconds. -The cap matches the Socket.IO settings. A mobile app in the background, or a device -in airplane mode, thus recovers without frequent retries. +`WebSocketTransport` reconnects through [partysocket](https://github.com/partykit/partykit/tree/main/packages/partysocket). It retries indefinitely. The delay doubles after each attempt, from one second up to ten seconds. The limits match `SocketIOTransport`, so a mobile app in the background, or a device in airplane mode, recovers without frequent retries. Set both delays in the options: ```typescript -new WebSocketTransport('wss://your-server.com/ws', { +new WebSocketTransport('wss://your-server.com', { reconnectionDelay: 1000, // first retry, in milliseconds - reconnectionDelayMax: 10000, // cap on the backoff, in milliseconds + reconnectionDelayMax: 10000, // upper bound, in milliseconds }); ``` -The server destroys the session when the connection closes. A reconnected client -therefore starts a new session. The client sends its conversation history with the -next `run_agent`. +The server destroys the session when the connection closes. A reconnected client therefore starts a new session. The client sends its conversation history with the next `run_agent`. -## Writing your own transport +## Your own transport `UseAITransport` has five members: diff --git a/packages/client/package.json b/packages/client/package.json index f9dea4cc..e9afdde9 100644 --- a/packages/client/package.json +++ b/packages/client/package.json @@ -34,6 +34,7 @@ }, "dependencies": { "@meetsmore-oss/use-ai-core": "workspace:^", + "partysocket": "^1.3.0", "react-markdown": "^8.0.0", "remark-gfm": "3", "socket.io-client": "^4.8.1", diff --git a/packages/client/src/client.test.ts b/packages/client/src/client.test.ts index d7274c80..34727abe 100644 --- a/packages/client/src/client.test.ts +++ b/packages/client/src/client.test.ts @@ -1,7 +1,6 @@ import { describe, test, expect, mock, beforeEach, afterEach, spyOn } from 'bun:test'; import type { Socket } from 'socket.io-client'; import type { UseAIClientMessage } from './types'; -import type { WebSocketLike } from './transport/WebSocketTransport'; /** * The UseAIClient suite, run over every bundled transport. @@ -42,18 +41,40 @@ mock.module('socket.io-client', () => ({ // ── Plain WebSocket harness ───────────────────────────────────────────────── -class FakeWebSocket implements WebSocketLike { +/** Enough of a WHATWG WebSocket for partysocket to drive. */ +class FakeWebSocket extends EventTarget { + static latest: FakeWebSocket | null = null; + + readyState = 0; + binaryType = 'blob'; sent: string[] = []; - onopen: ((event: unknown) => void) | null = null; - onmessage: ((event: { data: unknown }) => void) | null = null; - onclose: ((event: unknown) => void) | null = null; - onerror: ((event: unknown) => void) | null = null; + + constructor() { + super(); + FakeWebSocket.latest = this; + } send(data: string): void { this.sent.push(data); } - close(): void {} + close(): void { + this.readyState = 3; + } + + serverOpen(): void { + this.readyState = 1; + this.dispatchEvent(new Event('open')); + } + + serverSend(data: string): void { + this.dispatchEvent(new MessageEvent('message', { data })); + } + + serverClose(): void { + this.readyState = 3; + this.dispatchEvent(new CloseEvent('close', { code: 1006 })); + } } // Imported after the module mock so SocketIOTransport picks it up. @@ -123,37 +144,41 @@ const HARNESSES: Array<[string, () => Harness]> = [ [ 'WebSocketTransport', () => { - let socket!: FakeWebSocket; + FakeWebSocket.latest = null; const client = new UseAIClient( - new WebSocketTransport('wss://localhost:8081/ws', { + new WebSocketTransport('wss://localhost:8081', { reconnectionDelay: 1, reconnectionDelayMax: 1, - createWebSocket: () => (socket = new FakeWebSocket()), + WebSocket: FakeWebSocket as unknown as typeof WebSocket, }), ); client.connect(); + // partysocket opens the socket asynchronously, and opens a fresh one after a drop. + const liveSocket = async () => { + await waitUntil(() => FakeWebSocket.latest !== null && FakeWebSocket.latest.readyState < 2); + return FakeWebSocket.latest!; + }; + const sent: string[] = []; return { client, async open() { - const dropped = socket; - // A closed socket is detached; the replacement arrives on the backoff timer. - if (dropped.onopen === null) { - await waitUntil(() => socket !== dropped); - } - socket.onopen?.({}); + const socket = await liveSocket(); + socket.sent = sent; + socket.serverOpen(); }, close(reason: string) { // The reason a plain WebSocket reports comes from the close frame, not // from the caller; the client only logs it. void reason; - socket.onclose?.({}); + FakeWebSocket.latest?.serverClose(); }, deliver(name, data) { - socket.onmessage?.({ data: JSON.stringify({ name, data }) }); + const frame = name === 'event' ? data : { type: 'CUSTOM', name, value: data }; + FakeWebSocket.latest?.serverSend(JSON.stringify(frame)); }, sent() { - return socket.sent.map(frame => JSON.parse(frame) as UseAIClientMessage); + return sent.map(frame => JSON.parse(frame) as UseAIClientMessage); }, }; }, @@ -767,22 +792,24 @@ describe('UseAIClient construction', () => { client.disconnect(); }); - test('disconnect() unsubscribes from the transport', () => { - let socket!: FakeWebSocket; + test('disconnect() unsubscribes from the transport', async () => { + FakeWebSocket.latest = null; const client = new UseAIClient( - new WebSocketTransport('wss://localhost:8081/ws', { - createWebSocket: () => (socket = new FakeWebSocket()), + new WebSocketTransport('wss://localhost:8081', { + WebSocket: FakeWebSocket as unknown as typeof WebSocket, }), ); const stateChanges: boolean[] = []; client.onConnectionStateChange(connected => stateChanges.push(connected)); client.connect(); - socket.onopen?.({}); + await waitUntil(() => FakeWebSocket.latest !== null); + const socket = FakeWebSocket.latest!; + socket.serverOpen(); client.disconnect(); // The transport is closed, but a late frame from the old socket must not // reach a client that has stopped listening. - socket.onclose?.({}); + socket.serverClose(); expect(stateChanges).toEqual([false, true]); }); diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index 893cd270..14def423 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -112,7 +112,7 @@ export class UseAIClient { * @example * ```typescript * new UseAIClient('wss://your-server.com'); - * new UseAIClient(new WebSocketTransport('wss://your-server.com/ws')); + * new UseAIClient(new WebSocketTransport('wss://your-server.com')); * ``` */ constructor(target: string | UseAITransport) { diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts index e21f1ada..51932c7f 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -8,7 +8,6 @@ export type { UseAITransportEventName, SocketIOTransportOptions, WebSocketTransportOptions, - WebSocketLike, } from './transport'; export { defineTool, executeDefinedTool, convertToolsToDefinitions } from './defineTool'; /** @hidden */ @@ -81,6 +80,7 @@ export type { FloatingButtonProps, ChatPanelProps, UseAIProviderProps, + UseAIProviderOptions, } from './providers/useAIProvider'; export type { SendMessageOptions } from './hooks/useMessageQueue'; export type { DefinedTool, ToolsDefinition, ToolOptions, ToolAnnotations, ToolExecutionContext } from './defineTool'; diff --git a/packages/client/src/providers/useAIProvider.tsx b/packages/client/src/providers/useAIProvider.tsx index 2aa7637b..e1d430a6 100644 --- a/packages/client/src/providers/useAIProvider.tsx +++ b/packages/client/src/providers/useAIProvider.tsx @@ -5,7 +5,6 @@ import { UseAIChatPanel } from '../components/UseAIChatPanel'; import { UseAIFloatingChatWrapper, CloseButton } from '../components/UseAIFloatingChatWrapper'; import { __UseAIChatContext, type ChatUIContextValue } from '../components/UseAIChat'; import { UseAIClient } from '../client'; -import type { UseAITransport } from '../transport/types'; import { convertToolsToDefinitions, type ToolsDefinition } from '../defineTool'; import type { ChatRepository, Chat, ChatMetadata, CreateChatOptions, PersistedMessage, PersistedMessageContent } from './chatRepository/types'; import { LocalStorageChatRepository } from './chatRepository/LocalStorageChatRepository'; @@ -123,8 +122,8 @@ export interface PromptsContextValue { * Contains connection state and methods for managing tools and prompts. */ export interface UseAIContextValue { - /** The WebSocket URL of the UseAI server */ - serverUrl: string; + /** URL of the server, when the provider was given `serverUrl` rather than `transport`. */ + serverUrl?: string; /** Whether the client is connected to the server */ connected: boolean; /** The underlying WebSocket client instance */ @@ -163,7 +162,6 @@ let hasWarnedAboutMissingProvider = false; * Allows hooks to gracefully degrade instead of crashing. */ const noOpContextValue: UseAIContextValue = { - serverUrl: '', connected: false, client: null, tools: { @@ -261,27 +259,11 @@ export interface ChatPanelProps { onAgentChange?: (agentId: string | null) => void; } -export interface UseAIProviderProps extends UseAIConfig { +export type UseAIProviderProps = UseAIConfig & UseAIProviderOptions; + +export interface UseAIProviderOptions { children: ReactNode; systemPrompt?: string; - /** - * Transport used to reach the server. Defaults to a {@link SocketIOTransport} built - * from `serverUrl`, which is what the bundled server serves. Supply one to reach a - * server that speaks something else — {@link WebSocketTransport} carries documented - * JSON frames over a plain WebSocket. - * - * Only the value from the first render is read, so an inline object does not churn - * the connection. Change transports by remounting the provider. - * - * @example - * ```tsx - * - * ``` - */ - transport?: UseAITransport; CustomButton?: React.ComponentType | null; CustomChat?: React.ComponentType | null; /** Default component overrides for every built-in chat rendered by this provider. */ @@ -598,8 +580,9 @@ export function UseAIProvider({ const transportRef = useRef(transport); useEffect(() => { - console.log('[UseAIProvider] Initializing client with serverUrl:', serverUrl); - const client = new UseAIClient(transportRef.current ?? serverUrl); + const target = transportRef.current ?? serverUrl!; + console.log('[UseAIProvider] Initializing client with', typeof target === 'string' ? target : 'transport'); + const client = new UseAIClient(target); const unsubscribeConnection = client.onConnectionStateChange((isConnected) => { console.log('[UseAIProvider] Connection state changed:', isConnected); diff --git a/packages/client/src/transport/WebSocketTransport.test.ts b/packages/client/src/transport/WebSocketTransport.test.ts index f14dfa8a..2705d430 100644 --- a/packages/client/src/transport/WebSocketTransport.test.ts +++ b/packages/client/src/transport/WebSocketTransport.test.ts @@ -1,17 +1,17 @@ import { describe, test, expect, beforeEach, afterEach, spyOn } from 'bun:test'; -import { WebSocketTransport, type WebSocketLike } from './WebSocketTransport'; +import { WebSocketTransport } from './WebSocketTransport'; -class FakeWebSocket implements WebSocketLike { +/** Enough of a WHATWG WebSocket for partysocket to drive. */ +class FakeWebSocket extends EventTarget { static instances: FakeWebSocket[] = []; + readyState = 0; + binaryType = 'blob'; sent: string[] = []; closeCalls = 0; - onopen: ((event: unknown) => void) | null = null; - onmessage: ((event: { data: unknown }) => void) | null = null; - onclose: ((event: unknown) => void) | null = null; - onerror: ((event: unknown) => void) | null = null; constructor(readonly url: string) { + super(); FakeWebSocket.instances.push(this); } @@ -25,32 +25,50 @@ class FakeWebSocket implements WebSocketLike { close(): void { this.closeCalls++; + this.readyState = 3; } - /** Simulates the server accepting the connection. */ serverOpen(): void { - this.onopen?.({}); + this.readyState = 1; + this.dispatchEvent(new Event('open')); } - /** Simulates a frame arriving from the server. */ serverSend(data: unknown): void { - this.onmessage?.({ data }); + this.dispatchEvent(new MessageEvent('message', { data })); } - /** Simulates the connection dropping. */ serverClose(): void { - this.onclose?.({}); + this.readyState = 3; + this.dispatchEvent(new CloseEvent('close', { code: 1006 })); + } +} + +const tick = (ms = 0) => new Promise(resolve => setTimeout(resolve, ms)); + +async function waitUntil(condition: () => boolean, timeoutMs = 500): Promise { + const deadline = Date.now() + timeoutMs; + while (!condition()) { + if (Date.now() > deadline) throw new Error('Timed out'); + await tick(1); } } function makeTransport(options: { reconnectionDelay?: number; reconnectionDelayMax?: number } = {}) { - return new WebSocketTransport('wss://server.example/ws', { + return new WebSocketTransport('wss://server.example', { ...options, - createWebSocket: (url) => new FakeWebSocket(url), + WebSocket: FakeWebSocket as unknown as typeof WebSocket, }); } -const tick = (ms: number) => new Promise(resolve => setTimeout(resolve, ms)); +/** Connects and waits for the underlying socket, which partysocket opens asynchronously. */ +async function connect(transport: WebSocketTransport): Promise { + const before = FakeWebSocket.instances.length; + transport.connect(); + await waitUntil(() => FakeWebSocket.instances.length > before); + return FakeWebSocket.latest; +} + +const customEvent = (name: string, value: unknown) => JSON.stringify({ type: 'CUSTOM', name, value }); describe('WebSocketTransport', () => { let consoleLogSpy: ReturnType; @@ -67,47 +85,60 @@ describe('WebSocketTransport', () => { consoleWarnSpy.mockRestore(); }); - test('opens a socket at the configured url', () => { + test('opens a socket at the configured url', async () => { + const transport = makeTransport(); + const socket = await connect(transport); + + expect(socket.url).toBe('wss://server.example'); + + transport.disconnect(); + }); + + test('an AG-UI event frame is delivered on the event channel', async () => { const transport = makeTransport(); - transport.connect(); + const events: unknown[] = []; + transport.on('event', data => events.push(data)); - expect(FakeWebSocket.instances).toHaveLength(1); - expect(FakeWebSocket.latest.url).toBe('wss://server.example/ws'); + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(JSON.stringify({ type: 'RUN_STARTED', threadId: 't', runId: 'r' })); + + expect(events).toEqual([{ type: 'RUN_STARTED', threadId: 't', runId: 'r' }]); transport.disconnect(); }); - test('an incoming frame calls the handlers for its name', () => { + test('CUSTOM events named agents and config are delivered on their own channels', async () => { const transport = makeTransport(); const events: unknown[] = []; + const agents: unknown[] = []; const configs: unknown[] = []; transport.on('event', data => events.push(data)); + transport.on('agents', data => agents.push(data)); transport.on('config', data => configs.push(data)); - transport.connect(); - FakeWebSocket.latest.serverOpen(); + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(customEvent('agents', { agents: [], defaultAgent: 'claude' })); + socket.serverSend(customEvent('config', { langfuseEnabled: true })); - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: { type: 'RUN_STARTED' } })); - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'config', data: { langfuseEnabled: true } })); - - expect(events).toEqual([{ type: 'RUN_STARTED' }]); + expect(agents).toEqual([{ agents: [], defaultAgent: 'claude' }]); expect(configs).toEqual([{ langfuseEnabled: true }]); + expect(events).toEqual([]); transport.disconnect(); }); - test('delivers a frame to every subscriber of its name', () => { + test('delivers a frame to every subscriber of its channel', async () => { const transport = makeTransport(); const first: unknown[] = []; const second: unknown[] = []; transport.on('agents', data => first.push(data)); transport.on('agents', data => second.push(data)); - transport.connect(); - FakeWebSocket.latest.serverOpen(); - FakeWebSocket.latest.serverSend( - JSON.stringify({ name: 'agents', data: { agents: [], defaultAgent: 'claude' } }), - ); + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(customEvent('agents', { agents: [], defaultAgent: 'claude' })); expect(first).toEqual([{ agents: [], defaultAgent: 'claude' }]); expect(second).toEqual([{ agents: [], defaultAgent: 'claude' }]); @@ -115,141 +146,136 @@ describe('WebSocketTransport', () => { transport.disconnect(); }); - test('unsubscribing stops delivery', () => { + test('unsubscribing stops delivery', async () => { const transport = makeTransport(); const events: unknown[] = []; const unsubscribe = transport.on('event', data => events.push(data)); - transport.connect(); - FakeWebSocket.latest.serverOpen(); - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 1 })); + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(JSON.stringify({ type: 'STEP_STARTED', stepName: '1' })); unsubscribe(); - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 2 })); + socket.serverSend(JSON.stringify({ type: 'STEP_STARTED', stepName: '2' })); - expect(events).toEqual([1]); + expect(events).toEqual([{ type: 'STEP_STARTED', stepName: '1' }]); transport.disconnect(); }); - test('a frame with an unknown name is ignored, not an error', () => { + test('a CUSTOM event with an unknown name passes through as an event', async () => { const transport = makeTransport(); const events: unknown[] = []; transport.on('event', data => events.push(data)); - transport.connect(); - FakeWebSocket.latest.serverOpen(); - - expect(() => { - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'a_name_from_a_later_version', data: {} })); - }).not.toThrow(); + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(customEvent('a_name_from_a_later_version', {})); - // The connection survives, so the next known frame still arrives. - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 'still here' })); - expect(events).toEqual(['still here']); + // UseAIClient ignores event types it does not handle, so passing it on is safe. + expect(events).toEqual([{ type: 'CUSTOM', name: 'a_name_from_a_later_version', value: {} }]); expect(transport.connected).toBe(true); transport.disconnect(); }); - test('a malformed frame is ignored, not an error', () => { + test('a malformed frame is ignored, not an error', async () => { const transport = makeTransport(); const events: unknown[] = []; transport.on('event', data => events.push(data)); - transport.connect(); - FakeWebSocket.latest.serverOpen(); + const socket = await connect(transport); + socket.serverOpen(); - expect(() => FakeWebSocket.latest.serverSend('not json')).not.toThrow(); - expect(() => FakeWebSocket.latest.serverSend(JSON.stringify({ noNameHere: true }))).not.toThrow(); + expect(() => socket.serverSend('not json')).not.toThrow(); + expect(() => socket.serverSend(JSON.stringify({ noTypeHere: true }))).not.toThrow(); + expect(() => socket.serverSend(new ArrayBuffer(4))).not.toThrow(); - FakeWebSocket.latest.serverSend(JSON.stringify({ name: 'event', data: 'still here' })); - expect(events).toEqual(['still here']); + socket.serverSend(JSON.stringify({ type: 'RUN_FINISHED' })); + expect(events).toEqual([{ type: 'RUN_FINISHED' }]); + expect(transport.connected).toBe(true); transport.disconnect(); }); - test('send serializes the message with nothing wrapped around it', () => { + test('send serializes the message with nothing wrapped around it', async () => { const transport = makeTransport(); - transport.connect(); - FakeWebSocket.latest.serverOpen(); + const socket = await connect(transport); + socket.serverOpen(); transport.send({ type: 'abort_run', data: { runId: 'run-1' } }); - expect(FakeWebSocket.latest.sent).toEqual(['{"type":"abort_run","data":{"runId":"run-1"}}']); + expect(socket.sent).toEqual(['{"type":"abort_run","data":{"runId":"run-1"}}']); transport.disconnect(); }); - test('connected follows the socket opening and closing', () => { + test('connected follows the socket opening and closing', async () => { const transport = makeTransport({ reconnectionDelay: 10_000 }); expect(transport.connected).toBe(false); - transport.connect(); + const socket = await connect(transport); expect(transport.connected).toBe(false); - FakeWebSocket.latest.serverOpen(); + socket.serverOpen(); expect(transport.connected).toBe(true); - FakeWebSocket.latest.serverClose(); + socket.serverClose(); expect(transport.connected).toBe(false); transport.disconnect(); }); - test('a close after opening dispatches disconnect', () => { + test('a close after opening dispatches disconnect', async () => { const transport = makeTransport({ reconnectionDelay: 10_000 }); const states: string[] = []; transport.on('connect', () => states.push('connect')); transport.on('disconnect', () => states.push('disconnect')); - transport.connect(); - FakeWebSocket.latest.serverOpen(); - FakeWebSocket.latest.serverClose(); + const socket = await connect(transport); + socket.serverOpen(); + socket.serverClose(); expect(states).toEqual(['connect', 'disconnect']); transport.disconnect(); }); - test('a failed connection attempt does not dispatch disconnect', () => { + test('a failed connection attempt does not dispatch disconnect', async () => { const transport = makeTransport({ reconnectionDelay: 10_000 }); const states: string[] = []; transport.on('disconnect', () => states.push('disconnect')); - transport.connect(); + const socket = await connect(transport); // Never opened: the socket closes straight from the connecting state. - FakeWebSocket.latest.serverClose(); + socket.serverClose(); expect(states).toEqual([]); transport.disconnect(); }); - test('backoff reconnects after the socket drops', async () => { + test('reconnects after the socket drops', async () => { const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); - transport.connect(); - FakeWebSocket.latest.serverOpen(); - expect(FakeWebSocket.instances).toHaveLength(1); + const socket = await connect(transport); + socket.serverOpen(); - FakeWebSocket.latest.serverClose(); - await tick(20); + socket.serverClose(); + await waitUntil(() => FakeWebSocket.instances.length > 1); - expect(FakeWebSocket.instances.length).toBeGreaterThan(1); - - // The reconnected socket is live: opening it restores connected. FakeWebSocket.latest.serverOpen(); expect(transport.connected).toBe(true); transport.disconnect(); }); - test('backoff keeps retrying while attempts fail', async () => { + test('keeps retrying while attempts fail', async () => { const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); - transport.connect(); + await connect(transport); for (let i = 0; i < 3; i++) { + const count = FakeWebSocket.instances.length; FakeWebSocket.latest.serverClose(); - await tick(10); + await waitUntil(() => FakeWebSocket.instances.length > count); } expect(FakeWebSocket.instances.length).toBeGreaterThanOrEqual(4); @@ -257,12 +283,12 @@ describe('WebSocketTransport', () => { transport.disconnect(); }); - test('disconnect() stops the backoff', async () => { + test('disconnect() stops the retries', async () => { const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); - transport.connect(); - FakeWebSocket.latest.serverOpen(); + const socket = await connect(transport); + socket.serverOpen(); - FakeWebSocket.latest.serverClose(); + socket.serverClose(); transport.disconnect(); const openedByNow = FakeWebSocket.instances.length; @@ -272,12 +298,11 @@ describe('WebSocketTransport', () => { expect(transport.connected).toBe(false); }); - test('disconnect() closes the open socket', () => { + test('disconnect() closes the open socket', async () => { const transport = makeTransport(); - transport.connect(); - FakeWebSocket.latest.serverOpen(); + const socket = await connect(transport); + socket.serverOpen(); - const socket = FakeWebSocket.latest; transport.disconnect(); expect(socket.closeCalls).toBe(1); diff --git a/packages/client/src/transport/WebSocketTransport.ts b/packages/client/src/transport/WebSocketTransport.ts index 521a2d91..3b0180b7 100644 --- a/packages/client/src/transport/WebSocketTransport.ts +++ b/packages/client/src/transport/WebSocketTransport.ts @@ -1,108 +1,71 @@ +import ReconnectingWebSocket from 'partysocket/ws'; +import { EventType } from '@meetsmore-oss/use-ai-core'; import type { UseAIClientMessage } from '../types'; import { TransportHandlerRegistry } from './handlerRegistry'; import type { UseAITransport, UseAITransportEventName } from './types'; -/** - * The subset of the WHATWG `WebSocket` API that {@link WebSocketTransport} uses. - * Declared structurally so a test double, or a polyfill on a runtime without a - * global `WebSocket`, can stand in for the real thing. - */ -export interface WebSocketLike { - send(data: string): void; - close(): void; - onopen: ((event: unknown) => void) | null; - onmessage: ((event: { data: unknown }) => void) | null; - onclose: ((event: unknown) => void) | null; - onerror: ((event: unknown) => void) | null; -} - -/** - * A downstream frame, as sent by the server. - * - * @example - * ```json - * { "name": "config", "data": { "langfuseEnabled": true } } - * ``` - */ -interface DownstreamFrame { - name: string; - data: unknown; -} - /** * Options for {@link WebSocketTransport}. */ export interface WebSocketTransportOptions { /** * Delay before the first reconnection attempt, in milliseconds. - * Subsequent attempts double this, up to {@link reconnectionDelayMax}. + * Each later attempt doubles the delay, up to {@link reconnectionDelayMax}. * * @default 1000 */ reconnectionDelay?: number; /** - * Upper bound on the exponential backoff between reconnection attempts, in milliseconds. + * Upper bound on the delay between reconnection attempts, in milliseconds. * * @default 10000 */ reconnectionDelayMax?: number; /** - * Opens the underlying socket. + * WebSocket constructor to open the connection with. + * Supply one on a runtime without a global `WebSocket`, or in a test. * - * @default (url) => new WebSocket(url) + * @default globalThis.WebSocket */ - createWebSocket?: (url: string) => WebSocketLike; + WebSocket?: typeof WebSocket; } /** - * Transport over a plain WebSocket carrying JSON text frames. + * Transport over a plain WebSocket. Every frame is JSON text. * - * Use this to reach a server that does not speak Socket.IO. The framing is: + * Upstream, the client sends each `UseAIClientMessage` as one frame, with nothing + * around it. Downstream, the server sends one AG-UI event per frame. The `agents` + * and `config` payloads travel as AG-UI `CUSTOM` events named `agents` and `config`. + * The client ignores an event with a type or a custom name it does not know. * - * - **Upstream** — the `UseAIClientMessage`, serialized, with nothing wrapped around it: - * `{"type":"run_agent","data":{...}}` - * - **Downstream** — a named envelope, because a plain WebSocket has no event names of its own: - * `{"name":"event","data":{...}}`. The names are `event`, `agents` and `config`. - * A frame with any other name is ignored, so a server may add names without breaking - * older clients. - * - * A server should send `agents` and `config` once, after the connection opens. + * Reconnection is automatic: indefinite, with exponential backoff capped at + * `reconnectionDelayMax`. The defaults match {@link SocketIOTransport}. * * @example - * ```typescript - * + * ```tsx + * * ``` */ export class WebSocketTransport implements UseAITransport { - private socket: WebSocketLike | null = null; + private socket: ReconnectingWebSocket | null = null; private registry = new TransportHandlerRegistry(); private _connected = false; - // A plain WebSocket has no reconnection of its own, so this transport matches - // the Socket.IO settings: retry indefinitely with exponential backoff capped at - // reconnectionDelayMax, so a client recovers after an extended outage (mobile app - // backgrounded, airplane mode) without hammering the server in the meantime. - private reconnectionDelay: number; - private reconnectionDelayMax: number; - private reconnectAttempts = 0; - private reconnectTimer: ReturnType | null = null; - private reconnecting = false; - private createWebSocket: (url: string) => WebSocketLike; + private readonly options: Required> & + Pick; /** * @param url - WebSocket URL of the server * @example * ```typescript - * new WebSocketTransport('wss://your-server.com/ws'); + * new WebSocketTransport('wss://your-server.com'); * ``` */ constructor(private url: string, options: WebSocketTransportOptions = {}) { - this.reconnectionDelay = options.reconnectionDelay ?? 1000; - this.reconnectionDelayMax = options.reconnectionDelayMax ?? 10_000; - this.createWebSocket = - options.createWebSocket ?? ((url: string) => new WebSocket(url) as unknown as WebSocketLike); + this.options = { + reconnectionDelay: options.reconnectionDelay ?? 1000, + reconnectionDelayMax: options.reconnectionDelayMax ?? 10_000, + WebSocket: options.WebSocket, + }; } get connected(): boolean { @@ -110,115 +73,86 @@ export class WebSocketTransport implements UseAITransport { } connect(): void { - this.reconnecting = true; - this.open(); - } - - disconnect(): void { - this.reconnecting = false; - if (this.reconnectTimer !== null) { - clearTimeout(this.reconnectTimer); - this.reconnectTimer = null; - } - - const socket = this.socket; - this.socket = null; - this._connected = false; - if (socket) { - this.detach(socket); - socket.close(); - } - } - - send(message: UseAIClientMessage): void { - this.socket?.send(JSON.stringify(message)); - } - - on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void { - return this.registry.on(name, handler); - } - - private open(): void { - let socket: WebSocketLike; - try { - socket = this.createWebSocket(this.url); - } catch (error) { - // Use warn instead of error to avoid triggering Next.js error overlay - console.warn('[UseAI] Connection error:', error instanceof Error ? error.message : error); - this.scheduleReconnect(); - return; - } + if (this.socket) return; + + const socket = new ReconnectingWebSocket(this.url, [], { + WebSocket: this.options.WebSocket, + minReconnectionDelay: this.options.reconnectionDelay, + maxReconnectionDelay: this.options.reconnectionDelayMax, + reconnectionDelayGrowFactor: 2, + maxRetries: Infinity, + // UseAIClient only sends while connected, so nothing is queued for a later socket. + maxEnqueuedMessages: 0, + }); this.socket = socket; socket.onopen = () => { this._connected = true; - this.reconnectAttempts = 0; this.registry.dispatch('connect', undefined); }; socket.onmessage = (event) => { - const frame = this.parseFrame(event.data); + const frame = parseFrame(event.data); if (!frame) return; - this.registry.dispatch(frame.name, frame.data); + if (frame.type === EventType.CUSTOM && (frame.name === 'agents' || frame.name === 'config')) { + this.registry.dispatch(frame.name, frame.value); + return; + } + this.registry.dispatch('event', frame); }; - socket.onerror = () => { - // onclose always follows, and that is where reconnection is scheduled. - console.warn('[UseAI] Connection error:', this.url); + socket.onerror = (event) => { + // Use warn instead of error to avoid triggering Next.js error overlay + console.warn('[UseAI] Connection error:', event.message); }; socket.onclose = () => { - this.detach(socket); - if (this.socket !== socket) return; - this.socket = null; - - const wasConnected = this._connected; + // A close also fires for a failed attempt. Only a socket that opened reports a disconnection. + if (!this._connected) return; this._connected = false; - // A socket that never opened reports only a failed attempt, not a disconnection. - if (wasConnected) { - this.registry.dispatch('disconnect', 'transport close'); - } - this.scheduleReconnect(); + this.registry.dispatch('disconnect', 'transport close'); }; } - private parseFrame(data: unknown): DownstreamFrame | null { - if (typeof data !== 'string') { - console.warn('[UseAI] Ignoring non-text frame'); - return null; - } - let parsed: unknown; - try { - parsed = JSON.parse(data); - } catch { - console.warn('[UseAI] Ignoring malformed frame'); - return null; + disconnect(): void { + const socket = this.socket; + this.socket = null; + this._connected = false; + if (socket) { + socket.onopen = socket.onmessage = socket.onerror = socket.onclose = null; + socket.close(); } - if (typeof parsed !== 'object' || parsed === null) return null; - const frame = parsed as Partial; - if (typeof frame.name !== 'string') return null; - return { name: frame.name, data: frame.data }; } - private detach(socket: WebSocketLike): void { - socket.onopen = null; - socket.onmessage = null; - socket.onerror = null; - socket.onclose = null; + send(message: UseAIClientMessage): void { + this.socket?.send(JSON.stringify(message)); } - private scheduleReconnect(): void { - if (!this.reconnecting || this.reconnectTimer !== null) return; + on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void { + return this.registry.on(name, handler); + } +} - const delay = Math.min( - this.reconnectionDelay * 2 ** this.reconnectAttempts, - this.reconnectionDelayMax, - ); - this.reconnectAttempts++; +interface Frame { + type: string; + name?: string; + value?: unknown; +} - this.reconnectTimer = setTimeout(() => { - this.reconnectTimer = null; - if (this.reconnecting) this.open(); - }, delay); +function parseFrame(data: unknown): Frame | null { + if (typeof data !== 'string') { + console.warn('[UseAI] Ignoring non-text frame'); + return null; + } + let parsed: unknown; + try { + parsed = JSON.parse(data); + } catch { + console.warn('[UseAI] Ignoring malformed frame'); + return null; } + if (typeof parsed !== 'object' || parsed === null) return null; + const frame = parsed as Partial; + if (typeof frame.type !== 'string') return null; + return frame as Frame; } diff --git a/packages/client/src/transport/index.ts b/packages/client/src/transport/index.ts index 7dd56cc1..09d0e193 100644 --- a/packages/client/src/transport/index.ts +++ b/packages/client/src/transport/index.ts @@ -2,4 +2,4 @@ export type { UseAITransport, UseAITransportEventName } from './types'; export { SocketIOTransport } from './SocketIOTransport'; export type { SocketIOTransportOptions } from './SocketIOTransport'; export { WebSocketTransport } from './WebSocketTransport'; -export type { WebSocketTransportOptions, WebSocketLike } from './WebSocketTransport'; +export type { WebSocketTransportOptions } from './WebSocketTransport'; diff --git a/packages/client/src/types.ts b/packages/client/src/types.ts index 35606b18..41c2e044 100644 --- a/packages/client/src/types.ts +++ b/packages/client/src/types.ts @@ -1,13 +1,32 @@ +import type { UseAITransport } from './transport/types'; + /** - * Configuration for the UseAI client provider. + * How the provider reaches the server. Give one of the two. + * + * - `serverUrl` connects over Socket.IO, which the bundled server serves. + * - `transport` connects over anything that implements {@link UseAITransport}. + * + * @example + * ```tsx + * + * + * ``` */ -export interface UseAIConfig { - /** - * The WebSocket URL of the UseAI server. - * Unused when an explicit `transport` is supplied, but still reported on the context. - */ - serverUrl: string; -} +export type UseAIConfig = + | { + /** URL of a Socket.IO UseAI server. */ + serverUrl: string; + transport?: never; + } + | { + /** + * Transport to reach the server with. The provider reads it once, on the first + * render, so an inline object does not reconnect the client on every render. + * Remount the provider to change transports. + */ + transport: UseAITransport; + serverUrl?: never; + }; /** * Toggles for optional chat UI features. Opt-out features default to enabled diff --git a/packages/client/test/disabling-at-runtime.integration.test.tsx b/packages/client/test/disabling-at-runtime.integration.test.tsx index c3513630..720bec7f 100644 --- a/packages/client/test/disabling-at-runtime.integration.test.tsx +++ b/packages/client/test/disabling-at-runtime.integration.test.tsx @@ -23,7 +23,7 @@ describe('useAIContext without UseAIProvider', () => { const { result } = renderHook(() => useAIContext()); expect(result.current).toBeDefined(); - expect(result.current.serverUrl).toBe(''); + expect(result.current.serverUrl).toBeUndefined(); expect(result.current.connected).toBe(false); expect(result.current.client).toBeNull(); expect(result.current.chat.currentId).toBeNull(); diff --git a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts index 8efa6153..421d39c1 100644 --- a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts +++ b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts @@ -1,69 +1,27 @@ -import type { Server as SocketIOServer } from 'socket.io'; import { Server as BunEngine } from '@socket.io/bun-engine'; -import type { RuntimeAdapter, RuntimeServerConfig, RuntimeServerHandle } from '../types'; +import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle } from '../types'; import { resolveCorsHeaders, resolvePreflightHeaders } from './cors'; -import { BunRawWebSocket, isRawWebSocket, type RawWebSocketData } from './rawWebSocket'; +import { BunRawWebSocket, type RawWebSocketData } from './rawWebSocket'; + +type BunServer = Parameters[0]['fetch']>>[1]; +type WebSocketHandler = NonNullable[0]['websocket']>; /** * Runtime adapter for Bun. - * Uses @socket.io/bun-engine for native Bun WebSocket support. + * Serves Socket.IO through @socket.io/bun-engine, or a plain WebSocket through Bun.serve's + * own upgrade, on one HTTP server. */ export class BunRuntimeAdapter implements RuntimeAdapter { readonly name = 'bun' as const; private engine: BunEngine | null = null; - createServer(io: SocketIOServer, config: RuntimeServerConfig): RuntimeServerHandle { - // Create Bun-native engine - this.engine = new BunEngine({ - path: '/socket.io/', - maxHttpBufferSize: config.maxHttpBufferSize, - }); - - // Capture client IP for polling transport at engine connection time - this.engine.on('connection', (engineSocket, req, bunServer) => { - if (engineSocket.transport?.name === 'polling' && config.onPollingConnection) { - const clientIp = bunServer.requestIP(req); - if (clientIp) { - config.onPollingConnection(engineSocket.id, clientIp.address); - } - } - }); + createServer(listener: RuntimeListener, config: RuntimeServerConfig): RuntimeServerHandle { + const { upgrade, websocket } = + listener.transport === 'socketio' + ? this.socketIOHandlers(listener.io, config) + : this.webSocketHandlers(listener.onConnection); - // Bind Socket.IO to Bun engine - io.bind(this.engine); - - const handler = this.engine.handleRequest.bind(this.engine); - const engineWebSocket = this.engine.handler().websocket; - - // Bun.serve takes a single websocket handler table for the whole server, so the - // plain listener and the Socket.IO engine share it and dispatch on ws.data. - const rawSockets = new WeakMap(); - type EngineWebSocketHandler = typeof engineWebSocket; - type EngineWebSocket = Parameters[0]; - const websocketHandler: EngineWebSocketHandler = { - ...engineWebSocket, - open: (ws: EngineWebSocket) => { - if (!isRawWebSocket(ws)) return engineWebSocket.open(ws); - const { remoteAddress } = ws.data as unknown as RawWebSocketData; - const connection = new BunRawWebSocket(ws, remoteAddress); - rawSockets.set(ws, connection); - config.websocket?.onConnection(connection); - }, - message: (ws: EngineWebSocket, message: Parameters[1]) => { - if (!isRawWebSocket(ws)) return engineWebSocket.message(ws, message); - rawSockets.get(ws)?.receiveMessage( - typeof message === 'string' ? message : new TextDecoder().decode(message as Uint8Array), - ); - }, - close: (ws: EngineWebSocket, code: number, reason: string) => { - if (!isRawWebSocket(ws)) return engineWebSocket.close(ws, code, reason); - rawSockets.get(ws)?.receiveClose(); - rawSockets.delete(ws); - }, - }; - - // Start Bun server const bunServer = Bun.serve({ port: config.port, idleTimeout: config.idleTimeout ?? 30, @@ -81,7 +39,6 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }); } - // Helper to create response with CORS headers const corsHeaders = resolveCorsHeaders(requestOrigin, config.cors); // Health check endpoint @@ -91,24 +48,11 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }); } - // Plain WebSocket path - if (config.websocket && url.pathname === config.websocket.path) { - const data: RawWebSocketData = { - useAiRawWebSocket: true, - remoteAddress: server.requestIP(req)?.address, - }; - // The engine owns the server's WebSocket data type; the plain listener - // rides along on the same handler table and is told apart by useAiRawWebSocket. - if (server.upgrade(req, { data: data as unknown as EngineWebSocket['data'] })) return undefined; - return new Response('Expected a WebSocket upgrade', { status: 400, headers: corsHeaders }); - } - - // Socket.IO path - if (url.pathname.startsWith('/socket.io/')) { - const response = await handler(req, server); - - // Add CORS headers to Socket.IO responses - if (response && Object.keys(corsHeaders).length > 0) { + const response = await upgrade(req, server, url); + if (response === null) return undefined; + if (response) { + // Add CORS headers to the listener's responses + if (Object.keys(corsHeaders).length > 0) { const newHeaders = new Headers(response.headers); for (const [key, value] of Object.entries(corsHeaders)) { newHeaders.set(key, value); @@ -124,7 +68,7 @@ export class BunRuntimeAdapter implements RuntimeAdapter { return new Response('Not Found', { status: 404, headers: corsHeaders }); }, - websocket: websocketHandler, + websocket, }); return { @@ -134,4 +78,68 @@ export class BunRuntimeAdapter implements RuntimeAdapter { server: bunServer, }; } + + /** + * @returns `upgrade` yields a Response to send, `null` once the request was upgraded, + * or `undefined` when the path is not the listener's. + */ + private socketIOHandlers(io: RuntimeListener extends infer L ? (L extends { io: infer I } ? I : never) : never, config: RuntimeServerConfig) { + this.engine = new BunEngine({ + path: '/socket.io/', + maxHttpBufferSize: config.maxHttpBufferSize, + }); + + // Capture client IP for polling transport at engine connection time + this.engine.on('connection', (engineSocket, req, bunServer) => { + if (engineSocket.transport?.name === 'polling' && config.onPollingConnection) { + const clientIp = bunServer.requestIP(req); + if (clientIp) { + config.onPollingConnection(engineSocket.id, clientIp.address); + } + } + }); + + io.bind(this.engine); + + const handleRequest = this.engine.handleRequest.bind(this.engine); + return { + upgrade: async (req: Request, server: BunServer, url: URL): Promise => { + if (!url.pathname.startsWith('/socket.io/')) return undefined; + // The engine returns undefined once it has upgraded the request itself. + return (await handleRequest(req, server as never)) ?? null; + }, + websocket: this.engine.handler().websocket as unknown as WebSocketHandler, + }; + } + + private webSocketHandlers(onConnection: (connection: BunRawWebSocket) => void) { + const sockets = new WeakMap(); + const websocket: WebSocketHandler = { + open: (ws) => { + const { remoteAddress } = ws.data as RawWebSocketData; + const connection = new BunRawWebSocket(ws, remoteAddress); + sockets.set(ws, connection); + onConnection(connection); + }, + message: (ws, message) => { + sockets.get(ws)?.receiveMessage( + typeof message === 'string' ? message : new TextDecoder().decode(message), + ); + }, + close: (ws) => { + sockets.get(ws)?.receiveClose(); + sockets.delete(ws); + }, + }; + + return { + upgrade: async (req: Request, server: BunServer, url: URL): Promise => { + if (url.pathname !== '/') return undefined; + const data: RawWebSocketData = { remoteAddress: server.requestIP(req)?.address }; + if (server.upgrade(req, { data })) return null; + return new Response('Expected a WebSocket upgrade', { status: 426 }); + }, + websocket, + }; + } } diff --git a/packages/server/src/runtime/bun/rawWebSocket.ts b/packages/server/src/runtime/bun/rawWebSocket.ts index 79f93011..56743648 100644 --- a/packages/server/src/runtime/bun/rawWebSocket.ts +++ b/packages/server/src/runtime/bun/rawWebSocket.ts @@ -1,14 +1,12 @@ import type { RawWebSocket } from '../types'; -/** Marks a Bun WebSocket as belonging to the plain listener rather than to Socket.IO. */ +/** Per-connection data attached at upgrade time. */ export interface RawWebSocketData { - useAiRawWebSocket: true; remoteAddress?: string; } interface BunWebSocket { readonly readyState: number; - data: unknown; send(data: string): unknown; close(): void; } @@ -53,7 +51,3 @@ export class BunRawWebSocket implements RawWebSocket { this.closeHandler?.(); } } - -export function isRawWebSocket(ws: { data?: unknown }): boolean { - return (ws.data as Partial | undefined)?.useAiRawWebSocket === true; -} diff --git a/packages/server/src/runtime/index.ts b/packages/server/src/runtime/index.ts index 029f0ce5..40537b50 100644 --- a/packages/server/src/runtime/index.ts +++ b/packages/server/src/runtime/index.ts @@ -8,7 +8,7 @@ export type { RuntimeServerConfig, RuntimeServerHandle, RawWebSocket, - RawWebSocketListener, + RuntimeListener, } from './types'; export { detectRuntime } from './detection'; export { createClientIpTracker, type ClientIpTracker, type ClientIpConnection } from './clientIp'; diff --git a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts index 1c55b104..18e770de 100644 --- a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts +++ b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts @@ -1,21 +1,19 @@ import { createServer, type Server as HttpServer } from 'http'; -import type { Server as SocketIOServer } from 'socket.io'; import { WebSocketServer } from 'ws'; -import type { RuntimeAdapter, RuntimeServerConfig, RuntimeServerHandle } from '../types'; +import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle } from '../types'; import { NodeRawWebSocket } from './rawWebSocket'; /** * Runtime adapter for Node.js. - * Uses http.createServer with standard Socket.IO integration. + * Serves Socket.IO through its standard http.Server integration, or a plain WebSocket + * through `ws`, on one HTTP server. * - * CORS handling is delegated entirely to Socket.IO's built-in CORS support. - * This avoids redundancy and ensures consistent behavior with credentials. + * For Socket.IO, CORS handling is delegated entirely to Socket.IO's built-in support. */ export class NodeRuntimeAdapter implements RuntimeAdapter { readonly name = 'node' as const; - createServer(io: SocketIOServer, config: RuntimeServerConfig): RuntimeServerHandle { - // Create Node.js HTTP server + createServer(listener: RuntimeListener, config: RuntimeServerConfig): RuntimeServerHandle { const httpServer: HttpServer = createServer((req, res) => { const url = new URL(req.url || '/', `http://localhost:${config.port}`); @@ -27,36 +25,36 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { return; } - // Socket.IO handles /socket.io/* paths (including CORS) - // Other paths return 404 - if (!url.pathname.startsWith('/socket.io/')) { - res.statusCode = 404; - res.end('Not Found'); - } - // Socket.IO will handle the request via its internal listeners + // Socket.IO answers /socket.io/* itself through its own request listener. + if (listener.transport === 'socketio' && url.pathname.startsWith('/socket.io/')) return; + + res.statusCode = 404; + res.end('Not Found'); }); - // The plain listener claims its path before Socket.IO attaches, so engine.io's - // own upgrade handler sees a handshake already written and leaves the socket alone. - const websocketConfig = config.websocket; let wss: WebSocketServer | null = null; - if (websocketConfig) { - wss = new WebSocketServer({ noServer: true, maxPayload: config.maxHttpBufferSize }); - httpServer.on('upgrade', (req, socket, head) => { - const url = new URL(req.url || '/', `http://localhost:${config.port}`); - if (url.pathname !== websocketConfig.path) return; - wss!.handleUpgrade(req, socket, head, (ws) => { - const forwardedFor = req.headers['x-forwarded-for']; - const remoteAddress = typeof forwardedFor === 'string' - ? forwardedFor.split(',')[0].trim() - : req.socket.remoteAddress; - websocketConfig.onConnection(new NodeRawWebSocket(ws, remoteAddress)); - }); - }); + if (listener.transport === 'socketio') { + this.attachSocketIO(listener.io, httpServer, config); + } else { + wss = this.attachWebSocket(listener.onConnection, httpServer, config); } - // Attach Socket.IO to the HTTP server - // Socket.IO handles CORS internally for /socket.io/* paths + httpServer.listen(config.port); + + return { + stop: () => { + wss?.close(); + httpServer.close(); + }, + server: httpServer, + }; + } + + private attachSocketIO( + io: Extract['io'], + httpServer: HttpServer, + config: RuntimeServerConfig, + ) { io.attach(httpServer, { transports: ['polling', 'websocket'], maxHttpBufferSize: config.maxHttpBufferSize, @@ -81,16 +79,28 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { } }); } + } - // Start listening - httpServer.listen(config.port); - - return { - stop: () => { - wss?.close(); - httpServer.close(); - }, - server: httpServer, - }; + private attachWebSocket( + onConnection: Extract['onConnection'], + httpServer: HttpServer, + config: RuntimeServerConfig, + ): WebSocketServer { + const wss = new WebSocketServer({ noServer: true, maxPayload: config.maxHttpBufferSize }); + httpServer.on('upgrade', (req, socket, head) => { + const url = new URL(req.url || '/', `http://localhost:${config.port}`); + if (url.pathname !== '/') { + socket.destroy(); + return; + } + wss.handleUpgrade(req, socket, head, (ws) => { + const forwardedFor = req.headers['x-forwarded-for']; + const remoteAddress = typeof forwardedFor === 'string' + ? forwardedFor.split(',')[0].trim() + : req.socket.remoteAddress; + onConnection(new NodeRawWebSocket(ws, remoteAddress)); + }); + }); + return wss; } } diff --git a/packages/server/src/runtime/types.ts b/packages/server/src/runtime/types.ts index eeb3baa5..6f46f213 100644 --- a/packages/server/src/runtime/types.ts +++ b/packages/server/src/runtime/types.ts @@ -26,17 +26,11 @@ export interface RawWebSocket { } /** - * A plain WebSocket listener, served on the same port and HTTP server as Socket.IO. + * What the HTTP server hands connections to. A server runs exactly one. */ -export interface RawWebSocketListener { - /** - * Path that upgrades to a plain WebSocket. - * @example '/ws' - */ - path: string; - /** Called once per accepted connection. */ - onConnection(connection: RawWebSocket): void; -} +export type RuntimeListener = + | { transport: 'socketio'; io: SocketIOServer } + | { transport: 'websocket'; onConnection(connection: RawWebSocket): void }; /** * Configuration for creating a runtime server. @@ -59,11 +53,6 @@ export interface RuntimeServerConfig { * Called when a polling transport connection is established. */ onPollingConnection?: (sessionId: string, ip: string) => void; - /** - * Plain WebSocket listener to serve alongside Socket.IO. - * Omit to serve Socket.IO only. - */ - websocket?: RawWebSocketListener; } /** @@ -85,11 +74,11 @@ export interface RuntimeAdapter { readonly name: RuntimeType; /** - * Creates an HTTP server and binds Socket.IO to it. + * Creates an HTTP server and binds the listener to it. * - * @param io - Socket.IO server instance + * @param listener - Socket.IO server, or a plain WebSocket connection handler * @param config - Server configuration * @returns Handle to the running server */ - createServer(io: SocketIOServer, config: RuntimeServerConfig): RuntimeServerHandle; + createServer(listener: RuntimeListener, config: RuntimeServerConfig): RuntimeServerHandle; } diff --git a/packages/server/src/server.ts b/packages/server/src/server.ts index 0b2212a8..f4622627 100644 --- a/packages/server/src/server.ts +++ b/packages/server/src/server.ts @@ -37,6 +37,7 @@ import { type RuntimeServerHandle, type ClientIpTracker, type RawWebSocket, + type RuntimeListener, } from './runtime'; import { WebSocketClientConnection } from './webSocketConnection'; @@ -91,18 +92,18 @@ export type { ClientSession, ClientConnection } from './agents/types'; * ``` */ export class UseAIServer { - private io: SocketIOServer; + private io: SocketIOServer | null = null; private runtimeAdapter: RuntimeAdapter; private serverHandle: RuntimeServerHandle | null = null; private agent: Agent; // Default agent for chat (run_agent) private defaultAgentId: string; // ID of the default agent private agents: Record; // Registry of all agents private clients: Map = new Map(); - private config: Required> & { + private config: Required> & { maxHttpBufferSize: number; cors?: CorsOptions; idleTimeout: number; - webSocketPath: string | null; + transport: 'socketio' | 'websocket'; }; private rateLimiter: RateLimiter; private cleanupInterval: NodeJS.Timeout; @@ -134,7 +135,7 @@ export class UseAIServer { maxHttpBufferSize: config.maxHttpBufferSize ?? 20 * 1024 * 1024, // 20MB default cors: config.cors, idleTimeout: config.idleTimeout ?? 30, - webSocketPath: config.webSocketPath === undefined ? '/ws' : config.webSocketPath, + transport: config.transport ?? 'socketio', }; // Set agents registry @@ -198,13 +199,13 @@ export class UseAIServer { this.runtimeAdapter = createRuntimeAdapter(config.runtime ?? 'auto'); logger.info('Using runtime adapter', { runtime: this.runtimeAdapter.name }); - // Create Socket.IO server - this.io = new SocketIOServer({ - transports: ['polling', 'websocket'], - maxHttpBufferSize: this.config.maxHttpBufferSize, - }); - - this.setupSocketIOServer(); + if (this.config.transport === 'socketio') { + this.io = new SocketIOServer({ + transports: ['polling', 'websocket'], + maxHttpBufferSize: this.config.maxHttpBufferSize, + }); + this.setupSocketIOServer(this.io); + } if (this.rateLimiter.isEnabled()) { logger.info('Rate limiting enabled', { @@ -213,8 +214,11 @@ export class UseAIServer { }); } - // Start server using runtime adapter - this.serverHandle = this.runtimeAdapter.createServer(this.io, { + const listener: RuntimeListener = this.io + ? { transport: 'socketio', io: this.io } + : { transport: 'websocket', onConnection: (socket) => this.handleWebSocketConnection(socket) }; + + this.serverHandle = this.runtimeAdapter.createServer(listener, { port: this.config.port, idleTimeout: this.config.idleTimeout, cors: this.config.cors, @@ -222,13 +226,8 @@ export class UseAIServer { onPollingConnection: (sessionId, ip) => { this.clientIpTracker.trackPollingConnection(sessionId, ip); }, - websocket: this.config.webSocketPath - ? { - path: this.config.webSocketPath, - onConnection: (socket) => this.handleWebSocketConnection(socket), - } - : undefined, }); + logger.info('UseAI server ready', { port: this.config.port, transport: this.config.transport }); } /** @@ -293,8 +292,8 @@ export class UseAIServer { logger.debug('Registered message handler', { type }); } - private setupSocketIOServer() { - this.io.on('connection', (socket: Socket) => { + private setupSocketIOServer(io: SocketIOServer) { + io.on('connection', (socket: Socket) => { // Get connection info for IP address resolution const conn = socket.conn as unknown as { id: string; transport: { name: string; socket?: { remoteAddress?: string } } }; // Get IP address for rate limiting: @@ -327,14 +326,11 @@ export class UseAIServer { this.destroySession(session); }); }); - - logger.info('UseAI server ready', { port: this.config.port }); } /** - * Accepts a plain WebSocket connection, the alternative to Socket.IO served on the - * same port. Frames are JSON text: upstream the client message on its own, downstream - * a `{ name, data }` envelope. See docs/websocket-protocol.md. + * Accepts a plain WebSocket connection. Frames are JSON text: upstream the client + * message on its own, downstream one AG-UI event. See docs/websocket-protocol.md. */ private handleWebSocketConnection(socket: RawWebSocket) { const connection = new WebSocketClientConnection(`ws-${uuidv4()}`, socket); @@ -1170,7 +1166,7 @@ export class UseAIServer { this.plugins.map(plugin => plugin.close?.()) ); - this.io.close(); + this.io?.close(); if (this.serverHandle) { this.serverHandle.stop(); this.serverHandle = null; diff --git a/packages/server/src/types.ts b/packages/server/src/types.ts index e3d2380c..55765bce 100644 --- a/packages/server/src/types.ts +++ b/packages/server/src/types.ts @@ -135,15 +135,13 @@ export interface UseAIServerConfig { - const client = new UseAIClient(new WebSocketTransport(`ws://localhost:${port}${path}`)); +function connectClient(port: number): Promise<{ client: UseAIClient; events: AGUIEvent[] }> { + const client = new UseAIClient(new WebSocketTransport(`ws://localhost:${port}`)); const events: AGUIEvent[] = []; client.onEvent('test', event => events.push(event)); @@ -58,6 +58,7 @@ describe.each(RUNTIMES)('WebSocketTransport against a real server: %s runtime', server = new UseAIServer({ port, runtime, + transport: 'websocket', agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, defaultAgent: 'test-agent', }); @@ -125,10 +126,9 @@ describe.each(RUNTIMES)('WebSocketTransport against a real server: %s runtime', client.disconnect(); }); - test('Socket.IO still serves the same port', async () => { - const socket = await cleanup.createTestClient(port); - expect(socket.connected).toBe(true); - socket.disconnect(); + test('the health endpoint still answers', async () => { + const response = await fetch(`http://localhost:${port}/health`); + expect(response.ok).toBe(true); }); test('two plain WebSocket clients get isolated sessions', async () => { @@ -150,45 +150,24 @@ describe.each(RUNTIMES)('WebSocketTransport against a real server: %s runtime', }); }); -describe('webSocketPath', () => { +describe('transport defaults to socketio', () => { const cleanup = new TestCleanupManager(); afterAll(() => { cleanup.cleanup(); }); - test('serves the plain listener at a custom path', async () => { - const port = 9550; - const model = createSequentialMockModel([{ text: 'hi' }]); - cleanup.trackServer( - new UseAIServer({ - port, - webSocketPath: '/agent', - agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, - defaultAgent: 'test-agent', - }), - ); - - const { client } = await connectClient(port, '/agent'); - await waitFor(() => client.availableAgents.length > 0, 'the agents payload'); - - expect(client.defaultAgent).toBe('test-agent'); - client.disconnect(); - }); - - test('null serves Socket.IO only', async () => { + test('a default server serves Socket.IO and refuses a plain WebSocket', async () => { const port = 9560; const model = createSequentialMockModel([{ text: 'hi' }]); cleanup.trackServer( new UseAIServer({ port, - webSocketPath: null, agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, defaultAgent: 'test-agent', }), ); - // Socket.IO is unaffected. const socket = await cleanup.createTestClient(port); expect(socket.connected).toBe(true); socket.disconnect(); From 116f6cc8bb2095bdce66c81c78dda3b47c59df15 Mon Sep 17 00:00:00 2001 From: Zachary Davison Date: Fri, 4 Sep 2026 12:04:11 +0200 Subject: [PATCH 3/4] Report the transport's url as serverUrl on the provider context UseAITransport gains a readonly url. Both bundled transports already take one. UseAIContextValue.serverUrl is a required string again, populated from the transport when the provider was given one. Co-Authored-By: Claude Opus 5 (1M context) --- packages/client/src/client.ts | 5 +++++ packages/client/src/providers/useAIProvider.tsx | 7 ++++--- packages/client/src/transport/SocketIOTransport.ts | 6 +++--- packages/client/src/transport/WebSocketTransport.ts | 2 +- packages/client/src/transport/types.ts | 6 ++++++ packages/client/src/types.ts | 3 ++- .../client/test/disabling-at-runtime.integration.test.tsx | 2 +- 7 files changed, 22 insertions(+), 9 deletions(-) diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index 14def423..afb32590 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -119,6 +119,11 @@ export class UseAIClient { this.transport = typeof target === 'string' ? new SocketIOTransport(target) : target; } + /** URL of the server the transport connects to. */ + get serverUrl(): string { + return this.transport.url; + } + /** * Opens the transport's connection to the server. * Connection state changes are notified via onConnectionStateChange(). diff --git a/packages/client/src/providers/useAIProvider.tsx b/packages/client/src/providers/useAIProvider.tsx index e1d430a6..362eecc8 100644 --- a/packages/client/src/providers/useAIProvider.tsx +++ b/packages/client/src/providers/useAIProvider.tsx @@ -122,8 +122,8 @@ export interface PromptsContextValue { * Contains connection state and methods for managing tools and prompts. */ export interface UseAIContextValue { - /** URL of the server, when the provider was given `serverUrl` rather than `transport`. */ - serverUrl?: string; + /** URL of the server. With an explicit `transport`, this is the transport's `url`. */ + serverUrl: string; /** Whether the client is connected to the server */ connected: boolean; /** The underlying WebSocket client instance */ @@ -162,6 +162,7 @@ let hasWarnedAboutMissingProvider = false; * Allows hooks to gracefully degrade instead of crashing. */ const noOpContextValue: UseAIContextValue = { + serverUrl: '', connected: false, client: null, tools: { @@ -742,7 +743,7 @@ export function UseAIProvider({ // ── Context Values ────────────────────────────────────────────────────── const value: UseAIContextValue = { - serverUrl, + serverUrl: transportRef.current?.url ?? serverUrl!, connected, client: clientRef.current, tools: { diff --git a/packages/client/src/transport/SocketIOTransport.ts b/packages/client/src/transport/SocketIOTransport.ts index 99f0650e..3e2f5272 100644 --- a/packages/client/src/transport/SocketIOTransport.ts +++ b/packages/client/src/transport/SocketIOTransport.ts @@ -41,13 +41,13 @@ export class SocketIOTransport implements UseAITransport { private reconnectionDelayMax: number; /** - * @param serverUrl - The URL of the UseAI server + * @param url - The URL of the UseAI server * @example * ```typescript * new SocketIOTransport('ws://localhost:8081'); * ``` */ - constructor(private serverUrl: string, options: SocketIOTransportOptions = {}) { + constructor(readonly url: string, options: SocketIOTransportOptions = {}) { this.reconnectionDelay = options.reconnectionDelay ?? 1000; this.reconnectionDelayMax = options.reconnectionDelayMax ?? 10_000; } @@ -57,7 +57,7 @@ export class SocketIOTransport implements UseAITransport { } connect(): void { - const socket = io(this.serverUrl, { + const socket = io(this.url, { transports: ['polling', 'websocket'], reconnection: true, reconnectionAttempts: Infinity, diff --git a/packages/client/src/transport/WebSocketTransport.ts b/packages/client/src/transport/WebSocketTransport.ts index 3b0180b7..7bf45d36 100644 --- a/packages/client/src/transport/WebSocketTransport.ts +++ b/packages/client/src/transport/WebSocketTransport.ts @@ -60,7 +60,7 @@ export class WebSocketTransport implements UseAITransport { * new WebSocketTransport('wss://your-server.com'); * ``` */ - constructor(private url: string, options: WebSocketTransportOptions = {}) { + constructor(readonly url: string, options: WebSocketTransportOptions = {}) { this.options = { reconnectionDelay: options.reconnectionDelay ?? 1000, reconnectionDelayMax: options.reconnectionDelayMax ?? 10_000, diff --git a/packages/client/src/transport/types.ts b/packages/client/src/transport/types.ts index bf85c79c..d3c79d77 100644 --- a/packages/client/src/transport/types.ts +++ b/packages/client/src/transport/types.ts @@ -21,6 +21,12 @@ export type UseAITransportEventName = 'connect' | 'disconnect' | 'event' | 'agen * and {@link WebSocketTransport}. */ export interface UseAITransport { + /** + * The server this transport connects to. Reported as `serverUrl` on the provider context. + * @example 'wss://your-server.com' + */ + readonly url: string; + /** Opens the connection. Reconnection until {@link disconnect} is the transport's own responsibility. */ connect(): void; diff --git a/packages/client/src/types.ts b/packages/client/src/types.ts index 41c2e044..8f74501a 100644 --- a/packages/client/src/types.ts +++ b/packages/client/src/types.ts @@ -22,7 +22,8 @@ export type UseAIConfig = /** * Transport to reach the server with. The provider reads it once, on the first * render, so an inline object does not reconnect the client on every render. - * Remount the provider to change transports. + * Remount the provider to change transports. The context reports the + * transport's `url` as `serverUrl`. */ transport: UseAITransport; serverUrl?: never; diff --git a/packages/client/test/disabling-at-runtime.integration.test.tsx b/packages/client/test/disabling-at-runtime.integration.test.tsx index 720bec7f..c3513630 100644 --- a/packages/client/test/disabling-at-runtime.integration.test.tsx +++ b/packages/client/test/disabling-at-runtime.integration.test.tsx @@ -23,7 +23,7 @@ describe('useAIContext without UseAIProvider', () => { const { result } = renderHook(() => useAIContext()); expect(result.current).toBeDefined(); - expect(result.current.serverUrl).toBeUndefined(); + expect(result.current.serverUrl).toBe(''); expect(result.current.connected).toBe(false); expect(result.current.client).toBeNull(); expect(result.current.chat.currentId).toBeNull(); From 6970d4319b1c86cb7e9b1c0e2ff568d3726a18d6 Mon Sep 17 00:00:00 2001 From: Zachary Davison Date: Fri, 4 Sep 2026 12:23:18 +0200 Subject: [PATCH 4/4] Simplify the transport seam and the server connection boundary Client: UseAITransport is now onEvent(AGUIEvent) + onConnectionChange. The five named channels, the handler registry and the agents/config special case in WebSocketTransport are gone. SocketIOTransport presents its legacy agents/config events as AG-UI CUSTOM events, and UseAIClient.handleEvent gains one CUSTOM branch. SocketIOTransport loses options nothing used. The provider resolves its target once. Server: ClientConnection is the whole per-connection boundary (id, ipAddress, connected, emit, onMessage, onClose). SocketIOClientConnection and WebSocketClientConnection own their protocol details; server.ts has one acceptConnection. handleClientMessage takes the session it is given. The agents payload is built once. Bun adapter returns {path, upgrade, websocket} per listener and keeps the connection on ws.data. Binary frames are ignored on both runtimes. Listener member types are named. X-Forwarded-For parsing is one helper. Tests share one FakeWebSocket, one socket.io mock and one waitUntil. The negative transport test checks a raw socket close instead of waiting out a 5s timeout. Co-Authored-By: Claude Opus 5 (1M context) --- README.md | 2 +- docs/websocket-protocol.md | 10 +- packages/client/src/client.test.ts | 113 +++---------- packages/client/src/client.ts | 65 ++++---- packages/client/src/index.ts | 7 +- .../client/src/providers/useAIProvider.tsx | 13 +- .../src/transport/SocketIOTransport.test.ts | 135 +++++----------- .../client/src/transport/SocketIOTransport.ts | 91 ++++++----- .../src/transport/WebSocketTransport.test.ts | 152 ++++-------------- .../src/transport/WebSocketTransport.ts | 67 ++++---- .../client/src/transport/handlerRegistry.ts | 32 ---- packages/client/src/transport/index.ts | 3 +- packages/client/src/transport/types.ts | 39 +++-- packages/client/test/fakeWebSocket.ts | 56 +++++++ packages/client/test/socketIOMock.ts | 60 +++++++ packages/server/src/agents/types.ts | 16 +- .../src/runtime/bun/BunRuntimeAdapter.ts | 110 ++++++------- .../server/src/runtime/bun/rawWebSocket.ts | 5 +- packages/server/src/runtime/clientIp.ts | 10 ++ packages/server/src/runtime/index.ts | 4 +- .../src/runtime/node/NodeRuntimeAdapter.ts | 32 ++-- .../server/src/runtime/node/rawWebSocket.ts | 2 +- packages/server/src/runtime/types.ts | 20 ++- packages/server/src/server.ts | 118 ++++---------- packages/server/src/socketIOConnection.ts | 47 ++++++ packages/server/src/webSocketConnection.ts | 42 ++++- .../websocket-transport.integration.test.ts | 20 +-- packages/server/test/test-utils.ts | 11 ++ 28 files changed, 587 insertions(+), 695 deletions(-) delete mode 100644 packages/client/src/transport/handlerRegistry.ts create mode 100644 packages/client/test/fakeWebSocket.ts create mode 100644 packages/client/test/socketIOMock.ts create mode 100644 packages/server/src/socketIOConnection.ts diff --git a/README.md b/README.md index ab2e35d7..0154faf2 100644 --- a/README.md +++ b/README.md @@ -323,7 +323,7 @@ root.render( The bundled server serves one transport. Set `transport: 'websocket'` on `UseAIServer`, or `TRANSPORT=websocket` on the Docker image, to serve a plain WebSocket at `/` instead of Socket.IO. -To carry the same messages over something else, implement `UseAITransport` yourself. The interface has five members: `connect`, `disconnect`, `send`, `on` and `connected`. +To carry the same messages over something else, implement `UseAITransport` yourself. It opens and closes a connection, sends client messages, and delivers AG-UI events. See [docs/websocket-protocol.md](docs/websocket-protocol.md) for the frames, the turn sequence, and the reconnection behaviour. diff --git a/docs/websocket-protocol.md b/docs/websocket-protocol.md index 0050db9e..f8d11fde 100644 --- a/docs/websocket-protocol.md +++ b/docs/websocket-protocol.md @@ -121,15 +121,17 @@ The server destroys the session when the connection closes. A reconnected client ## Your own transport -`UseAITransport` has five members: +`UseAITransport` has seven members: +- `url` +- `connected` - `connect` - `disconnect` - `send` -- `on` -- `connected` +- `onEvent`, which delivers AG-UI events +- `onConnectionChange` -Implement `UseAITransport` to carry the same messages over something else. +Implement `UseAITransport` to carry the same messages over something else. Deliver the agent list and the server config as `CUSTOM` events, as the section above describes. ```typescript import { UseAIClient, type UseAITransport } from '@meetsmore-oss/use-ai-client'; diff --git a/packages/client/src/client.test.ts b/packages/client/src/client.test.ts index 34727abe..e1550eb5 100644 --- a/packages/client/src/client.test.ts +++ b/packages/client/src/client.test.ts @@ -1,6 +1,7 @@ -import { describe, test, expect, mock, beforeEach, afterEach, spyOn } from 'bun:test'; -import type { Socket } from 'socket.io-client'; +import { describe, test, expect, beforeEach, afterEach, spyOn } from 'bun:test'; import type { UseAIClientMessage } from './types'; +import { installSocketIOMock } from '../test/socketIOMock'; +import { FakeWebSocket, FakeWebSocketConstructor, waitUntil } from '../test/fakeWebSocket'; /** * The UseAIClient suite, run over every bundled transport. @@ -10,72 +11,7 @@ import type { UseAIClientMessage } from './types'; * next to it, in transport/*.test.ts. */ -// ── Socket.IO harness ─────────────────────────────────────────────────────── - -let socketHandlers: Record = {}; -let mockSocket: Partial & { connected: boolean }; - -function createMockSocket() { - socketHandlers = {}; - mockSocket = { - on: mock((event: string, handler: Function) => { - (socketHandlers[event] ??= []).push(handler); - return mockSocket as Socket; - }), - emit: mock(() => mockSocket as Socket), - connected: false, - disconnect: mock(() => mockSocket as Socket), - io: { - engine: { - transport: { name: 'polling' }, - on: mock(), - }, - } as never, - }; - return mockSocket as Socket; -} - -mock.module('socket.io-client', () => ({ - io: () => createMockSocket(), -})); - -// ── Plain WebSocket harness ───────────────────────────────────────────────── - -/** Enough of a WHATWG WebSocket for partysocket to drive. */ -class FakeWebSocket extends EventTarget { - static latest: FakeWebSocket | null = null; - - readyState = 0; - binaryType = 'blob'; - sent: string[] = []; - - constructor() { - super(); - FakeWebSocket.latest = this; - } - - send(data: string): void { - this.sent.push(data); - } - - close(): void { - this.readyState = 3; - } - - serverOpen(): void { - this.readyState = 1; - this.dispatchEvent(new Event('open')); - } - - serverSend(data: string): void { - this.dispatchEvent(new MessageEvent('message', { data })); - } - - serverClose(): void { - this.readyState = 3; - this.dispatchEvent(new CloseEvent('close', { code: 1006 })); - } -} +const sio = installSocketIOMock(); // Imported after the module mock so SocketIOTransport picks it up. const { UseAIClient } = await import('./client'); @@ -101,36 +37,26 @@ interface Harness { sent(): UseAIClientMessage[]; } -async function waitUntil(condition: () => boolean, timeoutMs = 500): Promise { - const deadline = Date.now() + timeoutMs; - while (!condition()) { - if (Date.now() > deadline) throw new Error('Timed out waiting for the transport to reconnect'); - await new Promise(resolve => setTimeout(resolve, 1)); - } -} - const HARNESSES: Array<[string, () => Harness]> = [ [ 'SocketIOTransport', () => { const client = new UseAIClient(new SocketIOTransport('http://localhost:8081')); client.connect(); - const socket = mockSocket; - const fire = (event: string, ...args: unknown[]) => - socketHandlers[event]?.forEach(handler => handler(...args)); + const socket = sio.socket; return { client, async open() { socket.connected = true; - fire('connect'); + sio.fire('connect'); }, close(reason: string) { socket.connected = false; - fire('disconnect', reason); + sio.fire('disconnect', reason); }, deliver(name, data) { - fire(name, data); + sio.fire(name, data); }, sent() { const emit = socket.emit as unknown as { mock: { calls: unknown[][] } }; @@ -144,18 +70,18 @@ const HARNESSES: Array<[string, () => Harness]> = [ [ 'WebSocketTransport', () => { - FakeWebSocket.latest = null; + FakeWebSocket.reset(); const client = new UseAIClient( new WebSocketTransport('wss://localhost:8081', { reconnectionDelay: 1, reconnectionDelayMax: 1, - WebSocket: FakeWebSocket as unknown as typeof WebSocket, + WebSocket: FakeWebSocketConstructor, }), ); client.connect(); // partysocket opens the socket asynchronously, and opens a fresh one after a drop. const liveSocket = async () => { - await waitUntil(() => FakeWebSocket.latest !== null && FakeWebSocket.latest.readyState < 2); + await waitUntil(() => FakeWebSocket.latest !== undefined && FakeWebSocket.latest.readyState < 2); return FakeWebSocket.latest!; }; const sent: string[] = []; @@ -192,7 +118,10 @@ describe.each(HARNESSES)('UseAIClient over %s', (_name, createHarness) => { let client: Client; /** The last message the client sent upstream. */ - const lastSent = () => harness.sent()[harness.sent().length - 1]; + const lastSent = () => { + const all = harness.sent(); + return all[all.length - 1]; + }; const emitEvent = (event: Record) => harness.deliver('event', event); beforeEach(() => { @@ -783,26 +712,24 @@ describe('UseAIClient construction', () => { const client = new UseAIClient('http://localhost:8081'); client.connect(); - expect(mockSocket).toBeDefined(); + expect(sio.socket).toBeDefined(); expect(client.isConnected()).toBe(false); - mockSocket.connected = true; + sio.socket.connected = true; expect(client.isConnected()).toBe(true); client.disconnect(); }); test('disconnect() unsubscribes from the transport', async () => { - FakeWebSocket.latest = null; + FakeWebSocket.reset(); const client = new UseAIClient( - new WebSocketTransport('wss://localhost:8081', { - WebSocket: FakeWebSocket as unknown as typeof WebSocket, - }), + new WebSocketTransport('wss://localhost:8081', { WebSocket: FakeWebSocketConstructor }), ); const stateChanges: boolean[] = []; client.onConnectionStateChange(connected => stateChanges.push(connected)); client.connect(); - await waitUntil(() => FakeWebSocket.latest !== null); + await waitUntil(() => FakeWebSocket.latest !== undefined); const socket = FakeWebSocket.latest!; socket.serverOpen(); diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index afb32590..4bea0449 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -1,4 +1,4 @@ -import { EventType } from '@meetsmore-oss/use-ai-core'; +import { EventType, type CustomEvent } from '@meetsmore-oss/use-ai-core'; import type { ToolDefinition, Message, @@ -119,11 +119,6 @@ export class UseAIClient { this.transport = typeof target === 'string' ? new SocketIOTransport(target) : target; } - /** URL of the server the transport connects to. */ - get serverUrl(): string { - return this.transport.url; - } - /** * Opens the transport's connection to the server. * Connection state changes are notified via onConnectionStateChange(). @@ -131,47 +126,49 @@ export class UseAIClient { */ connect(): void { this.transportUnsubscribes.push( - this.transport.on('connect', () => { - console.log('[UseAI] Connected to server'); - this.connectionStateHandlers.forEach(handler => handler(true)); + this.transport.onConnectionChange((connected, reason) => { + console.log(connected ? '[UseAI] Connected to server' : '[UseAI] Disconnected:', reason ?? ''); + this.connectionStateHandlers.forEach(handler => handler(connected)); }), - - this.transport.on('event', (data) => { - const aguiEvent = data as AGUIEvent; + this.transport.onEvent((event) => { try { - console.log('[Client] Received event:', aguiEvent.type); - this.handleEvent(aguiEvent); + console.log('[Client] Received event:', event.type); + this.handleEvent(event); } catch (error) { console.error('[UseAI] Error handling event:', error); } }), - - this.transport.on('agents', (data) => { - const { agents, defaultAgent } = data as { agents: AgentInfo[]; defaultAgent: string }; - console.log('[Client] Received available agents:', data); - this._availableAgents = agents; - this._defaultAgent = defaultAgent; - this.agentsChangeHandlers.forEach(handler => handler(agents, defaultAgent)); - }), - - this.transport.on('config', (data) => { - const { langfuseEnabled } = data as { langfuseEnabled?: boolean }; - console.log('[Client] Received server config:', data); - this._langfuseEnabled = langfuseEnabled ?? false; - this.langfuseConfigHandlers.forEach(handler => handler(this._langfuseEnabled)); - }), - - this.transport.on('disconnect', (reason) => { - console.log('[UseAI] Disconnected:', reason); - this.connectionStateHandlers.forEach(handler => handler(false)); - }), ); this.transport.connect(); } + /** + * The agent list and the server config arrive as CUSTOM events. They update the + * client and are not forwarded to onEvent subscribers. + */ + private handleCustomEvent(event: CustomEvent): boolean { + if (event.name === 'agents') { + const { agents, defaultAgent } = event.value as { agents: AgentInfo[]; defaultAgent: string }; + console.log('[Client] Received available agents:', event.value); + this._availableAgents = agents; + this._defaultAgent = defaultAgent; + this.agentsChangeHandlers.forEach(handler => handler(agents, defaultAgent)); + return true; + } + if (event.name === 'config') { + const { langfuseEnabled } = event.value as { langfuseEnabled?: boolean }; + console.log('[Client] Received server config:', event.value); + this._langfuseEnabled = langfuseEnabled ?? false; + this.langfuseConfigHandlers.forEach(handler => handler(this._langfuseEnabled)); + return true; + } + return false; + } private handleEvent(event: AGUIEvent) { + if (event.type === EventType.CUSTOM && this.handleCustomEvent(event as CustomEvent)) return; + // Track assistant message lifecycle for conversation history if (event.type === EventType.RUN_STARTED) { // Start of a new assistant response - initialize message diff --git a/packages/client/src/index.ts b/packages/client/src/index.ts index 51932c7f..c4cd079d 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -3,12 +3,7 @@ export { useAIWorkflow } from './useAIWorkflow'; export { UseAIProvider, useAIContext } from './providers/useAIProvider'; export { UseAIClient } from './client'; export { SocketIOTransport, WebSocketTransport } from './transport'; -export type { - UseAITransport, - UseAITransportEventName, - SocketIOTransportOptions, - WebSocketTransportOptions, -} from './transport'; +export type { UseAITransport, WebSocketTransportOptions } from './transport'; export { defineTool, executeDefinedTool, convertToolsToDefinitions } from './defineTool'; /** @hidden */ export { z } from 'zod'; diff --git a/packages/client/src/providers/useAIProvider.tsx b/packages/client/src/providers/useAIProvider.tsx index 362eecc8..1f9f77d2 100644 --- a/packages/client/src/providers/useAIProvider.tsx +++ b/packages/client/src/providers/useAIProvider.tsx @@ -576,14 +576,13 @@ export function UseAIProvider({ const handleDisconnectRef = useRef(serverEvents.handleDisconnect); handleDisconnectRef.current = serverEvents.handleDisconnect; - // Read once: an inline `transport` object would otherwise re-create the client - // on every render. - const transportRef = useRef(transport); + const resolvedServerUrl = transport?.url ?? serverUrl!; useEffect(() => { - const target = transportRef.current ?? serverUrl!; - console.log('[UseAIProvider] Initializing client with', typeof target === 'string' ? target : 'transport'); - const client = new UseAIClient(target); + console.log('[UseAIProvider] Initializing client for', resolvedServerUrl); + // `transport` is deliberately not a dependency: an inline object must not + // re-create the client on every render. + const client = new UseAIClient(transport ?? serverUrl!); const unsubscribeConnection = client.onConnectionStateChange((isConnected) => { console.log('[UseAIProvider] Connection state changed:', isConnected); @@ -743,7 +742,7 @@ export function UseAIProvider({ // ── Context Values ────────────────────────────────────────────────────── const value: UseAIContextValue = { - serverUrl: transportRef.current?.url ?? serverUrl!, + serverUrl: resolvedServerUrl, connected, client: clientRef.current, tools: { diff --git a/packages/client/src/transport/SocketIOTransport.test.ts b/packages/client/src/transport/SocketIOTransport.test.ts index 5d75bc4d..3ac94a73 100644 --- a/packages/client/src/transport/SocketIOTransport.test.ts +++ b/packages/client/src/transport/SocketIOTransport.test.ts @@ -1,46 +1,10 @@ -import { describe, test, expect, mock, beforeEach, afterEach, spyOn } from 'bun:test'; -import type { Socket } from 'socket.io-client'; +import { describe, test, expect, beforeEach, afterEach, spyOn } from 'bun:test'; +import { installSocketIOMock } from '../../test/socketIOMock'; // Socket.IO builds its own socket, so the module is the only seam for testing // this transport's wiring. Everything above it is tested through a transport // instead: see client.test.ts. -let handlers: Record = {}; -let ioOptions: Record | undefined; -let mockSocket: Partial & { connected: boolean }; - -function createMockSocket() { - handlers = {}; - mockSocket = { - on: mock((event: string, handler: Function) => { - (handlers[event] ??= []).push(handler); - return mockSocket as Socket; - }), - emit: mock(() => mockSocket as Socket), - connected: false, - disconnect: mock(() => mockSocket as Socket), - io: { - engine: { - transport: { name: 'polling' }, - on: mock((event: string, handler: Function) => { - (handlers[`engine:${event}`] ??= []).push(handler); - }), - }, - } as never, - }; - return mockSocket as Socket; -} - -function emitSocketEvent(event: string, ...args: unknown[]) { - handlers[event]?.forEach(handler => handler(...args)); -} - -mock.module('socket.io-client', () => ({ - io: (_url: string, options: Record) => { - ioOptions = options; - return createMockSocket(); - }, -})); - +const sio = installSocketIOMock(); const { SocketIOTransport } = await import('./SocketIOTransport'); describe('SocketIOTransport', () => { @@ -48,7 +12,6 @@ describe('SocketIOTransport', () => { let consoleWarnSpy: ReturnType; beforeEach(() => { - ioOptions = undefined; consoleLogSpy = spyOn(console, 'log').mockImplementation(() => {}); consoleWarnSpy = spyOn(console, 'warn').mockImplementation(() => {}); }); @@ -61,7 +24,7 @@ describe('SocketIOTransport', () => { test('reconnects indefinitely with backoff capped at ten seconds', () => { new SocketIOTransport('http://localhost:8081').connect(); - expect(ioOptions).toMatchObject({ + expect(sio.ioOptions).toMatchObject({ transports: ['polling', 'websocket'], reconnection: true, reconnectionAttempts: Infinity, @@ -71,91 +34,79 @@ describe('SocketIOTransport', () => { }); }); - test('reconnection delays are configurable', () => { - new SocketIOTransport('http://localhost:8081', { - reconnectionDelay: 250, - reconnectionDelayMax: 2000, - }).connect(); + test('reports connection changes with the Socket.IO reason', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + const changes: Array<[boolean, string | undefined]> = []; + transport.onConnectionChange((connected, reason) => changes.push([connected, reason])); + + transport.connect(); + sio.fire('connect'); + sio.fire('disconnect', 'transport close'); - expect(ioOptions).toMatchObject({ reconnectionDelay: 250, reconnectionDelayMax: 2000 }); + expect(changes).toEqual([[true, undefined], [false, 'transport close']]); }); - test('dispatches connect, disconnect and named payloads', () => { + test('delivers AG-UI events, and presents agents and config as CUSTOM events', () => { const transport = new SocketIOTransport('http://localhost:8081'); - const received: Array<[string, unknown]> = []; - for (const name of ['connect', 'disconnect', 'event', 'agents', 'config'] as const) { - transport.on(name, data => received.push([name, data])); - } + const events: unknown[] = []; + transport.onEvent(event => events.push(event)); transport.connect(); - - emitSocketEvent('connect'); - emitSocketEvent('event', { type: 'RUN_STARTED' }); - emitSocketEvent('agents', { agents: [], defaultAgent: 'claude' }); - emitSocketEvent('config', { langfuseEnabled: true }); - emitSocketEvent('disconnect', 'transport close'); - - expect(received).toEqual([ - ['connect', undefined], - ['event', { type: 'RUN_STARTED' }], - ['agents', { agents: [], defaultAgent: 'claude' }], - ['config', { langfuseEnabled: true }], - ['disconnect', 'transport close'], + sio.fire('event', { type: 'RUN_STARTED' }); + sio.fire('agents', { agents: [], defaultAgent: 'claude' }); + sio.fire('config', { langfuseEnabled: true }); + + expect(events).toEqual([ + { type: 'RUN_STARTED' }, + expect.objectContaining({ type: 'CUSTOM', name: 'agents', value: { agents: [], defaultAgent: 'claude' } }), + expect.objectContaining({ type: 'CUSTOM', name: 'config', value: { langfuseEnabled: true } }), ]); }); test('logs a warning on connection error without throwing', () => { - const transport = new SocketIOTransport('http://localhost:8081'); - transport.connect(); + new SocketIOTransport('http://localhost:8081').connect(); - emitSocketEvent('connect_error', new Error('Connection refused')); + sio.fire('connect_error', new Error('Connection refused')); // Warn, not error: console.error triggers the Next.js error overlay. expect(consoleWarnSpy).toHaveBeenCalledWith('[UseAI] Connection error:', 'Connection refused'); }); - test('repeated connection errors do not dispatch a connection state change', () => { + test('repeated connection errors do not report a connection change', () => { const transport = new SocketIOTransport('http://localhost:8081'); - const states: string[] = []; - transport.on('connect', () => states.push('connect')); - transport.on('disconnect', () => states.push('disconnect')); + const changes: boolean[] = []; + transport.onConnectionChange(connected => changes.push(connected)); transport.connect(); + sio.fire('connect_error', new Error('Attempt 1 failed')); + sio.fire('connect_error', new Error('Attempt 2 failed')); + sio.fire('connect'); - emitSocketEvent('connect_error', new Error('Attempt 1 failed')); - emitSocketEvent('connect_error', new Error('Attempt 2 failed')); - emitSocketEvent('connect'); - - expect(states).toEqual(['connect']); + expect(changes).toEqual([true]); expect(consoleWarnSpy).toHaveBeenCalledTimes(2); }); test('logs transport upgrades', () => { - const transport = new SocketIOTransport('http://localhost:8081'); - transport.connect(); - emitSocketEvent('connect'); + new SocketIOTransport('http://localhost:8081').connect(); + sio.fire('connect'); expect(consoleLogSpy).toHaveBeenCalledWith('[UseAI] Transport:', 'polling'); - emitSocketEvent('engine:upgrade', { name: 'websocket' }); + sio.fire('engine:upgrade', { name: 'websocket' }); expect(consoleLogSpy).toHaveBeenCalledWith('[UseAI] Upgraded to transport:', 'websocket'); - emitSocketEvent('engine:upgradeError', { message: 'upgrade failed' }); + sio.fire('engine:upgradeError', { message: 'upgrade failed' }); expect(consoleWarnSpy).toHaveBeenCalledWith('[UseAI] Upgrade error:', 'upgrade failed'); }); test('send emits the message on the message channel', () => { const transport = new SocketIOTransport('http://localhost:8081'); transport.connect(); - mockSocket.connected = true; - emitSocketEvent('connect'); + sio.socket.connected = true; transport.send({ type: 'abort_run', data: { runId: 'run-1' } }); - expect(mockSocket.emit).toHaveBeenCalledWith('message', { - type: 'abort_run', - data: { runId: 'run-1' }, - }); + expect(sio.socket.emit).toHaveBeenCalledWith('message', { type: 'abort_run', data: { runId: 'run-1' } }); }); test('connected follows the socket', () => { @@ -163,21 +114,21 @@ describe('SocketIOTransport', () => { expect(transport.connected).toBe(false); transport.connect(); - mockSocket.connected = true; + sio.socket.connected = true; expect(transport.connected).toBe(true); - mockSocket.connected = false; + sio.socket.connected = false; expect(transport.connected).toBe(false); }); test('disconnect closes the socket and reports disconnected', () => { const transport = new SocketIOTransport('http://localhost:8081'); transport.connect(); - mockSocket.connected = true; + sio.socket.connected = true; transport.disconnect(); - expect(mockSocket.disconnect).toHaveBeenCalled(); + expect(sio.socket.disconnect).toHaveBeenCalled(); expect(transport.connected).toBe(false); }); }); diff --git a/packages/client/src/transport/SocketIOTransport.ts b/packages/client/src/transport/SocketIOTransport.ts index 3e2f5272..8fef24de 100644 --- a/packages/client/src/transport/SocketIOTransport.ts +++ b/packages/client/src/transport/SocketIOTransport.ts @@ -1,29 +1,18 @@ import { io, Socket } from 'socket.io-client'; -import type { UseAIClientMessage } from '../types'; -import { TransportHandlerRegistry } from './handlerRegistry'; -import type { UseAITransport, UseAITransportEventName } from './types'; +import { EventType, type CustomEvent } from '@meetsmore-oss/use-ai-core'; +import type { AGUIEvent, UseAIClientMessage } from '../types'; +import type { UseAITransport } from './types'; -/** - * Options for {@link SocketIOTransport}. - */ -export interface SocketIOTransportOptions { - /** - * Delay before the first reconnection attempt, in milliseconds. - * - * @default 1000 - */ - reconnectionDelay?: number; - /** - * Upper bound on the exponential backoff between reconnection attempts, in milliseconds. - * - * @default 10000 - */ - reconnectionDelayMax?: number; -} +// Reconnect indefinitely so clients recover after extended outages (mobile +// app backgrounded long enough for server pingTimeout, airplane mode, etc.). +// Socket.IO applies exponential backoff capped at RECONNECTION_DELAY_MAX, +// so steady-state retry frequency is ~one attempt per 10s. +const RECONNECTION_DELAY = 1000; +const RECONNECTION_DELAY_MAX = 10_000; /** - * Transport over Socket.IO. This is what {@link UseAIProvider} uses when given only a `serverUrl`, - * and what the bundled `@meetsmore-oss/use-ai-server` serves. + * Transport over Socket.IO. This is what {@link UseAIProvider} uses when given a `serverUrl`, + * and what the bundled `@meetsmore-oss/use-ai-server` serves by default. * * @example * ```typescript @@ -32,13 +21,8 @@ export interface SocketIOTransportOptions { */ export class SocketIOTransport implements UseAITransport { private socket: Socket | null = null; - private registry = new TransportHandlerRegistry(); - // Reconnect indefinitely so clients recover after extended outages (mobile - // app backgrounded long enough for server pingTimeout, airplane mode, etc.). - // Socket.IO applies exponential backoff capped at reconnectionDelayMax, - // so steady-state retry frequency is ~one attempt per 10s. - private reconnectionDelay: number; - private reconnectionDelayMax: number; + private eventHandlers = new Set<(event: AGUIEvent) => void>(); + private connectionHandlers = new Set<(connected: boolean, reason?: string) => void>(); /** * @param url - The URL of the UseAI server @@ -47,10 +31,7 @@ export class SocketIOTransport implements UseAITransport { * new SocketIOTransport('ws://localhost:8081'); * ``` */ - constructor(readonly url: string, options: SocketIOTransportOptions = {}) { - this.reconnectionDelay = options.reconnectionDelay ?? 1000; - this.reconnectionDelayMax = options.reconnectionDelayMax ?? 10_000; - } + constructor(readonly url: string) {} get connected(): boolean { return this.socket !== null && this.socket.connected; @@ -61,8 +42,8 @@ export class SocketIOTransport implements UseAITransport { transports: ['polling', 'websocket'], reconnection: true, reconnectionAttempts: Infinity, - reconnectionDelay: this.reconnectionDelay, - reconnectionDelayMax: this.reconnectionDelayMax, + reconnectionDelay: RECONNECTION_DELAY, + reconnectionDelayMax: RECONNECTION_DELAY_MAX, withCredentials: true, }); this.socket = socket; @@ -81,12 +62,14 @@ export class SocketIOTransport implements UseAITransport { }); } - this.registry.dispatch('connect', undefined); + this.connectionHandlers.forEach(handler => handler(true)); }); - socket.on('event', (event: unknown) => this.registry.dispatch('event', event)); - socket.on('agents', (data: unknown) => this.registry.dispatch('agents', data)); - socket.on('config', (data: unknown) => this.registry.dispatch('config', data)); + socket.on('event', (event: AGUIEvent) => this.dispatch(event)); + // The Socket.IO server sends these two on their own channels; every other + // transport carries them as AG-UI CUSTOM events, so present them the same way. + socket.on('agents', (value: unknown) => this.dispatch(customEvent('agents', value))); + socket.on('config', (value: unknown) => this.dispatch(customEvent('config', value))); socket.on('connect_error', (error: Error) => { // Use warn instead of error to avoid triggering Next.js error overlay @@ -94,22 +77,38 @@ export class SocketIOTransport implements UseAITransport { }); socket.on('disconnect', (reason: string) => { - this.registry.dispatch('disconnect', reason); + this.connectionHandlers.forEach(handler => handler(false, reason)); }); } disconnect(): void { - if (this.socket) { - this.socket.disconnect(); - this.socket = null; - } + this.socket?.disconnect(); + this.socket = null; } send(message: UseAIClientMessage): void { this.socket?.emit('message', message); } - on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void { - return this.registry.on(name, handler); + onEvent(handler: (event: AGUIEvent) => void): () => void { + this.eventHandlers.add(handler); + return () => { + this.eventHandlers.delete(handler); + }; + } + + onConnectionChange(handler: (connected: boolean, reason?: string) => void): () => void { + this.connectionHandlers.add(handler); + return () => { + this.connectionHandlers.delete(handler); + }; } + + private dispatch(event: AGUIEvent): void { + this.eventHandlers.forEach(handler => handler(event)); + } +} + +function customEvent(name: string, value: unknown): CustomEvent { + return { type: EventType.CUSTOM, name, value, timestamp: Date.now() }; } diff --git a/packages/client/src/transport/WebSocketTransport.test.ts b/packages/client/src/transport/WebSocketTransport.test.ts index 2705d430..3e89b8be 100644 --- a/packages/client/src/transport/WebSocketTransport.test.ts +++ b/packages/client/src/transport/WebSocketTransport.test.ts @@ -1,63 +1,9 @@ import { describe, test, expect, beforeEach, afterEach, spyOn } from 'bun:test'; import { WebSocketTransport } from './WebSocketTransport'; - -/** Enough of a WHATWG WebSocket for partysocket to drive. */ -class FakeWebSocket extends EventTarget { - static instances: FakeWebSocket[] = []; - - readyState = 0; - binaryType = 'blob'; - sent: string[] = []; - closeCalls = 0; - - constructor(readonly url: string) { - super(); - FakeWebSocket.instances.push(this); - } - - static get latest(): FakeWebSocket { - return FakeWebSocket.instances[FakeWebSocket.instances.length - 1]; - } - - send(data: string): void { - this.sent.push(data); - } - - close(): void { - this.closeCalls++; - this.readyState = 3; - } - - serverOpen(): void { - this.readyState = 1; - this.dispatchEvent(new Event('open')); - } - - serverSend(data: unknown): void { - this.dispatchEvent(new MessageEvent('message', { data })); - } - - serverClose(): void { - this.readyState = 3; - this.dispatchEvent(new CloseEvent('close', { code: 1006 })); - } -} - -const tick = (ms = 0) => new Promise(resolve => setTimeout(resolve, ms)); - -async function waitUntil(condition: () => boolean, timeoutMs = 500): Promise { - const deadline = Date.now() + timeoutMs; - while (!condition()) { - if (Date.now() > deadline) throw new Error('Timed out'); - await tick(1); - } -} +import { FakeWebSocket, FakeWebSocketConstructor, waitUntil } from '../../test/fakeWebSocket'; function makeTransport(options: { reconnectionDelay?: number; reconnectionDelayMax?: number } = {}) { - return new WebSocketTransport('wss://server.example', { - ...options, - WebSocket: FakeWebSocket as unknown as typeof WebSocket, - }); + return new WebSocketTransport('wss://server.example', { ...options, WebSocket: FakeWebSocketConstructor }); } /** Connects and waits for the underlying socket, which partysocket opens asynchronously. */ @@ -65,17 +11,15 @@ async function connect(transport: WebSocketTransport): Promise { const before = FakeWebSocket.instances.length; transport.connect(); await waitUntil(() => FakeWebSocket.instances.length > before); - return FakeWebSocket.latest; + return FakeWebSocket.latest!; } -const customEvent = (name: string, value: unknown) => JSON.stringify({ type: 'CUSTOM', name, value }); - describe('WebSocketTransport', () => { let consoleLogSpy: ReturnType; let consoleWarnSpy: ReturnType; beforeEach(() => { - FakeWebSocket.instances = []; + FakeWebSocket.reset(); consoleLogSpy = spyOn(console, 'log').mockImplementation(() => {}); consoleWarnSpy = spyOn(console, 'warn').mockImplementation(() => {}); }); @@ -94,54 +38,37 @@ describe('WebSocketTransport', () => { transport.disconnect(); }); - test('an AG-UI event frame is delivered on the event channel', async () => { + test('delivers each JSON frame as one AG-UI event', async () => { const transport = makeTransport(); const events: unknown[] = []; - transport.on('event', data => events.push(data)); + transport.onEvent(event => events.push(event)); const socket = await connect(transport); socket.serverOpen(); socket.serverSend(JSON.stringify({ type: 'RUN_STARTED', threadId: 't', runId: 'r' })); + socket.serverSend(JSON.stringify({ type: 'CUSTOM', name: 'agents', value: { agents: [], defaultAgent: 'claude' } })); - expect(events).toEqual([{ type: 'RUN_STARTED', threadId: 't', runId: 'r' }]); + expect(events).toEqual([ + { type: 'RUN_STARTED', threadId: 't', runId: 'r' }, + { type: 'CUSTOM', name: 'agents', value: { agents: [], defaultAgent: 'claude' } }, + ]); transport.disconnect(); }); - test('CUSTOM events named agents and config are delivered on their own channels', async () => { - const transport = makeTransport(); - const events: unknown[] = []; - const agents: unknown[] = []; - const configs: unknown[] = []; - transport.on('event', data => events.push(data)); - transport.on('agents', data => agents.push(data)); - transport.on('config', data => configs.push(data)); - - const socket = await connect(transport); - socket.serverOpen(); - socket.serverSend(customEvent('agents', { agents: [], defaultAgent: 'claude' })); - socket.serverSend(customEvent('config', { langfuseEnabled: true })); - - expect(agents).toEqual([{ agents: [], defaultAgent: 'claude' }]); - expect(configs).toEqual([{ langfuseEnabled: true }]); - expect(events).toEqual([]); - - transport.disconnect(); - }); - - test('delivers a frame to every subscriber of its channel', async () => { + test('delivers a frame to every subscriber', async () => { const transport = makeTransport(); const first: unknown[] = []; const second: unknown[] = []; - transport.on('agents', data => first.push(data)); - transport.on('agents', data => second.push(data)); + transport.onEvent(event => first.push(event)); + transport.onEvent(event => second.push(event)); const socket = await connect(transport); socket.serverOpen(); - socket.serverSend(customEvent('agents', { agents: [], defaultAgent: 'claude' })); + socket.serverSend(JSON.stringify({ type: 'RUN_FINISHED' })); - expect(first).toEqual([{ agents: [], defaultAgent: 'claude' }]); - expect(second).toEqual([{ agents: [], defaultAgent: 'claude' }]); + expect(first).toEqual([{ type: 'RUN_FINISHED' }]); + expect(second).toEqual([{ type: 'RUN_FINISHED' }]); transport.disconnect(); }); @@ -149,7 +76,7 @@ describe('WebSocketTransport', () => { test('unsubscribing stops delivery', async () => { const transport = makeTransport(); const events: unknown[] = []; - const unsubscribe = transport.on('event', data => events.push(data)); + const unsubscribe = transport.onEvent(event => events.push(event)); const socket = await connect(transport); socket.serverOpen(); @@ -162,26 +89,10 @@ describe('WebSocketTransport', () => { transport.disconnect(); }); - test('a CUSTOM event with an unknown name passes through as an event', async () => { - const transport = makeTransport(); - const events: unknown[] = []; - transport.on('event', data => events.push(data)); - - const socket = await connect(transport); - socket.serverOpen(); - socket.serverSend(customEvent('a_name_from_a_later_version', {})); - - // UseAIClient ignores event types it does not handle, so passing it on is safe. - expect(events).toEqual([{ type: 'CUSTOM', name: 'a_name_from_a_later_version', value: {} }]); - expect(transport.connected).toBe(true); - - transport.disconnect(); - }); - - test('a malformed frame is ignored, not an error', async () => { + test('a malformed or non-text frame is ignored, not an error', async () => { const transport = makeTransport(); const events: unknown[] = []; - transport.on('event', data => events.push(data)); + transport.onEvent(event => events.push(event)); const socket = await connect(transport); socket.serverOpen(); @@ -225,31 +136,30 @@ describe('WebSocketTransport', () => { transport.disconnect(); }); - test('a close after opening dispatches disconnect', async () => { + test('reports a connection change on open and on close', async () => { const transport = makeTransport({ reconnectionDelay: 10_000 }); - const states: string[] = []; - transport.on('connect', () => states.push('connect')); - transport.on('disconnect', () => states.push('disconnect')); + const changes: boolean[] = []; + transport.onConnectionChange(connected => changes.push(connected)); const socket = await connect(transport); socket.serverOpen(); socket.serverClose(); - expect(states).toEqual(['connect', 'disconnect']); + expect(changes).toEqual([true, false]); transport.disconnect(); }); - test('a failed connection attempt does not dispatch disconnect', async () => { + test('a failed connection attempt does not report a disconnection', async () => { const transport = makeTransport({ reconnectionDelay: 10_000 }); - const states: string[] = []; - transport.on('disconnect', () => states.push('disconnect')); + const changes: boolean[] = []; + transport.onConnectionChange(connected => changes.push(connected)); const socket = await connect(transport); // Never opened: the socket closes straight from the connecting state. socket.serverClose(); - expect(states).toEqual([]); + expect(changes).toEqual([]); transport.disconnect(); }); @@ -262,7 +172,7 @@ describe('WebSocketTransport', () => { socket.serverClose(); await waitUntil(() => FakeWebSocket.instances.length > 1); - FakeWebSocket.latest.serverOpen(); + FakeWebSocket.latest!.serverOpen(); expect(transport.connected).toBe(true); transport.disconnect(); @@ -274,7 +184,7 @@ describe('WebSocketTransport', () => { for (let i = 0; i < 3; i++) { const count = FakeWebSocket.instances.length; - FakeWebSocket.latest.serverClose(); + FakeWebSocket.latest!.serverClose(); await waitUntil(() => FakeWebSocket.instances.length > count); } @@ -292,7 +202,7 @@ describe('WebSocketTransport', () => { transport.disconnect(); const openedByNow = FakeWebSocket.instances.length; - await tick(20); + await new Promise(resolve => setTimeout(resolve, 20)); expect(FakeWebSocket.instances).toHaveLength(openedByNow); expect(transport.connected).toBe(false); diff --git a/packages/client/src/transport/WebSocketTransport.ts b/packages/client/src/transport/WebSocketTransport.ts index 7bf45d36..f31053bd 100644 --- a/packages/client/src/transport/WebSocketTransport.ts +++ b/packages/client/src/transport/WebSocketTransport.ts @@ -1,8 +1,6 @@ import ReconnectingWebSocket from 'partysocket/ws'; -import { EventType } from '@meetsmore-oss/use-ai-core'; -import type { UseAIClientMessage } from '../types'; -import { TransportHandlerRegistry } from './handlerRegistry'; -import type { UseAITransport, UseAITransportEventName } from './types'; +import type { AGUIEvent, UseAIClientMessage } from '../types'; +import type { UseAITransport } from './types'; /** * Options for {@link WebSocketTransport}. @@ -34,9 +32,8 @@ export interface WebSocketTransportOptions { * Transport over a plain WebSocket. Every frame is JSON text. * * Upstream, the client sends each `UseAIClientMessage` as one frame, with nothing - * around it. Downstream, the server sends one AG-UI event per frame. The `agents` - * and `config` payloads travel as AG-UI `CUSTOM` events named `agents` and `config`. - * The client ignores an event with a type or a custom name it does not know. + * around it. Downstream, the server sends one AG-UI event per frame. The client + * ignores an event type it does not handle. * * Reconnection is automatic: indefinite, with exponential backoff capped at * `reconnectionDelayMax`. The defaults match {@link SocketIOTransport}. @@ -48,10 +45,9 @@ export interface WebSocketTransportOptions { */ export class WebSocketTransport implements UseAITransport { private socket: ReconnectingWebSocket | null = null; - private registry = new TransportHandlerRegistry(); private _connected = false; - private readonly options: Required> & - Pick; + private eventHandlers = new Set<(event: AGUIEvent) => void>(); + private connectionHandlers = new Set<(connected: boolean, reason?: string) => void>(); /** * @param url - WebSocket URL of the server @@ -60,13 +56,7 @@ export class WebSocketTransport implements UseAITransport { * new WebSocketTransport('wss://your-server.com'); * ``` */ - constructor(readonly url: string, options: WebSocketTransportOptions = {}) { - this.options = { - reconnectionDelay: options.reconnectionDelay ?? 1000, - reconnectionDelayMax: options.reconnectionDelayMax ?? 10_000, - WebSocket: options.WebSocket, - }; - } + constructor(readonly url: string, private readonly options: WebSocketTransportOptions = {}) {} get connected(): boolean { return this._connected; @@ -77,8 +67,8 @@ export class WebSocketTransport implements UseAITransport { const socket = new ReconnectingWebSocket(this.url, [], { WebSocket: this.options.WebSocket, - minReconnectionDelay: this.options.reconnectionDelay, - maxReconnectionDelay: this.options.reconnectionDelayMax, + minReconnectionDelay: this.options.reconnectionDelay ?? 1000, + maxReconnectionDelay: this.options.reconnectionDelayMax ?? 10_000, reconnectionDelayGrowFactor: 2, maxRetries: Infinity, // UseAIClient only sends while connected, so nothing is queued for a later socket. @@ -88,17 +78,12 @@ export class WebSocketTransport implements UseAITransport { socket.onopen = () => { this._connected = true; - this.registry.dispatch('connect', undefined); + this.connectionHandlers.forEach(handler => handler(true)); }; socket.onmessage = (event) => { const frame = parseFrame(event.data); - if (!frame) return; - if (frame.type === EventType.CUSTOM && (frame.name === 'agents' || frame.name === 'config')) { - this.registry.dispatch(frame.name, frame.value); - return; - } - this.registry.dispatch('event', frame); + if (frame) this.eventHandlers.forEach(handler => handler(frame)); }; socket.onerror = (event) => { @@ -110,7 +95,7 @@ export class WebSocketTransport implements UseAITransport { // A close also fires for a failed attempt. Only a socket that opened reports a disconnection. if (!this._connected) return; this._connected = false; - this.registry.dispatch('disconnect', 'transport close'); + this.connectionHandlers.forEach(handler => handler(false, 'transport close')); }; } @@ -128,18 +113,22 @@ export class WebSocketTransport implements UseAITransport { this.socket?.send(JSON.stringify(message)); } - on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void { - return this.registry.on(name, handler); + onEvent(handler: (event: AGUIEvent) => void): () => void { + this.eventHandlers.add(handler); + return () => { + this.eventHandlers.delete(handler); + }; } -} -interface Frame { - type: string; - name?: string; - value?: unknown; + onConnectionChange(handler: (connected: boolean, reason?: string) => void): () => void { + this.connectionHandlers.add(handler); + return () => { + this.connectionHandlers.delete(handler); + }; + } } -function parseFrame(data: unknown): Frame | null { +function parseFrame(data: unknown): AGUIEvent | null { if (typeof data !== 'string') { console.warn('[UseAI] Ignoring non-text frame'); return null; @@ -151,8 +140,8 @@ function parseFrame(data: unknown): Frame | null { console.warn('[UseAI] Ignoring malformed frame'); return null; } - if (typeof parsed !== 'object' || parsed === null) return null; - const frame = parsed as Partial; - if (typeof frame.type !== 'string') return null; - return frame as Frame; + if (typeof parsed !== 'object' || parsed === null || typeof (parsed as { type?: unknown }).type !== 'string') { + return null; + } + return parsed as AGUIEvent; } diff --git a/packages/client/src/transport/handlerRegistry.ts b/packages/client/src/transport/handlerRegistry.ts deleted file mode 100644 index 2938237a..00000000 --- a/packages/client/src/transport/handlerRegistry.ts +++ /dev/null @@ -1,32 +0,0 @@ -import type { UseAITransportEventName } from './types'; - -type Handler = (data: unknown) => void; - -/** - * The subscribe/dispatch bookkeeping shared by the bundled transports. - * A name with no subscribers dispatches to nobody, which is what makes an - * unrecognised downstream frame a no-op rather than an error. - */ -export class TransportHandlerRegistry { - private handlers: Map> = new Map(); - - on(name: UseAITransportEventName, handler: Handler): () => void { - let set = this.handlers.get(name); - if (!set) { - set = new Set(); - this.handlers.set(name, set); - } - set.add(handler); - return () => { - set.delete(handler); - }; - } - - dispatch(name: string, data: unknown): void { - const set = this.handlers.get(name); - if (!set) return; - for (const handler of [...set]) { - handler(data); - } - } -} diff --git a/packages/client/src/transport/index.ts b/packages/client/src/transport/index.ts index 09d0e193..cf685bda 100644 --- a/packages/client/src/transport/index.ts +++ b/packages/client/src/transport/index.ts @@ -1,5 +1,4 @@ -export type { UseAITransport, UseAITransportEventName } from './types'; +export type { UseAITransport } from './types'; export { SocketIOTransport } from './SocketIOTransport'; -export type { SocketIOTransportOptions } from './SocketIOTransport'; export { WebSocketTransport } from './WebSocketTransport'; export type { WebSocketTransportOptions } from './WebSocketTransport'; diff --git a/packages/client/src/transport/types.ts b/packages/client/src/transport/types.ts index d3c79d77..3be0141a 100644 --- a/packages/client/src/transport/types.ts +++ b/packages/client/src/transport/types.ts @@ -1,24 +1,15 @@ -import type { UseAIClientMessage } from '../types'; - -/** - * Names of the downstream channels a transport delivers to {@link UseAIClient}. - * - * - `connect` / `disconnect` — connection lifecycle. `disconnect` carries a reason string. - * - `event` — an AG-UI event. - * - `agents` — the server's agent list, `{ agents, defaultAgent }`. - * - `config` — server capability flags, `{ langfuseEnabled }`. - */ -export type UseAITransportEventName = 'connect' | 'disconnect' | 'event' | 'agents' | 'config'; +import type { AGUIEvent, UseAIClientMessage } from '../types'; /** * The pipe between {@link UseAIClient} and a server. * - * A transport carries {@link UseAIClientMessage} upstream and named payloads downstream. - * It owns everything protocol-specific: how a connection is opened, how it reconnects, - * and how a named payload is framed on the wire. + * Upstream it carries {@link UseAIClientMessage}. Downstream it delivers AG-UI events. + * Two payloads that are not AG-UI events, the agent list and the server config, arrive + * as AG-UI `CUSTOM` events named `agents` and `config`. * - * Two implementations ship with the library: {@link SocketIOTransport} (the default) - * and {@link WebSocketTransport}. + * A transport owns everything protocol-specific: how the connection opens, how it + * reconnects, and how an event is framed on the wire. Two implementations ship with + * the library: {@link SocketIOTransport} (the default) and {@link WebSocketTransport}. */ export interface UseAITransport { /** @@ -27,6 +18,9 @@ export interface UseAITransport { */ readonly url: string; + /** Whether the connection is currently open. */ + readonly connected: boolean; + /** Opens the connection. Reconnection until {@link disconnect} is the transport's own responsibility. */ connect(): void; @@ -37,12 +31,15 @@ export interface UseAITransport { send(message: UseAIClientMessage): void; /** - * Subscribes to a downstream channel. - * + * Subscribes to downstream AG-UI events. * @returns Cleanup function to unsubscribe */ - on(name: UseAITransportEventName, handler: (data: unknown) => void): () => void; + onEvent(handler: (event: AGUIEvent) => void): () => void; - /** Whether the connection is currently open. */ - readonly connected: boolean; + /** + * Subscribes to the connection opening and closing. + * @param handler - Receives `true` on connect and `false` on disconnect, with the transport's reason + * @returns Cleanup function to unsubscribe + */ + onConnectionChange(handler: (connected: boolean, reason?: string) => void): () => void; } diff --git a/packages/client/test/fakeWebSocket.ts b/packages/client/test/fakeWebSocket.ts new file mode 100644 index 00000000..53163f0a --- /dev/null +++ b/packages/client/test/fakeWebSocket.ts @@ -0,0 +1,56 @@ +/** Enough of a WHATWG WebSocket for partysocket to drive. */ +export class FakeWebSocket extends EventTarget { + static instances: FakeWebSocket[] = []; + + readyState = 0; + binaryType = 'blob'; + sent: string[] = []; + closeCalls = 0; + + constructor(readonly url: string) { + super(); + FakeWebSocket.instances.push(this); + } + + static reset(): void { + FakeWebSocket.instances = []; + } + + static get latest(): FakeWebSocket | undefined { + return FakeWebSocket.instances[FakeWebSocket.instances.length - 1]; + } + + send(data: string): void { + this.sent.push(data); + } + + close(): void { + this.closeCalls++; + this.readyState = 3; + } + + serverOpen(): void { + this.readyState = 1; + this.dispatchEvent(new Event('open')); + } + + serverSend(data: unknown): void { + this.dispatchEvent(new MessageEvent('message', { data })); + } + + serverClose(): void { + this.readyState = 3; + this.dispatchEvent(new CloseEvent('close', { code: 1006 })); + } +} + +/** The constructor as `WebSocketTransportOptions.WebSocket` expects it. */ +export const FakeWebSocketConstructor = FakeWebSocket as unknown as typeof WebSocket; + +export async function waitUntil(condition: () => boolean, timeoutMs = 500): Promise { + const deadline = Date.now() + timeoutMs; + while (!condition()) { + if (Date.now() > deadline) throw new Error('Timed out waiting for condition'); + await new Promise(resolve => setTimeout(resolve, 1)); + } +} diff --git a/packages/client/test/socketIOMock.ts b/packages/client/test/socketIOMock.ts new file mode 100644 index 00000000..72196bed --- /dev/null +++ b/packages/client/test/socketIOMock.ts @@ -0,0 +1,60 @@ +import { mock } from 'bun:test'; +import type { Socket } from 'socket.io-client'; + +export type MockSocket = Partial & { connected: boolean }; + +export interface SocketIOMock { + /** The socket the most recent `io()` call returned. */ + readonly socket: MockSocket; + /** The options the most recent `io()` call received. */ + readonly ioOptions: Record | undefined; + /** Fires a socket event, or an engine event as `engine:`. */ + fire(event: string, ...args: unknown[]): void; +} + +/** + * Replaces `socket.io-client` so `io()` returns a scriptable socket. + * Call at the top of the test file, before importing anything that imports the transport. + */ +export function installSocketIOMock(): SocketIOMock { + let handlers: Record = {}; + let socket: MockSocket; + let ioOptions: Record | undefined; + + mock.module('socket.io-client', () => ({ + io: (_url: string, options: Record) => { + ioOptions = options; + handlers = {}; + socket = { + on: mock((event: string, handler: Function) => { + (handlers[event] ??= []).push(handler); + return socket as Socket; + }), + emit: mock(() => socket as Socket), + connected: false, + disconnect: mock(() => socket as Socket), + io: { + engine: { + transport: { name: 'polling' }, + on: mock((event: string, handler: Function) => { + (handlers[`engine:${event}`] ??= []).push(handler); + }), + }, + } as never, + }; + return socket; + }, + })); + + return { + get socket() { + return socket; + }, + get ioOptions() { + return ioOptions; + }, + fire(event, ...args) { + handlers[event]?.forEach(handler => handler(...args)); + }, + }; +} diff --git a/packages/server/src/agents/types.ts b/packages/server/src/agents/types.ts index 5f4a6ed9..dc96ef33 100644 --- a/packages/server/src/agents/types.ts +++ b/packages/server/src/agents/types.ts @@ -1,19 +1,23 @@ import type { ModelMessage } from 'ai'; -import type { ToolDefinition, AGUIEvent, ToolApprovalRequestEvent } from '../types'; +import type { ToolDefinition, AGUIEvent, ToolApprovalRequestEvent, UseAIClientMessage } from '../types'; /** - * A client connection, as much of it as a session needs. - * - * Both connection kinds the server accepts satisfy this: a Socket.IO socket, and a - * plain WebSocket where `emit(name, data)` is written out as `{"name":...,"data":...}`. + * One client connection, whatever protocol carries it. The server accepts every + * connection through this interface, so protocol details stay in its implementations. */ export interface ClientConnection { /** Identifies the connection for the lifetime of the session. */ readonly id: string; + /** Address the connection came from, used for rate limiting. */ + readonly ipAddress: string; /** Whether the connection is still open. */ readonly connected: boolean; - /** Sends a named payload to the client. */ + /** Sends a named payload to the client. `event` carries an AG-UI event. */ emit(name: string, data?: unknown): void; + /** Registers the handler for messages from the client. */ + onMessage(handler: (message: UseAIClientMessage) => void): void; + /** Registers the handler for the connection closing, for any reason. */ + onClose(handler: () => void): void; } /** diff --git a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts index 421d39c1..ae359dcd 100644 --- a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts +++ b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts @@ -1,10 +1,21 @@ +import type { Server as SocketIOServer } from 'socket.io'; import { Server as BunEngine } from '@socket.io/bun-engine'; -import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle } from '../types'; +import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle, WebSocketListener } from '../types'; import { resolveCorsHeaders, resolvePreflightHeaders } from './cors'; -import { BunRawWebSocket, type RawWebSocketData } from './rawWebSocket'; - -type BunServer = Parameters[0]['fetch']>>[1]; -type WebSocketHandler = NonNullable[0]['websocket']>; +import { BunRawWebSocket, type BunWebSocketData } from './rawWebSocket'; + +type ServeOptions = Parameters[0]; +type BunServer = Parameters>[1]; +type WebSocketHandler = NonNullable; + +/** The part of the HTTP server a listener owns. */ +interface ListenerHandlers { + /** Requests under this path go to `upgrade`; everything else is 404. */ + path: string; + /** Answers the request, or returns undefined once the request was upgraded. */ + upgrade(req: Request, server: BunServer): Promise; + websocket: WebSocketHandler; +} /** * Runtime adapter for Bun. @@ -17,7 +28,7 @@ export class BunRuntimeAdapter implements RuntimeAdapter { private engine: BunEngine | null = null; createServer(listener: RuntimeListener, config: RuntimeServerConfig): RuntimeServerHandle { - const { upgrade, websocket } = + const handlers = listener.transport === 'socketio' ? this.socketIOHandlers(listener.io, config) : this.webSocketHandlers(listener.onConnection); @@ -48,27 +59,20 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }); } - const response = await upgrade(req, server, url); - if (response === null) return undefined; - if (response) { - // Add CORS headers to the listener's responses - if (Object.keys(corsHeaders).length > 0) { - const newHeaders = new Headers(response.headers); - for (const [key, value] of Object.entries(corsHeaders)) { - newHeaders.set(key, value); - } - return new Response(response.body, { - status: response.status, - statusText: response.statusText, - headers: newHeaders, - }); - } - return response; + if (!url.pathname.startsWith(handlers.path)) { + return new Response('Not Found', { status: 404, headers: corsHeaders }); } - return new Response('Not Found', { status: 404, headers: corsHeaders }); + const response = await handlers.upgrade(req, server); + if (!response || Object.keys(corsHeaders).length === 0) return response; + + const headers = new Headers(response.headers); + for (const [key, value] of Object.entries(corsHeaders)) { + headers.set(key, value); + } + return new Response(response.body, { status: response.status, statusText: response.statusText, headers }); }, - websocket, + websocket: handlers.websocket, }); return { @@ -79,11 +83,7 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }; } - /** - * @returns `upgrade` yields a Response to send, `null` once the request was upgraded, - * or `undefined` when the path is not the listener's. - */ - private socketIOHandlers(io: RuntimeListener extends infer L ? (L extends { io: infer I } ? I : never) : never, config: RuntimeServerConfig) { + private socketIOHandlers(io: SocketIOServer, config: RuntimeServerConfig): ListenerHandlers { this.engine = new BunEngine({ path: '/socket.io/', maxHttpBufferSize: config.maxHttpBufferSize, @@ -103,43 +103,35 @@ export class BunRuntimeAdapter implements RuntimeAdapter { const handleRequest = this.engine.handleRequest.bind(this.engine); return { - upgrade: async (req: Request, server: BunServer, url: URL): Promise => { - if (!url.pathname.startsWith('/socket.io/')) return undefined; - // The engine returns undefined once it has upgraded the request itself. - return (await handleRequest(req, server as never)) ?? null; - }, + path: '/socket.io/', + upgrade: (req, server) => handleRequest(req, server as never), websocket: this.engine.handler().websocket as unknown as WebSocketHandler, }; } - private webSocketHandlers(onConnection: (connection: BunRawWebSocket) => void) { - const sockets = new WeakMap(); - const websocket: WebSocketHandler = { - open: (ws) => { - const { remoteAddress } = ws.data as RawWebSocketData; - const connection = new BunRawWebSocket(ws, remoteAddress); - sockets.set(ws, connection); - onConnection(connection); - }, - message: (ws, message) => { - sockets.get(ws)?.receiveMessage( - typeof message === 'string' ? message : new TextDecoder().decode(message), - ); - }, - close: (ws) => { - sockets.get(ws)?.receiveClose(); - sockets.delete(ws); - }, - }; - + private webSocketHandlers(onConnection: WebSocketListener['onConnection']): ListenerHandlers { + const dataOf = (ws: { data: unknown }) => ws.data as BunWebSocketData; return { - upgrade: async (req: Request, server: BunServer, url: URL): Promise => { - if (url.pathname !== '/') return undefined; - const data: RawWebSocketData = { remoteAddress: server.requestIP(req)?.address }; - if (server.upgrade(req, { data })) return null; + path: '/', + upgrade: async (req, server) => { + if (new URL(req.url).pathname !== '/') return new Response('Not Found', { status: 404 }); + const data: BunWebSocketData = { remoteAddress: server.requestIP(req)?.address }; + if (server.upgrade(req, { data })) return undefined; return new Response('Expected a WebSocket upgrade', { status: 426 }); }, - websocket, + websocket: { + open: (ws) => { + const data = dataOf(ws); + data.connection = new BunRawWebSocket(ws, data.remoteAddress); + onConnection(data.connection); + }, + message: (ws, message) => { + if (typeof message === 'string') dataOf(ws).connection?.receiveMessage(message); + }, + close: (ws) => { + dataOf(ws).connection?.receiveClose(); + }, + }, }; } } diff --git a/packages/server/src/runtime/bun/rawWebSocket.ts b/packages/server/src/runtime/bun/rawWebSocket.ts index 56743648..95931111 100644 --- a/packages/server/src/runtime/bun/rawWebSocket.ts +++ b/packages/server/src/runtime/bun/rawWebSocket.ts @@ -1,8 +1,9 @@ import type { RawWebSocket } from '../types'; -/** Per-connection data attached at upgrade time. */ -export interface RawWebSocketData { +/** Per-connection data on `ws.data`: set at upgrade, completed on open. */ +export interface BunWebSocketData { remoteAddress?: string; + connection?: BunRawWebSocket; } interface BunWebSocket { diff --git a/packages/server/src/runtime/clientIp.ts b/packages/server/src/runtime/clientIp.ts index 9d54833a..796f6b87 100644 --- a/packages/server/src/runtime/clientIp.ts +++ b/packages/server/src/runtime/clientIp.ts @@ -41,6 +41,16 @@ export interface ClientIpTracker { getClientIp(conn: ClientIpConnection): string | undefined; } +/** + * The client address behind a proxy: the first hop of `X-Forwarded-For`, else the fallback. + */ +export function forwardedClientIp( + forwardedFor: string | string[] | undefined, + fallback: string | undefined, +): string | undefined { + return typeof forwardedFor === 'string' ? forwardedFor.split(',')[0].trim() : fallback; +} + /** * Creates a ClientIpTracker instance. * diff --git a/packages/server/src/runtime/index.ts b/packages/server/src/runtime/index.ts index 40537b50..503b0ec5 100644 --- a/packages/server/src/runtime/index.ts +++ b/packages/server/src/runtime/index.ts @@ -9,9 +9,11 @@ export type { RuntimeServerHandle, RawWebSocket, RuntimeListener, + SocketIOListener, + WebSocketListener, } from './types'; export { detectRuntime } from './detection'; -export { createClientIpTracker, type ClientIpTracker, type ClientIpConnection } from './clientIp'; +export { createClientIpTracker, forwardedClientIp, type ClientIpTracker, type ClientIpConnection } from './clientIp'; /** * Creates a runtime adapter for the specified or detected runtime. diff --git a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts index 18e770de..59efffec 100644 --- a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts +++ b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts @@ -1,6 +1,8 @@ import { createServer, type Server as HttpServer } from 'http'; +import type { Server as SocketIOServer } from 'socket.io'; import { WebSocketServer } from 'ws'; -import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle } from '../types'; +import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle, WebSocketListener } from '../types'; +import { forwardedClientIp } from '../clientIp'; import { NodeRawWebSocket } from './rawWebSocket'; /** @@ -32,12 +34,9 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { res.end('Not Found'); }); - let wss: WebSocketServer | null = null; - if (listener.transport === 'socketio') { - this.attachSocketIO(listener.io, httpServer, config); - } else { - wss = this.attachWebSocket(listener.onConnection, httpServer, config); - } + const wss = listener.transport === 'socketio' + ? this.attachSocketIO(listener.io, httpServer, config) + : this.attachWebSocket(listener.onConnection, httpServer, config); httpServer.listen(config.port); @@ -50,11 +49,7 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { }; } - private attachSocketIO( - io: Extract['io'], - httpServer: HttpServer, - config: RuntimeServerConfig, - ) { + private attachSocketIO(io: SocketIOServer, httpServer: HttpServer, config: RuntimeServerConfig): null { io.attach(httpServer, { transports: ['polling', 'websocket'], maxHttpBufferSize: config.maxHttpBufferSize, @@ -69,20 +64,18 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { if (config.onPollingConnection) { io.engine.on('connection', (socket) => { if (socket.transport.name === 'polling') { - const xForwardedFor = socket.request.headers['x-forwarded-for']; - const ip = typeof xForwardedFor === 'string' - ? xForwardedFor.split(',')[0].trim() - : socket.request.socket?.remoteAddress; + const ip = forwardedClientIp(socket.request.headers['x-forwarded-for'], socket.request.socket?.remoteAddress); if (ip) { config.onPollingConnection!(socket.id, ip); } } }); } + return null; } private attachWebSocket( - onConnection: Extract['onConnection'], + onConnection: WebSocketListener['onConnection'], httpServer: HttpServer, config: RuntimeServerConfig, ): WebSocketServer { @@ -94,10 +87,7 @@ export class NodeRuntimeAdapter implements RuntimeAdapter { return; } wss.handleUpgrade(req, socket, head, (ws) => { - const forwardedFor = req.headers['x-forwarded-for']; - const remoteAddress = typeof forwardedFor === 'string' - ? forwardedFor.split(',')[0].trim() - : req.socket.remoteAddress; + const remoteAddress = forwardedClientIp(req.headers['x-forwarded-for'], req.socket.remoteAddress); onConnection(new NodeRawWebSocket(ws, remoteAddress)); }); }); diff --git a/packages/server/src/runtime/node/rawWebSocket.ts b/packages/server/src/runtime/node/rawWebSocket.ts index db26abcc..b2b1a617 100644 --- a/packages/server/src/runtime/node/rawWebSocket.ts +++ b/packages/server/src/runtime/node/rawWebSocket.ts @@ -21,7 +21,7 @@ export class NodeRawWebSocket implements RawWebSocket { onMessage(handler: (data: string) => void): void { this.ws.on('message', (data: unknown, isBinary: boolean) => { - handler(isBinary ? '' : String(data)); + if (!isBinary) handler(String(data)); }); } diff --git a/packages/server/src/runtime/types.ts b/packages/server/src/runtime/types.ts index 6f46f213..d694ea78 100644 --- a/packages/server/src/runtime/types.ts +++ b/packages/server/src/runtime/types.ts @@ -25,12 +25,20 @@ export interface RawWebSocket { onClose(handler: () => void): void; } -/** - * What the HTTP server hands connections to. A server runs exactly one. - */ -export type RuntimeListener = - | { transport: 'socketio'; io: SocketIOServer } - | { transport: 'websocket'; onConnection(connection: RawWebSocket): void }; +/** Hands the HTTP server's `/socket.io/` traffic to a Socket.IO server. */ +export interface SocketIOListener { + transport: 'socketio'; + io: SocketIOServer; +} + +/** Accepts plain WebSocket upgrades at `/`. */ +export interface WebSocketListener { + transport: 'websocket'; + onConnection(connection: RawWebSocket): void; +} + +/** What the HTTP server hands connections to. A server runs exactly one. */ +export type RuntimeListener = SocketIOListener | WebSocketListener; /** * Configuration for creating a runtime server. diff --git a/packages/server/src/server.ts b/packages/server/src/server.ts index f4622627..c2424bc0 100644 --- a/packages/server/src/server.ts +++ b/packages/server/src/server.ts @@ -1,4 +1,4 @@ -import { Server as SocketIOServer, Socket } from 'socket.io'; +import { Server as SocketIOServer } from 'socket.io'; import { ModelMessage, ToolModelMessage } from 'ai'; import { createHash } from 'crypto'; import { EventType, type McpHeadersMap, type UseAIForwardedProps, type ResolveAttachments } from '@meetsmore-oss/use-ai-core'; @@ -36,9 +36,9 @@ import { type RuntimeAdapter, type RuntimeServerHandle, type ClientIpTracker, - type RawWebSocket, type RuntimeListener, } from './runtime'; +import { SocketIOClientConnection } from './socketIOConnection'; import { WebSocketClientConnection } from './webSocketConnection'; // Re-export session types for external use @@ -99,6 +99,8 @@ export class UseAIServer { private defaultAgentId: string; // ID of the default agent private agents: Record; // Registry of all agents private clients: Map = new Map(); + // Sent to every client on connect; constant after construction. + private agentsPayload: { agents: Array<{ id: string; name: string; annotation?: string }>; defaultAgent: string }; private config: Required> & { maxHttpBufferSize: number; cors?: CorsOptions; @@ -150,6 +152,14 @@ export class UseAIServer { } this.agent = defaultAgent; this.defaultAgentId = config.defaultAgent; + this.agentsPayload = { + agents: Object.entries(this.agents).map(([id, agent]) => ({ + id, + name: agent.getName?.() || id, + annotation: agent.getAnnotation?.(), + })), + defaultAgent: this.defaultAgentId, + }; this.rateLimiter = new RateLimiter({ maxRequests: this.config.rateLimitMaxRequests, @@ -199,12 +209,21 @@ export class UseAIServer { this.runtimeAdapter = createRuntimeAdapter(config.runtime ?? 'auto'); logger.info('Using runtime adapter', { runtime: this.runtimeAdapter.name }); + let listener: RuntimeListener; if (this.config.transport === 'socketio') { this.io = new SocketIOServer({ transports: ['polling', 'websocket'], maxHttpBufferSize: this.config.maxHttpBufferSize, }); - this.setupSocketIOServer(this.io); + this.io.on('connection', (socket) => { + this.acceptConnection(new SocketIOClientConnection(socket, this.clientIpTracker)); + }); + listener = { transport: 'socketio', io: this.io }; + } else { + listener = { + transport: 'websocket', + onConnection: (socket) => this.acceptConnection(new WebSocketClientConnection(`ws-${uuidv4()}`, socket)), + }; } if (this.rateLimiter.isEnabled()) { @@ -214,10 +233,6 @@ export class UseAIServer { }); } - const listener: RuntimeListener = this.io - ? { transport: 'socketio', io: this.io } - : { transport: 'websocket', onConnection: (socket) => this.handleWebSocketConnection(socket) }; - this.serverHandle = this.runtimeAdapter.createServer(listener, { port: this.config.port, idleTimeout: this.config.idleTimeout, @@ -292,83 +307,32 @@ export class UseAIServer { logger.debug('Registered message handler', { type }); } - private setupSocketIOServer(io: SocketIOServer) { - io.on('connection', (socket: Socket) => { - // Get connection info for IP address resolution - const conn = socket.conn as unknown as { id: string; transport: { name: string; socket?: { remoteAddress?: string } } }; - // Get IP address for rate limiting: - // 1. Try clientIpTracker (works for polling transport) - // 2. Fall back to socket.handshake.address (works for WebSocket) - // 3. Last resort: use socket.id - const ipAddress = this.clientIpTracker.getClientIp(conn) - || socket.handshake.address - || socket.id; - - const session = this.createSession(socket, ipAddress); - logger.info('Client connected', { - clientId: session.clientId, - threadId: session.threadId, - ipAddress, - transport: conn.transport.name, - }); - - // Log transport upgrades - socket.conn.on('upgrade', (transport) => { - logger.info('Client upgraded transport', { clientId: session.clientId, transport: transport.name }); - }); - - socket.on('message', (message: UseAIClientMessage) => this.receiveClientMessage(session, message)); - - socket.on('disconnect', () => { - logger.info('Client disconnected', { clientId: session.clientId, ipAddress }); - // Clean up polling IP entry - this.clientIpTracker.removePollingConnection(conn.id); - this.destroySession(session); - }); - }); - } - - /** - * Accepts a plain WebSocket connection. Frames are JSON text: upstream the client - * message on its own, downstream one AG-UI event. See docs/websocket-protocol.md. - */ - private handleWebSocketConnection(socket: RawWebSocket) { - const connection = new WebSocketClientConnection(`ws-${uuidv4()}`, socket); - const ipAddress = socket.remoteAddress || connection.id; - - const session = this.createSession(connection, ipAddress); + private acceptConnection(connection: ClientConnection) { + const session = this.createSession(connection); logger.info('Client connected', { clientId: session.clientId, threadId: session.threadId, - ipAddress, - transport: 'websocket', + ipAddress: session.ipAddress, }); - socket.onMessage((data) => { - let message: UseAIClientMessage; - try { - message = JSON.parse(data) as UseAIClientMessage; - } catch { - logger.warn('Discarding malformed frame', { clientId: session.clientId }); - return; - } + connection.onMessage((message) => { void this.receiveClientMessage(session, message); }); - socket.onClose(() => { - logger.info('Client disconnected', { clientId: session.clientId, ipAddress }); + connection.onClose(() => { + logger.info('Client disconnected', { clientId: session.clientId, ipAddress: session.ipAddress }); this.destroySession(session); }); } /** * Creates the session for a newly accepted connection and announces the server's - * agents to it. Both connection kinds land here. + * agents to it. */ - private createSession(connection: ClientConnection, ipAddress: string): ClientSession { + private createSession(connection: ClientConnection): ClientSession { const session: ClientSession = { clientId: `client-${++this.clientIdCounter}`, - ipAddress, + ipAddress: connection.ipAddress, socket: connection, threadId: uuidv4(), tools: [], @@ -379,14 +343,7 @@ export class UseAIServer { this.clients.set(connection.id, session); - connection.emit('agents', { - agents: Object.entries(this.agents).map(([id, agent]) => ({ - id, - name: agent.getName?.() || id, - annotation: agent.getAnnotation?.(), - })), - defaultAgent: this.defaultAgentId, - }); + connection.emit('agents', this.agentsPayload); for (const plugin of this.plugins) { plugin.onClientConnect?.(session); @@ -415,7 +372,7 @@ export class UseAIServer { */ private async receiveClientMessage(session: ClientSession, message: UseAIClientMessage) { try { - await this.handleClientMessage(session.socket, message); + await this.handleClientMessage(session, message); } catch (error) { logger.error('Error handling message', { error: error instanceof Error ? error.message : 'Unknown error', @@ -442,10 +399,7 @@ export class UseAIServer { } } - private async handleClientMessage(connection: ClientConnection, message: UseAIClientMessage) { - const session = this.clients.get(connection.id); - if (!session) return; - + private async handleClientMessage(session: ClientSession, message: UseAIClientMessage) { // Check if a plugin has registered a handler for this message type const pluginHandler = this.messageHandlers.get(message.type); if (pluginHandler) { @@ -989,9 +943,7 @@ export class UseAIServer { } private sendEvent(connection: ClientConnection, event: T) { - if (connection.connected) { - connection.emit('event', event); - } + connection.emit('event', event); } /** diff --git a/packages/server/src/socketIOConnection.ts b/packages/server/src/socketIOConnection.ts new file mode 100644 index 00000000..372e2536 --- /dev/null +++ b/packages/server/src/socketIOConnection.ts @@ -0,0 +1,47 @@ +import type { Socket } from 'socket.io'; +import type { ClientConnection } from './agents/types'; +import type { UseAIClientMessage } from './types'; +import type { ClientIpConnection, ClientIpTracker } from './runtime'; +import { logger } from './logger'; + +/** + * A Socket.IO socket presented as a {@link ClientConnection}. + */ +export class SocketIOClientConnection implements ClientConnection { + readonly id: string; + readonly ipAddress: string; + + private readonly conn: ClientIpConnection; + + constructor(private socket: Socket, private ipTracker: ClientIpTracker) { + this.id = socket.id; + this.conn = socket.conn as unknown as ClientIpConnection; + // Polling connections record their address at engine level, since the transport + // socket is not available later; WebSocket connections carry it on the handshake. + this.ipAddress = ipTracker.getClientIp(this.conn) || socket.handshake.address || socket.id; + + logger.info('Socket.IO connection', { connectionId: this.id, transport: socket.conn.transport.name }); + socket.conn.on('upgrade', (transport) => { + logger.info('Socket.IO connection upgraded transport', { connectionId: this.id, transport: transport.name }); + }); + } + + get connected(): boolean { + return this.socket.connected; + } + + emit(name: string, data?: unknown): void { + if (this.socket.connected) this.socket.emit(name, data); + } + + onMessage(handler: (message: UseAIClientMessage) => void): void { + this.socket.on('message', handler); + } + + onClose(handler: () => void): void { + this.socket.on('disconnect', () => { + this.ipTracker.removePollingConnection(this.conn.id); + handler(); + }); + } +} diff --git a/packages/server/src/webSocketConnection.ts b/packages/server/src/webSocketConnection.ts index 4b02a251..df878e77 100644 --- a/packages/server/src/webSocketConnection.ts +++ b/packages/server/src/webSocketConnection.ts @@ -1,16 +1,23 @@ -import { EventType } from '@meetsmore-oss/use-ai-core'; +import { EventType, type CustomEvent } from '@meetsmore-oss/use-ai-core'; import type { ClientConnection } from './agents/types'; +import type { UseAIClientMessage } from './types'; import type { RawWebSocket } from './runtime'; +import { logger } from './logger'; /** * A plain WebSocket presented as a {@link ClientConnection}. * - * The downstream stream is AG-UI: `emit('event', e)` writes the event as one JSON text - * frame, and any other name goes out as an AG-UI `CUSTOM` event carrying that name. - * `WebSocketTransport` on the client reads exactly this. + * Both directions are JSON text frames. Upstream, each frame is one client message. + * Downstream, each frame is one AG-UI event: `emit('event', e)` writes `e` as is, and + * any other name goes out as an AG-UI `CUSTOM` event carrying that name. The client's + * `WebSocketTransport` reads exactly this. */ export class WebSocketClientConnection implements ClientConnection { - constructor(readonly id: string, private socket: RawWebSocket) {} + readonly ipAddress: string; + + constructor(readonly id: string, private socket: RawWebSocket) { + this.ipAddress = socket.remoteAddress || id; + } get connected(): boolean { return this.socket.open; @@ -18,9 +25,28 @@ export class WebSocketClientConnection implements ClientConnection { emit(name: string, data?: unknown): void { if (!this.socket.open) return; - const frame = name === 'event' - ? data - : { type: EventType.CUSTOM, name, value: data, timestamp: Date.now() }; + const frame: unknown = name === 'event' ? data : customEvent(name, data); this.socket.send(JSON.stringify(frame)); } + + onMessage(handler: (message: UseAIClientMessage) => void): void { + this.socket.onMessage((data) => { + let message: UseAIClientMessage; + try { + message = JSON.parse(data) as UseAIClientMessage; + } catch { + logger.warn('Discarding malformed frame', { connectionId: this.id }); + return; + } + handler(message); + }); + } + + onClose(handler: () => void): void { + this.socket.onClose(handler); + } +} + +function customEvent(name: string, value: unknown): CustomEvent { + return { type: EventType.CUSTOM, name, value, timestamp: Date.now() }; } diff --git a/packages/server/src/websocket-transport.integration.test.ts b/packages/server/src/websocket-transport.integration.test.ts index 790e6c59..a3769921 100644 --- a/packages/server/src/websocket-transport.integration.test.ts +++ b/packages/server/src/websocket-transport.integration.test.ts @@ -7,6 +7,7 @@ import { createSequentialMockModel, TestCleanupManager, } from '../test/integration-test-utils'; +import { waitFor } from '../test/test-utils'; import { AISDKAgent } from './agents/AISDKAgent'; /** @@ -36,14 +37,6 @@ function connectClient(port: number): Promise<{ client: UseAIClient; events: AGU }); } -async function waitFor(condition: () => boolean, message: string, timeoutMs = 5000): Promise { - const deadline = Date.now() + timeoutMs; - while (!condition()) { - if (Date.now() > deadline) throw new Error(`Timed out waiting for ${message}`); - await new Promise(resolve => setTimeout(resolve, 10)); - } -} - describe.each(RUNTIMES)('WebSocketTransport against a real server: %s runtime', (runtime) => { const cleanup = new TestCleanupManager(); const port = runtime === 'bun' ? 9530 : 9540; @@ -172,6 +165,13 @@ describe('transport defaults to socketio', () => { expect(socket.connected).toBe(true); socket.disconnect(); - await expect(connectClient(port)).rejects.toThrow('Timed out connecting'); - }, 10_000); + // A Socket.IO server answers a plain upgrade at / with 404, so the socket closes at once. + const closed = await new Promise((resolve) => { + const ws = new WebSocket(`ws://localhost:${port}`); + ws.onopen = () => resolve(false); + ws.onclose = () => resolve(true); + ws.onerror = () => {}; + }); + expect(closed).toBe(true); + }); }); diff --git a/packages/server/test/test-utils.ts b/packages/server/test/test-utils.ts index da3c2f9c..f4e62f25 100644 --- a/packages/server/test/test-utils.ts +++ b/packages/server/test/test-utils.ts @@ -36,6 +36,17 @@ export function waitForConnection(socket: Socket): Promise { }); } +/** + * Polls until the condition holds. + */ +export async function waitFor(condition: () => boolean, message: string, timeoutMs = 5000): Promise { + const deadline = Date.now() + timeoutMs; + while (!condition()) { + if (Date.now() > deadline) throw new Error(`Timed out waiting for ${message}`); + await new Promise(resolve => setTimeout(resolve, 10)); + } +} + /** * Wait for an AG-UI event from the Socket.IO server */