diff --git a/CLAUDE.md b/CLAUDE.md index e51248f5e..0f0b6bdda 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 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 ### Data Flow diff --git a/README.md b/README.md index 84ce4d367..0154faf2b 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,35 @@ 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 | 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'` | + +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. + +```tsx +import { UseAIProvider, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; + +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. 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. + ## Client ### `useAI` hook @@ -353,6 +383,8 @@ root.render( ); ``` +Pass `transport` instead of `serverUrl` 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 +1080,7 @@ const server = new UseAIServer({ }) }, defaultAgent: 'claude', + 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 4e5965817..8ee640502 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 5178962af..e7498d141 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", @@ -137,15 +138,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 +549,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=="], @@ -767,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=="], @@ -1097,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 new file mode 100644 index 000000000..f8d11fde9 --- /dev/null +++ b/docs/websocket-protocol.md @@ -0,0 +1,140 @@ +# Plain WebSocket protocol + +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 `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 + +```tsx +import { UseAIProvider, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; + +root.render( + + + +); +``` + +Give the provider `serverUrl` or `transport`, not both. `serverUrl` connects over Socket.IO. `transport` connects over the transport that you pass. + +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. + +## Bundled server + +The bundled server serves one transport. The default is Socket.IO. + +```typescript +const server = new UseAIServer({ + agents: { claude }, + defaultAgent: 'claude', + 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 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` +- `message_feedback` + +Plugins add more. See `UseAIClientMessage` in `@meetsmore-oss/use-ai-core` for each payload. + +## Downstream frames + +The server sends one AG-UI event per frame. Each event has a `type` field. + +```json +{ "type": "RUN_STARTED", "threadId": "...", "runId": "...", "timestamp": 1700000000000 } +{ "type": "TEXT_MESSAGE_CONTENT", "messageId": "...", "delta": "Hello" } +{ "type": "RUN_FINISHED", "threadId": "...", "runId": "..." } +``` + +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 handles these event types: + +- `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 + +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 `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 streams the rest of the output. +8. The server sends `RUN_FINISHED`. + +## Reconnection + +`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', { + reconnectionDelay: 1000, // first retry, 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`. + +## Your own transport + +`UseAITransport` has seven members: + +- `url` +- `connected` +- `connect` +- `disconnect` +- `send` +- `onEvent`, which delivers AG-UI events +- `onConnectionChange` + +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'; + +const client = new UseAIClient(myTransport); +``` diff --git a/packages/client/package.json b/packages/client/package.json index f9dea4cc5..e9afdde9e 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 5036feee7..e1550eb56 100644 --- a/packages/client/src/client.test.ts +++ b/packages/client/src/client.test.ts @@ -1,265 +1,219 @@ -import { describe, test, expect, mock, beforeEach, afterEach, spyOn } from 'bun:test'; -import type { Socket } from 'socket.io-client'; - -// Store event handlers registered via socket.on() -let eventHandlers: Record = {}; -let mockSocket: Partial & { connected: boolean }; - -function createMockSocket() { - eventHandlers = {}; - mockSocket = { - on: mock((event: string, handler: Function) => { - if (!eventHandlers[event]) eventHandlers[event] = []; - eventHandlers[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 any, - }; - return mockSocket as Socket; -} - -// Helper to emit socket events in tests -function emitSocketEvent(event: string, ...args: any[]) { - eventHandlers[event]?.forEach(handler => handler(...args)); +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. + * + * 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. + */ + +const sio = installSocketIOMock(); + +// 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[]; } -// Mock socket.io-client module -mock.module('socket.io-client', () => ({ - io: () => createMockSocket(), -})); +const HARNESSES: Array<[string, () => Harness]> = [ + [ + 'SocketIOTransport', + () => { + const client = new UseAIClient(new SocketIOTransport('http://localhost:8081')); + client.connect(); + const socket = sio.socket; -// Import after mocking -const { UseAIClient } = await import('./client'); + return { + client, + async open() { + socket.connected = true; + sio.fire('connect'); + }, + close(reason: string) { + socket.connected = false; + sio.fire('disconnect', reason); + }, + deliver(name, data) { + sio.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', + () => { + FakeWebSocket.reset(); + const client = new UseAIClient( + new WebSocketTransport('wss://localhost:8081', { + reconnectionDelay: 1, + reconnectionDelayMax: 1, + WebSocket: FakeWebSocketConstructor, + }), + ); + client.connect(); + // partysocket opens the socket asynchronously, and opens a fresh one after a drop. + const liveSocket = async () => { + await waitUntil(() => FakeWebSocket.latest !== undefined && FakeWebSocket.latest.readyState < 2); + return FakeWebSocket.latest!; + }; + const sent: string[] = []; + + return { + client, + async open() { + 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; + FakeWebSocket.latest?.serverClose(); + }, + deliver(name, data) { + const frame = name === 'event' ? data : { type: 'CUSTOM', name, value: data }; + FakeWebSocket.latest?.serverSend(JSON.stringify(frame)); + }, + sent() { + return 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 = () => { + const all = harness.sent(); + return all[all.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 +221,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)); - mockSocket.connected = true; - emitSocketEvent('connect'); + await harness.open(); + harness.deliver('config', { langfuseEnabled: true }); - mockSocket.connected = false; - emitSocketEvent('disconnect', 'transport close'); + expect(received).toEqual([false, true]); + }); - expect(client.isConnected()).toBe(false); + test('submitFeedback sends feedback once Langfuse is enabled', async () => { + await harness.open(); + harness.deliver('config', { langfuseEnabled: true }); + + 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 +383,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 +404,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 +421,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 +444,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 +538,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 +567,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 +596,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 +625,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 +643,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 +687,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 +706,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(sio.socket).toBeDefined(); + expect(client.isConnected()).toBe(false); + + sio.socket.connected = true; + expect(client.isConnected()).toBe(true); + + client.disconnect(); + }); + + test('disconnect() unsubscribes from the transport', async () => { + FakeWebSocket.reset(); + const client = new UseAIClient( + new WebSocketTransport('wss://localhost:8081', { WebSocket: FakeWebSocketConstructor }), + ); + const stateChanges: boolean[] = []; + client.onConnectionStateChange(connected => stateChanges.push(connected)); + client.connect(); + await waitUntil(() => FakeWebSocket.latest !== undefined); + 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.serverClose(); + + expect(stateChanges).toEqual([false, true]); + }); +}); diff --git a/packages/client/src/client.ts b/packages/client/src/client.ts index c8f1fa002..4bea04498 100644 --- a/packages/client/src/client.ts +++ b/packages/client/src/client.ts @@ -1,5 +1,4 @@ -import { io, Socket } from 'socket.io-client'; -import { EventType } from '@meetsmore-oss/use-ai-core'; +import { EventType, type CustomEvent } from '@meetsmore-oss/use-ai-core'; import type { ToolDefinition, Message, @@ -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,85 +107,68 @@ 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')); + * ``` */ - 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.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.transportUnsubscribes.push( + this.transport.onConnectionChange((connected, reason) => { + console.log(connected ? '[UseAI] Connected to server' : '[UseAI] Disconnected:', reason ?? ''); + this.connectionStateHandlers.forEach(handler => handler(connected)); + }), + this.transport.onEvent((event) => { + try { + console.log('[Client] Received event:', event.type); + this.handleEvent(event); + } catch (error) { + console.error('[UseAI] Error handling event:', error); + } + }), + ); - // 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.connect(); + } - // 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 + /** + * 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)); - }); - - 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)); - }); + 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 @@ -869,21 +848,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 +870,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 +897,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 6caed550d..c4cd079da 100644 --- a/packages/client/src/index.ts +++ b/packages/client/src/index.ts @@ -2,6 +2,8 @@ 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, WebSocketTransportOptions } from './transport'; export { defineTool, executeDefinedTool, convertToolsToDefinitions } from './defineTool'; /** @hidden */ export { z } from 'zod'; @@ -73,6 +75,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 b7b747395..1f9f77d2d 100644 --- a/packages/client/src/providers/useAIProvider.tsx +++ b/packages/client/src/providers/useAIProvider.tsx @@ -122,7 +122,7 @@ export interface PromptsContextValue { * Contains connection state and methods for managing tools and prompts. */ export interface UseAIContextValue { - /** The WebSocket URL of the UseAI server */ + /** 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; @@ -260,7 +260,9 @@ 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; CustomButton?: React.ComponentType | null; @@ -471,6 +473,7 @@ const DEFAULT_FILE_UPLOAD_CONFIG: FileUploadConfig = { */ export function UseAIProvider({ serverUrl, + transport, children, systemPrompt, CustomButton, @@ -573,16 +576,20 @@ export function UseAIProvider({ const handleDisconnectRef = useRef(serverEvents.handleDisconnect); handleDisconnectRef.current = serverEvents.handleDisconnect; + const resolvedServerUrl = transport?.url ?? serverUrl!; + useEffect(() => { - console.log('[UseAIProvider] Initializing client with serverUrl:', serverUrl); - const client = new UseAIClient(serverUrl); + 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); 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(); @@ -735,7 +742,7 @@ export function UseAIProvider({ // ── Context Values ────────────────────────────────────────────────────── const value: UseAIContextValue = { - 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 new file mode 100644 index 000000000..3ac94a733 --- /dev/null +++ b/packages/client/src/transport/SocketIOTransport.test.ts @@ -0,0 +1,134 @@ +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. +const sio = installSocketIOMock(); +const { SocketIOTransport } = await import('./SocketIOTransport'); + +describe('SocketIOTransport', () => { + let consoleLogSpy: ReturnType; + let consoleWarnSpy: ReturnType; + + beforeEach(() => { + 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(sio.ioOptions).toMatchObject({ + transports: ['polling', 'websocket'], + reconnection: true, + reconnectionAttempts: Infinity, + reconnectionDelay: 1000, + reconnectionDelayMax: 10_000, + withCredentials: true, + }); + }); + + 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(changes).toEqual([[true, undefined], [false, 'transport close']]); + }); + + test('delivers AG-UI events, and presents agents and config as CUSTOM events', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + const events: unknown[] = []; + transport.onEvent(event => events.push(event)); + + transport.connect(); + 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', () => { + new SocketIOTransport('http://localhost:8081').connect(); + + 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 report a connection change', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + 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'); + + expect(changes).toEqual([true]); + expect(consoleWarnSpy).toHaveBeenCalledTimes(2); + }); + + test('logs transport upgrades', () => { + new SocketIOTransport('http://localhost:8081').connect(); + sio.fire('connect'); + + expect(consoleLogSpy).toHaveBeenCalledWith('[UseAI] Transport:', 'polling'); + + sio.fire('engine:upgrade', { name: 'websocket' }); + expect(consoleLogSpy).toHaveBeenCalledWith('[UseAI] Upgraded to transport:', 'websocket'); + + 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(); + sio.socket.connected = true; + + transport.send({ 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', () => { + const transport = new SocketIOTransport('http://localhost:8081'); + expect(transport.connected).toBe(false); + + transport.connect(); + sio.socket.connected = true; + expect(transport.connected).toBe(true); + + 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(); + sio.socket.connected = true; + + transport.disconnect(); + + 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 new file mode 100644 index 000000000..8fef24dec --- /dev/null +++ b/packages/client/src/transport/SocketIOTransport.ts @@ -0,0 +1,114 @@ +import { io, Socket } from 'socket.io-client'; +import { EventType, type CustomEvent } from '@meetsmore-oss/use-ai-core'; +import type { AGUIEvent, UseAIClientMessage } from '../types'; +import type { UseAITransport } from './types'; + +// 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 a `serverUrl`, + * and what the bundled `@meetsmore-oss/use-ai-server` serves by default. + * + * @example + * ```typescript + * const transport = new SocketIOTransport('wss://your-server.com'); + * ``` + */ +export class SocketIOTransport implements UseAITransport { + private socket: Socket | null = null; + private eventHandlers = new Set<(event: AGUIEvent) => void>(); + private connectionHandlers = new Set<(connected: boolean, reason?: string) => void>(); + + /** + * @param url - The URL of the UseAI server + * @example + * ```typescript + * new SocketIOTransport('ws://localhost:8081'); + * ``` + */ + constructor(readonly url: string) {} + + get connected(): boolean { + return this.socket !== null && this.socket.connected; + } + + connect(): void { + const socket = io(this.url, { + transports: ['polling', 'websocket'], + reconnection: true, + reconnectionAttempts: Infinity, + reconnectionDelay: RECONNECTION_DELAY, + reconnectionDelayMax: RECONNECTION_DELAY_MAX, + 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.connectionHandlers.forEach(handler => handler(true)); + }); + + 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 + console.warn('[UseAI] Connection error:', error.message); + }); + + socket.on('disconnect', (reason: string) => { + this.connectionHandlers.forEach(handler => handler(false, reason)); + }); + } + + disconnect(): void { + this.socket?.disconnect(); + this.socket = null; + } + + send(message: UseAIClientMessage): void { + this.socket?.emit('message', message); + } + + 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 new file mode 100644 index 000000000..3e89b8bec --- /dev/null +++ b/packages/client/src/transport/WebSocketTransport.test.ts @@ -0,0 +1,221 @@ +import { describe, test, expect, beforeEach, afterEach, spyOn } from 'bun:test'; +import { WebSocketTransport } from './WebSocketTransport'; +import { FakeWebSocket, FakeWebSocketConstructor, waitUntil } from '../../test/fakeWebSocket'; + +function makeTransport(options: { reconnectionDelay?: number; reconnectionDelayMax?: number } = {}) { + return new WebSocketTransport('wss://server.example', { ...options, WebSocket: FakeWebSocketConstructor }); +} + +/** 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!; +} + +describe('WebSocketTransport', () => { + let consoleLogSpy: ReturnType; + let consoleWarnSpy: ReturnType; + + beforeEach(() => { + FakeWebSocket.reset(); + consoleLogSpy = spyOn(console, 'log').mockImplementation(() => {}); + consoleWarnSpy = spyOn(console, 'warn').mockImplementation(() => {}); + }); + + afterEach(() => { + consoleLogSpy.mockRestore(); + consoleWarnSpy.mockRestore(); + }); + + 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('delivers each JSON frame as one AG-UI event', async () => { + const transport = makeTransport(); + const events: unknown[] = []; + 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' }, + { type: 'CUSTOM', name: 'agents', value: { agents: [], defaultAgent: 'claude' } }, + ]); + + transport.disconnect(); + }); + + test('delivers a frame to every subscriber', async () => { + const transport = makeTransport(); + const first: unknown[] = []; + const second: unknown[] = []; + transport.onEvent(event => first.push(event)); + transport.onEvent(event => second.push(event)); + + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(JSON.stringify({ type: 'RUN_FINISHED' })); + + expect(first).toEqual([{ type: 'RUN_FINISHED' }]); + expect(second).toEqual([{ type: 'RUN_FINISHED' }]); + + transport.disconnect(); + }); + + test('unsubscribing stops delivery', async () => { + const transport = makeTransport(); + const events: unknown[] = []; + const unsubscribe = transport.onEvent(event => events.push(event)); + + const socket = await connect(transport); + socket.serverOpen(); + socket.serverSend(JSON.stringify({ type: 'STEP_STARTED', stepName: '1' })); + unsubscribe(); + socket.serverSend(JSON.stringify({ type: 'STEP_STARTED', stepName: '2' })); + + expect(events).toEqual([{ type: 'STEP_STARTED', stepName: '1' }]); + + transport.disconnect(); + }); + + test('a malformed or non-text frame is ignored, not an error', async () => { + const transport = makeTransport(); + const events: unknown[] = []; + transport.onEvent(event => events.push(event)); + + const socket = await connect(transport); + socket.serverOpen(); + + expect(() => socket.serverSend('not json')).not.toThrow(); + expect(() => socket.serverSend(JSON.stringify({ noTypeHere: true }))).not.toThrow(); + expect(() => socket.serverSend(new ArrayBuffer(4))).not.toThrow(); + + 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', async () => { + const transport = makeTransport(); + const socket = await connect(transport); + socket.serverOpen(); + + transport.send({ 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', async () => { + const transport = makeTransport({ reconnectionDelay: 10_000 }); + expect(transport.connected).toBe(false); + + const socket = await connect(transport); + expect(transport.connected).toBe(false); + + socket.serverOpen(); + expect(transport.connected).toBe(true); + + socket.serverClose(); + expect(transport.connected).toBe(false); + + transport.disconnect(); + }); + + test('reports a connection change on open and on close', async () => { + const transport = makeTransport({ reconnectionDelay: 10_000 }); + const changes: boolean[] = []; + transport.onConnectionChange(connected => changes.push(connected)); + + const socket = await connect(transport); + socket.serverOpen(); + socket.serverClose(); + + expect(changes).toEqual([true, false]); + + transport.disconnect(); + }); + + test('a failed connection attempt does not report a disconnection', async () => { + const transport = makeTransport({ reconnectionDelay: 10_000 }); + 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(changes).toEqual([]); + + transport.disconnect(); + }); + + test('reconnects after the socket drops', async () => { + const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); + const socket = await connect(transport); + socket.serverOpen(); + + socket.serverClose(); + await waitUntil(() => FakeWebSocket.instances.length > 1); + + FakeWebSocket.latest!.serverOpen(); + expect(transport.connected).toBe(true); + + transport.disconnect(); + }); + + test('keeps retrying while attempts fail', async () => { + const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); + await connect(transport); + + for (let i = 0; i < 3; i++) { + const count = FakeWebSocket.instances.length; + FakeWebSocket.latest!.serverClose(); + await waitUntil(() => FakeWebSocket.instances.length > count); + } + + expect(FakeWebSocket.instances.length).toBeGreaterThanOrEqual(4); + + transport.disconnect(); + }); + + test('disconnect() stops the retries', async () => { + const transport = makeTransport({ reconnectionDelay: 1, reconnectionDelayMax: 2 }); + const socket = await connect(transport); + socket.serverOpen(); + + socket.serverClose(); + transport.disconnect(); + const openedByNow = FakeWebSocket.instances.length; + + await new Promise(resolve => setTimeout(resolve, 20)); + + expect(FakeWebSocket.instances).toHaveLength(openedByNow); + expect(transport.connected).toBe(false); + }); + + test('disconnect() closes the open socket', async () => { + const transport = makeTransport(); + const socket = await connect(transport); + socket.serverOpen(); + + 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 000000000..f31053bd3 --- /dev/null +++ b/packages/client/src/transport/WebSocketTransport.ts @@ -0,0 +1,147 @@ +import ReconnectingWebSocket from 'partysocket/ws'; +import type { AGUIEvent, UseAIClientMessage } from '../types'; +import type { UseAITransport } from './types'; + +/** + * Options for {@link WebSocketTransport}. + */ +export interface WebSocketTransportOptions { + /** + * Delay before the first reconnection attempt, in milliseconds. + * Each later attempt doubles the delay, up to {@link reconnectionDelayMax}. + * + * @default 1000 + */ + reconnectionDelay?: number; + /** + * Upper bound on the delay between reconnection attempts, in milliseconds. + * + * @default 10000 + */ + reconnectionDelayMax?: number; + /** + * WebSocket constructor to open the connection with. + * Supply one on a runtime without a global `WebSocket`, or in a test. + * + * @default globalThis.WebSocket + */ + WebSocket?: typeof WebSocket; +} + +/** + * 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 client + * ignores an event type it does not handle. + * + * Reconnection is automatic: indefinite, with exponential backoff capped at + * `reconnectionDelayMax`. The defaults match {@link SocketIOTransport}. + * + * @example + * ```tsx + * + * ``` + */ +export class WebSocketTransport implements UseAITransport { + private socket: ReconnectingWebSocket | null = null; + private _connected = false; + private eventHandlers = new Set<(event: AGUIEvent) => void>(); + private connectionHandlers = new Set<(connected: boolean, reason?: string) => void>(); + + /** + * @param url - WebSocket URL of the server + * @example + * ```typescript + * new WebSocketTransport('wss://your-server.com'); + * ``` + */ + constructor(readonly url: string, private readonly options: WebSocketTransportOptions = {}) {} + + get connected(): boolean { + return this._connected; + } + + connect(): void { + if (this.socket) return; + + const socket = new ReconnectingWebSocket(this.url, [], { + WebSocket: this.options.WebSocket, + 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. + maxEnqueuedMessages: 0, + }); + this.socket = socket; + + socket.onopen = () => { + this._connected = true; + this.connectionHandlers.forEach(handler => handler(true)); + }; + + socket.onmessage = (event) => { + const frame = parseFrame(event.data); + if (frame) this.eventHandlers.forEach(handler => handler(frame)); + }; + + socket.onerror = (event) => { + // Use warn instead of error to avoid triggering Next.js error overlay + console.warn('[UseAI] Connection error:', event.message); + }; + + socket.onclose = () => { + // A close also fires for a failed attempt. Only a socket that opened reports a disconnection. + if (!this._connected) return; + this._connected = false; + this.connectionHandlers.forEach(handler => handler(false, 'transport close')); + }; + } + + 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(); + } + } + + send(message: UseAIClientMessage): void { + this.socket?.send(JSON.stringify(message)); + } + + 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); + }; + } +} + +function parseFrame(data: unknown): AGUIEvent | 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 || typeof (parsed as { type?: unknown }).type !== 'string') { + return null; + } + return parsed as AGUIEvent; +} diff --git a/packages/client/src/transport/index.ts b/packages/client/src/transport/index.ts new file mode 100644 index 000000000..cf685bdaf --- /dev/null +++ b/packages/client/src/transport/index.ts @@ -0,0 +1,4 @@ +export type { UseAITransport } from './types'; +export { SocketIOTransport } 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 new file mode 100644 index 000000000..3be0141ab --- /dev/null +++ b/packages/client/src/transport/types.ts @@ -0,0 +1,45 @@ +import type { AGUIEvent, UseAIClientMessage } from '../types'; + +/** + * The pipe between {@link UseAIClient} and a server. + * + * 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`. + * + * 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 { + /** + * The server this transport connects to. Reported as `serverUrl` on the provider context. + * @example 'wss://your-server.com' + */ + 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; + + /** 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 downstream AG-UI events. + * @returns Cleanup function to unsubscribe + */ + onEvent(handler: (event: AGUIEvent) => void): () => void; + + /** + * 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/src/types.ts b/packages/client/src/types.ts index 71d7c3100..8f74501a6 100644 --- a/packages/client/src/types.ts +++ b/packages/client/src/types.ts @@ -1,10 +1,33 @@ +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 */ - 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. The context reports the + * transport's `url` as `serverUrl`. + */ + transport: UseAITransport; + serverUrl?: never; + }; /** * Toggles for optional chat UI features. Opt-out features default to enabled diff --git a/packages/client/test/fakeWebSocket.ts b/packages/client/test/fakeWebSocket.ts new file mode 100644 index 000000000..53163f0a2 --- /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 000000000..72196bedb --- /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/package.json b/packages/server/package.json index f2ad6d07c..78cd4ce4f 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 79c0adf74..dc96ef33f 100644 --- a/packages/server/src/agents/types.ts +++ b/packages/server/src/agents/types.ts @@ -1,6 +1,24 @@ -import type { Socket } from 'socket.io'; import type { ModelMessage } from 'ai'; -import type { ToolDefinition, AGUIEvent, ToolApprovalRequestEvent } from '../types'; +import type { ToolDefinition, AGUIEvent, ToolApprovalRequestEvent, UseAIClientMessage } from '../types'; + +/** + * 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. `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; +} /** * Context for a single client session. @@ -11,8 +29,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 314cad831..8dec0d9d8 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 2874db438..ae359dcdb 100644 --- a/packages/server/src/runtime/bun/BunRuntimeAdapter.ts +++ b/packages/server/src/runtime/bun/BunRuntimeAdapter.ts @@ -1,41 +1,38 @@ 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, WebSocketListener } from '../types'; import { resolveCorsHeaders, resolvePreflightHeaders } from './cors'; +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. - * 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, - }); + createServer(listener: RuntimeListener, config: RuntimeServerConfig): RuntimeServerHandle { + const handlers = + listener.transport === 'socketio' + ? this.socketIOHandlers(listener.io, config) + : this.webSocketHandlers(listener.onConnection); - // 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); - } - } - }); - - // Bind Socket.IO to Bun engine - io.bind(this.engine); - - const handler = this.engine.handleRequest.bind(this.engine); - const websocketHandler = this.engine.handler().websocket; - - // Start Bun server const bunServer = Bun.serve({ port: config.port, idleTimeout: config.idleTimeout ?? 30, @@ -53,7 +50,6 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }); } - // Helper to create response with CORS headers const corsHeaders = resolveCorsHeaders(requestOrigin, config.cors); // Health check endpoint @@ -63,28 +59,20 @@ export class BunRuntimeAdapter implements RuntimeAdapter { }); } - // 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 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: websocketHandler, + websocket: handlers.websocket, }); return { @@ -94,4 +82,56 @@ export class BunRuntimeAdapter implements RuntimeAdapter { server: bunServer, }; } + + private socketIOHandlers(io: SocketIOServer, config: RuntimeServerConfig): ListenerHandlers { + 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 { + path: '/socket.io/', + upgrade: (req, server) => handleRequest(req, server as never), + websocket: this.engine.handler().websocket as unknown as WebSocketHandler, + }; + } + + private webSocketHandlers(onConnection: WebSocketListener['onConnection']): ListenerHandlers { + const dataOf = (ws: { data: unknown }) => ws.data as BunWebSocketData; + return { + 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: { + 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 new file mode 100644 index 000000000..95931111d --- /dev/null +++ b/packages/server/src/runtime/bun/rawWebSocket.ts @@ -0,0 +1,54 @@ +import type { RawWebSocket } from '../types'; + +/** Per-connection data on `ws.data`: set at upgrade, completed on open. */ +export interface BunWebSocketData { + remoteAddress?: string; + connection?: BunRawWebSocket; +} + +interface BunWebSocket { + readonly readyState: number; + 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?.(); + } +} diff --git a/packages/server/src/runtime/clientIp.ts b/packages/server/src/runtime/clientIp.ts index 9d54833a2..796f6b87c 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 26bf7afa8..503b0ec59 100644 --- a/packages/server/src/runtime/index.ts +++ b/packages/server/src/runtime/index.ts @@ -7,9 +7,13 @@ export type { RuntimeType, RuntimeServerConfig, 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 7e4529843..59efffecf 100644 --- a/packages/server/src/runtime/node/NodeRuntimeAdapter.ts +++ b/packages/server/src/runtime/node/NodeRuntimeAdapter.ts @@ -1,19 +1,21 @@ import { createServer, type Server as HttpServer } from 'http'; import type { Server as SocketIOServer } from 'socket.io'; -import type { RuntimeAdapter, RuntimeServerConfig, RuntimeServerHandle } from '../types'; +import { WebSocketServer } from 'ws'; +import type { RuntimeAdapter, RuntimeListener, RuntimeServerConfig, RuntimeServerHandle, WebSocketListener } from '../types'; +import { forwardedClientIp } from '../clientIp'; +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}`); @@ -25,17 +27,29 @@ 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'); }); - // Attach Socket.IO to the HTTP server - // Socket.IO handles CORS internally for /socket.io/* paths + const wss = listener.transport === 'socketio' + ? this.attachSocketIO(listener.io, httpServer, config) + : this.attachWebSocket(listener.onConnection, httpServer, config); + + httpServer.listen(config.port); + + return { + stop: () => { + wss?.close(); + httpServer.close(); + }, + server: httpServer, + }; + } + + private attachSocketIO(io: SocketIOServer, httpServer: HttpServer, config: RuntimeServerConfig): null { io.attach(httpServer, { transports: ['polling', 'websocket'], maxHttpBufferSize: config.maxHttpBufferSize, @@ -50,25 +64,33 @@ 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; + } - // Start listening - httpServer.listen(config.port); - - return { - stop: () => { - httpServer.close(); - }, - server: httpServer, - }; + private attachWebSocket( + onConnection: WebSocketListener['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 remoteAddress = forwardedClientIp(req.headers['x-forwarded-for'], req.socket.remoteAddress); + onConnection(new NodeRawWebSocket(ws, remoteAddress)); + }); + }); + return wss; } } diff --git a/packages/server/src/runtime/node/rawWebSocket.ts b/packages/server/src/runtime/node/rawWebSocket.ts new file mode 100644 index 000000000..b2b1a617a --- /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) => { + if (!isBinary) handler(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 76f58c814..d694ea786 100644 --- a/packages/server/src/runtime/types.ts +++ b/packages/server/src/runtime/types.ts @@ -6,6 +6,40 @@ 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; +} + +/** 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. */ @@ -48,11 +82,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 f91691181..c2424bc05 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'; @@ -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,13 @@ import { type RuntimeAdapter, type RuntimeServerHandle, type ClientIpTracker, + type RuntimeListener, } from './runtime'; +import { SocketIOClientConnection } from './socketIOConnection'; +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. @@ -89,17 +92,20 @@ export type { ClientSession } 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> & { + // 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; idleTimeout: number; + transport: 'socketio' | 'websocket'; }; private rateLimiter: RateLimiter; private cleanupInterval: NodeJS.Timeout; @@ -131,6 +137,7 @@ export class UseAIServer { maxHttpBufferSize: config.maxHttpBufferSize ?? 20 * 1024 * 1024, // 20MB default cors: config.cors, idleTimeout: config.idleTimeout ?? 30, + transport: config.transport ?? 'socketio', }; // Set agents registry @@ -145,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, @@ -194,13 +209,22 @@ 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(); + let listener: RuntimeListener; + if (this.config.transport === 'socketio') { + this.io = new SocketIOServer({ + transports: ['polling', 'websocket'], + maxHttpBufferSize: this.config.maxHttpBufferSize, + }); + 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()) { logger.info('Rate limiting enabled', { @@ -209,8 +233,7 @@ export class UseAIServer { }); } - // Start server using runtime adapter - this.serverHandle = this.runtimeAdapter.createServer(this.io, { + this.serverHandle = this.runtimeAdapter.createServer(listener, { port: this.config.port, idleTimeout: this.config.idleTimeout, cors: this.config.cors, @@ -219,6 +242,7 @@ export class UseAIServer { this.clientIpTracker.trackPollingConnection(sessionId, ip); }, }); + logger.info('UseAI server ready', { port: this.config.port, transport: this.config.transport }); } /** @@ -283,111 +307,99 @@ export class UseAIServer { logger.debug('Registered message handler', { type }); } - 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: - // 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 transport = conn.transport.name; - logger.info('Client connected', { clientId, threadId, ipAddress, transport }); - - // Log transport upgrades - socket.conn.on('upgrade', (transport) => { - logger.info('Client upgraded transport', { clientId, transport: transport.name }); - }); - - const session: ClientSession = { - clientId, - ipAddress, - socket, - threadId, - tools: [], - state: null, - pendingToolCalls: new Map(), - pendingToolApprovals: new Map(), - }; + private acceptConnection(connection: ClientConnection) { + const session = this.createSession(connection); + logger.info('Client connected', { + clientId: session.clientId, + threadId: session.threadId, + ipAddress: session.ipAddress, + }); - this.clients.set(socket.id, session); + connection.onMessage((message) => { + void this.receiveClientMessage(session, message); + }); - // Send available agents to client - const availableAgents = Object.entries(this.agents).map(([id, agent]) => ({ - id, - name: agent.getName?.() || id, - annotation: agent.getAnnotation?.(), - })); - socket.emit('agents', { - agents: availableAgents, - defaultAgent: this.defaultAgentId, - }); + connection.onClose(() => { + logger.info('Client disconnected', { clientId: session.clientId, ipAddress: session.ipAddress }); + this.destroySession(session); + }); + } - // Call plugin lifecycle hooks - for (const plugin of this.plugins) { - plugin.onClientConnect?.(session); - } + /** + * Creates the session for a newly accepted connection and announces the server's + * agents to it. + */ + private createSession(connection: ClientConnection): ClientSession { + const session: ClientSession = { + clientId: `client-${++this.clientIdCounter}`, + ipAddress: connection.ipAddress, + socket: connection, + threadId: uuidv4(), + tools: [], + state: null, + pendingToolCalls: new Map(), + pendingToolApprovals: new Map(), + }; - 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(), - }); - } - }); + this.clients.set(connection.id, session); - socket.on('disconnect', () => { - logger.info('Client disconnected', { clientId, ipAddress }); + connection.emit('agents', this.agentsPayload); - // Abort any pending tool calls/approvals for this session - abortRun(session.abortController, new RunAbortedByClientDisconnect()); + for (const plugin of this.plugins) { + plugin.onClientConnect?.(session); + } - // Clean up polling IP entry - this.clientIpTracker.removePollingConnection(conn.id); + return session; + } - // Call plugin lifecycle hooks - for (const plugin of this.plugins) { - plugin.onClientDisconnect?.(session); - } + /** + * 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()); - // Note: Rate limiting persists by IP address across connections - this.clients.delete(socket.id); - }); - }); + for (const plugin of this.plugins) { + plugin.onClientDisconnect?.(session); + } - logger.info('UseAI server ready', { port: this.config.port }); + this.clients.delete(session.socket.id); } - private async handleClientMessage(socket: Socket, message: UseAIClientMessage) { - const session = this.clients.get(socket.id); - if (!session) return; + /** + * 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, message); + } catch (error) { + logger.error('Error handling message', { + error: error instanceof Error ? error.message : 'Unknown error', + clientId: session.clientId, + }); + 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(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) { @@ -930,10 +942,8 @@ 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) { + connection.emit('event', event); } /** @@ -1108,7 +1118,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/socketIOConnection.ts b/packages/server/src/socketIOConnection.ts new file mode 100644 index 000000000..372e25360 --- /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/types.ts b/packages/server/src/types.ts index a7bc5f750..55765bce9 100644 --- a/packages/server/src/types.ts +++ b/packages/server/src/types.ts @@ -134,6 +134,14 @@ export interface UseAIServerConfig 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 new file mode 100644 index 000000000..a3769921b --- /dev/null +++ b/packages/server/src/websocket-transport.integration.test.ts @@ -0,0 +1,177 @@ +import { describe, expect, test, beforeAll, afterAll } from 'bun:test'; +import { UseAIClient, WebSocketTransport } from '@meetsmore-oss/use-ai-client'; +import { UseAIServer } from './server'; +import { EventType } from '@meetsmore-oss/use-ai-core'; +import type { AGUIEvent } from './types'; +import { + createSequentialMockModel, + TestCleanupManager, +} from '../test/integration-test-utils'; +import { waitFor } from '../test/test-utils'; +import { AISDKAgent } from './agents/AISDKAgent'; + +/** + * Drives the bundled WebSocketTransport against a real server (transport: 'websocket') + * through a full run: prompt → tool call → tool result → RUN_FINISHED. + * + * Nothing is stubbed on either side, so this is what proves the documented framing + * in docs/websocket-protocol.md is implementable. + */ + +const RUNTIMES: ('bun' | 'node')[] = ['bun', 'node']; + +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)); + + 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(); + }); +} + +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, + transport: 'websocket', + 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('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 () => { + 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('transport defaults to socketio', () => { + const cleanup = new TestCleanupManager(); + + afterAll(() => { + cleanup.cleanup(); + }); + + 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, + agents: { 'test-agent': new AISDKAgent({ hooks: { loadConfig: () => ({ model }) } }) }, + defaultAgent: 'test-agent', + }), + ); + + const socket = await cleanup.createTestClient(port); + expect(socket.connected).toBe(true); + socket.disconnect(); + + // 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 da3c2f9c8..f4e62f257 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 */