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
*/