From 47bfe02a87dc3703a9b162bbe27b6e6184baccd6 Mon Sep 17 00:00:00 2001 From: Neko Date: Thu, 18 Jun 2026 15:11:35 +0800 Subject: [PATCH] feat(better-ws): added new package (#1989) --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by-agent: Codex --- apps/stage-pocket/src/App.vue | 4 +- .../src/modules/server-channel-qr-probe.ts | 14 +- .../src/modules/websocket-bridge.ts | 114 +- packages/better-ws/README.md | 141 + packages/better-ws/package.json | 55 + .../src/client/crossws/index.test.ts | 118 + .../better-ws/src/client/crossws/index.ts | 132 + packages/better-ws/src/client/index.ts | 917 ++++++ packages/better-ws/src/index.test.ts | 2489 +++++++++++++++++ packages/better-ws/src/index.ts | 2 + .../better-ws/src/server/h3/index.test.ts | 55 + packages/better-ws/src/server/h3/index.ts | 42 + packages/better-ws/src/server/index.ts | 621 ++++ .../better-ws/src/server/liveness.test.ts | 190 ++ packages/better-ws/src/server/peers.ts | 341 +++ packages/better-ws/src/shared/index.ts | 80 + .../src/shared/utils/event-wait-for.test.ts | 52 + .../src/shared/utils/event-wait-for.ts | 185 ++ packages/better-ws/src/shared/utils/index.ts | 1 + packages/better-ws/tsconfig.json | 22 + packages/better-ws/tsdown.config.ts | 11 + packages/better-ws/vitest.config.ts | 7 + packages/server-runtime/package.json | 4 +- packages/server-runtime/src/index.test.ts | 8 +- packages/server-runtime/src/index.ts | 1183 ++++---- .../airi/{index.test.ts => codec.test.ts} | 83 +- .../src/server-ws/airi/codec.ts | 81 + .../index.test.ts => airi/consumers.test.ts} | 12 +- .../src/server-ws/airi/consumers.ts | 301 ++ .../src/server-ws/airi/index.ts | 381 +-- .../src/server-ws/airi/liveness.test.ts | 21 + .../src/server-ws/airi/liveness.ts | 13 + .../src/server-ws/airi/responses.test.ts | 70 + .../src/server-ws/airi/responses.ts | 76 + .../src/server-ws/airi/routing.ts | 67 + .../src/server-ws/core/index.ts | 543 ---- packages/server-runtime/src/server.test.ts | 10 +- packages/server-runtime/src/server/index.ts | 10 +- .../src/setupApp.liveness.test.ts | 235 ++ packages/server-runtime/src/types/conn.ts | 5 + packages/server-sdk/package.json | 5 +- packages/server-sdk/src/client.ts | 1064 +++---- packages/server-sdk/src/codec.ts | 81 + packages/server-sdk/src/extension-peer.ts | 129 +- packages/server-sdk/src/index.ts | 2 +- packages/server-sdk/src/websocket-like.ts | 30 - packages/server-sdk/test/client.test.ts | 723 +++-- packages/server-sdk/test/codec.test.ts | 71 + .../server-sdk/test/extension-peer.test.ts | 103 +- .../src/stores/mods/api/channel-server.ts | 23 +- pnpm-lock.yaml | 30 +- vitest.config.ts | 1 + 52 files changed, 8112 insertions(+), 2846 deletions(-) create mode 100644 packages/better-ws/README.md create mode 100644 packages/better-ws/package.json create mode 100644 packages/better-ws/src/client/crossws/index.test.ts create mode 100644 packages/better-ws/src/client/crossws/index.ts create mode 100644 packages/better-ws/src/client/index.ts create mode 100644 packages/better-ws/src/index.test.ts create mode 100644 packages/better-ws/src/index.ts create mode 100644 packages/better-ws/src/server/h3/index.test.ts create mode 100644 packages/better-ws/src/server/h3/index.ts create mode 100644 packages/better-ws/src/server/index.ts create mode 100644 packages/better-ws/src/server/liveness.test.ts create mode 100644 packages/better-ws/src/server/peers.ts create mode 100644 packages/better-ws/src/shared/index.ts create mode 100644 packages/better-ws/src/shared/utils/event-wait-for.test.ts create mode 100644 packages/better-ws/src/shared/utils/event-wait-for.ts create mode 100644 packages/better-ws/src/shared/utils/index.ts create mode 100644 packages/better-ws/tsconfig.json create mode 100644 packages/better-ws/tsdown.config.ts create mode 100644 packages/better-ws/vitest.config.ts rename packages/server-runtime/src/server-ws/airi/{index.test.ts => codec.test.ts} (50%) create mode 100644 packages/server-runtime/src/server-ws/airi/codec.ts rename packages/server-runtime/src/server-ws/{core/index.test.ts => airi/consumers.test.ts} (96%) create mode 100644 packages/server-runtime/src/server-ws/airi/consumers.ts create mode 100644 packages/server-runtime/src/server-ws/airi/liveness.test.ts create mode 100644 packages/server-runtime/src/server-ws/airi/liveness.ts create mode 100644 packages/server-runtime/src/server-ws/airi/responses.test.ts create mode 100644 packages/server-runtime/src/server-ws/airi/responses.ts create mode 100644 packages/server-runtime/src/server-ws/airi/routing.ts delete mode 100644 packages/server-runtime/src/server-ws/core/index.ts create mode 100644 packages/server-runtime/src/setupApp.liveness.test.ts create mode 100644 packages/server-sdk/src/codec.ts delete mode 100644 packages/server-sdk/src/websocket-like.ts create mode 100644 packages/server-sdk/test/codec.test.ts diff --git a/apps/stage-pocket/src/App.vue b/apps/stage-pocket/src/App.vue index dd57d30e8..10498f4ba 100644 --- a/apps/stage-pocket/src/App.vue +++ b/apps/stage-pocket/src/App.vue @@ -18,7 +18,7 @@ import { toast, Toaster } from 'vue-sonner' import OnboardingPermissionsStep from './components/onboarding/step-permissions.vue' -import { getHostWebSocketConstructor } from './modules/websocket-bridge' +import { getHostWebSocketConnector } from './modules/websocket-bridge' const contextBridgeStore = useContextBridgeStore() const i18n = useI18n() @@ -80,7 +80,7 @@ onMounted(async () => { await serverChannelStore.initialize({ possibleEvents: ['ui:configure'], - websocketConstructor: getHostWebSocketConstructor(), + connector: getHostWebSocketConnector, }).catch(err => console.error('Failed to initialize Mods Server Channel in App.vue:', err)) contextBridgeStore.initialize() characterOrchestratorStore.initialize() diff --git a/apps/stage-pocket/src/modules/server-channel-qr-probe.ts b/apps/stage-pocket/src/modules/server-channel-qr-probe.ts index 57b8693bf..e78da1804 100644 --- a/apps/stage-pocket/src/modules/server-channel-qr-probe.ts +++ b/apps/stage-pocket/src/modules/server-channel-qr-probe.ts @@ -1,19 +1,23 @@ import type { ServerChannelQrPayload } from '@proj-airi/stage-shared/server-channel-qr' import { errorMessageFrom } from '@moeru/std' -import { Client, WebSocketEventSource } from '@proj-airi/server-sdk' +import { Client, createTextProtocolConnector, WebSocketEventSource } from '@proj-airi/server-sdk' -import { getHostWebSocketConstructor } from './websocket-bridge' +import { getHostWebSocketConnector } from './websocket-bridge' export async function probeServerChannelQrPayload(payload: ServerChannelQrPayload) { - const websocketConstructor = getHostWebSocketConstructor() - if (!websocketConstructor) { + if (!payload.urls.some(url => getHostWebSocketConnector(url))) { throw new Error('AIRI host websocket bridge is unavailable') } const errors: string[] = [] for (const url of payload.urls) { + const connector = getHostWebSocketConnector(url) + if (!connector) { + throw new Error('AIRI host websocket bridge is unavailable') + } + const client = new Client({ autoConnect: false, autoReconnect: false, @@ -21,7 +25,7 @@ export async function probeServerChannelQrPayload(payload: ServerChannelQrPayloa name: WebSocketEventSource.StageWeb, token: payload.authToken, url, - websocketConstructor, + connector: createTextProtocolConnector(connector), }) try { diff --git a/apps/stage-pocket/src/modules/websocket-bridge.ts b/apps/stage-pocket/src/modules/websocket-bridge.ts index 2a787d3e4..1590148df 100644 --- a/apps/stage-pocket/src/modules/websocket-bridge.ts +++ b/apps/stage-pocket/src/modules/websocket-bridge.ts @@ -1,9 +1,4 @@ -import type { - WebSocketErrorEventLike, - WebSocketLike, - WebSocketLikeConstructor, - WebSocketMessageEventLike, -} from '@proj-airi/server-sdk' +import type { ClientConnector, ClientEvents } from '@proj-airi/server-sdk' type HostBridgeCommand = | { kind: 'connect', id: string, url: string } @@ -34,7 +29,7 @@ declare global { } } -const sockets = new Map() +const connections = new Map() function postBridgeMessage(command: HostBridgeCommand) { if (window.AiriHostBridge) { @@ -52,43 +47,37 @@ function postBridgeMessage(command: HostBridgeCommand) { function dispatchNativeEvent(payload: string) { const event = JSON.parse(payload) as HostBridgeEvent - const socket = sockets.get(event.id) - if (!socket) { + const connection = connections.get(event.id) + if (!connection) { return } - socket.handleNativeEvent(event) + connection.handleNativeEvent(event) } -class HostWebSocket implements WebSocketLike { - static readonly CONNECTING = 0 - static readonly OPEN = 1 - static readonly CLOSING = 2 - static readonly CLOSED = 3 - +class HostBridgeConnection { readonly id = crypto.randomUUID() - readyState = HostWebSocket.CONNECTING - onopen?: (event?: unknown) => void - onmessage?: (event: WebSocketMessageEventLike) => void - onerror?: (event: WebSocketErrorEventLike | unknown) => void - onclose?: (event?: unknown) => void + private opened = false + private settled = false + + constructor( + private readonly url: string, + private readonly events: ClientEvents, + private readonly resolve: () => void, + private readonly reject: (error: Error) => void, + ) { + connections.set(this.id, this) - constructor(url: string) { - sockets.set(this.id, this) postBridgeMessage({ kind: 'connect', id: this.id, - url, + url: this.url, }) } - send(data: string | ArrayBufferLike | ArrayBufferView) { - if (typeof data !== 'string') { - throw new TypeError('HostWebSocket only supports text frames') - } - - if (this.readyState !== HostWebSocket.OPEN) { - throw new Error('WebSocket is not open') + send(data: string) { + if (!this.opened) { + return false } postBridgeMessage({ @@ -96,14 +85,15 @@ class HostWebSocket implements WebSocketLike { id: this.id, data, }) + + return true } close(code?: number, reason?: string) { - if (this.readyState === HostWebSocket.CLOSED) { + if (this.settled && !this.opened) { return } - this.readyState = HostWebSocket.CLOSING postBridgeMessage({ kind: 'close', id: this.id, @@ -115,33 +105,73 @@ class HostWebSocket implements WebSocketLike { handleNativeEvent(event: HostBridgeEvent) { switch (event.kind) { case 'open': - this.readyState = HostWebSocket.OPEN - this.onopen?.() + this.opened = true + this.settled = true + this.resolve() break case 'message': - this.onmessage?.({ data: event.data }) + this.events.message(event.data) break case 'error': - this.onerror?.({ error: new Error(event.message) }) + if (!this.settled) { + this.settled = true + connections.delete(this.id) + this.reject(new Error(event.message)) + return + } + + this.events.error(new Error(event.message)) break case 'close': - this.readyState = HostWebSocket.CLOSED - sockets.delete(this.id) - this.onclose?.({ code: event.code, reason: event.reason }) + connections.delete(this.id) + if (!this.settled) { + this.settled = true + this.reject(createCloseBeforeOpenError(event)) + return + } + + this.opened = false + this.events.close({ code: event.code, reason: event.reason }) break } } } -export function getHostWebSocketConstructor() { +function createCloseBeforeOpenError(event: Extract) { + const reason = event.reason ? ` ${event.reason}` : '' + const code = typeof event.code === 'number' ? ` with code ${event.code}` : '' + return new Error(`AIRI host websocket bridge closed before opening${code}.${reason}`) +} + +export function getHostWebSocketConnector(url: string): ClientConnector | undefined { if (!window.AiriHostBridge && !window.webkit?.messageHandlers?.airiHostBridge) { return undefined } window.__airiHostBridge = window.__airiHostBridge ?? {} window.__airiHostBridge.onNativeMessage = dispatchNativeEvent - return HostWebSocket as unknown as WebSocketLikeConstructor + + return { + connect(events) { + let connection: HostBridgeConnection | undefined + const opened = new Promise((resolve, reject) => { + connection = new HostBridgeConnection(url, events, resolve, reject) + }) + + return opened.then(() => { + const activeConnection = connection + if (!activeConnection) { + throw new Error('AIRI host websocket bridge connection was not created') + } + + return { + send: message => activeConnection.send(message), + close: (code?: number, reason?: string) => activeConnection.close(code, reason), + } + }) + }, + } } diff --git a/packages/better-ws/README.md b/packages/better-ws/README.md new file mode 100644 index 000000000..8819d6124 --- /dev/null +++ b/packages/better-ws/README.md @@ -0,0 +1,141 @@ +# @proj-airi/better-ws + +Runtime-agnostic WebSocket primitives for reliable realtime connections. + +## What it does + +- Provides `createClient(...)` and `createServer(...)` runtime primitives. +- Tracks client connection state and schedules reconnects after unexpected closes. +- Tracks server peers and supports peer send, broadcast, and named groups. +- Dispatches raw caller-owned messages without imposing an event, RPC, or extension protocol. + +## How to use + +```ts +import { createClient } from '@proj-airi/better-ws' +import { createServer } from '@proj-airi/better-ws/server' + +const client = createClient({ + url: 'ws://localhost:3000/ws', + // Reconnect is enabled by default; pass an object to customize the policy. + reconnect: { + retries: Number.POSITIVE_INFINITY, + delay: attempt => Math.min(1000 * 2 ** (attempt - 1), 30_000), + }, +}) + +await client.connect() +client.send('hello') + +const server = createServer() + +server.onMessage(({ server, message }) => { + server.broadcast(message) +}) + +const peer = server.accept({ + id: 'peer-1', + send(message) { + console.info('send to runtime', message) + return true + }, +}) + +peer.receive('hello') +``` + +Server peer liveness is opt-in and scheduler-driven. It tracks inbound traffic +without imposing any protocol event shape. `checkLiveness()` evaluates +`peers.unhealthyTimeout` and `peers.closeTimeout`; server-side heartbeat +`mode`, `interval`, `message`, and `isResponse` are reserved for app or +adapter driven scheduling and do not automatically send pings or classify +responses. + +```ts +const server = createServer({ + peers: { + unhealthyTimeout: 60_000, + closeTimeout: 120_000, + }, + heartbeat: { + timeout: 60_000, + }, +}) + +server.onPeerHealthChange(({ peer, healthy, silentFor }) => { + console.info('peer health changed', peer.id, healthy, silentFor) +}) + +setInterval(() => { + server.checkLiveness() +}, 30_000) +``` + +Heartbeat is opt-in. This example uses an application-level message heartbeat. +Message heartbeat mode requires `message` so the client knows what to send: + +```ts +const client = createClient({ + url: 'ws://localhost:3000/ws', + heartbeat: { + mode: 'message', + message: 'ping', + isResponse: message => message === 'pong', + interval: 30_000, + timeout: 10_000, + }, +}) +``` + +Native ping heartbeat is used only when a connector exposes `ping()`. In `auto` +mode, the client uses native `ping()` when available and falls back to message +heartbeat only when `message` is provided. The built-in browser `WebSocket` +adapter does not expose native ping frames. + +When `isResponse` is provided, heartbeat uses strict response matching: only a +matching inbound message clears the pending heartbeat timeout. When `isResponse` +is omitted, any inbound message is treated as liveness and clears the pending +timeout. + +For non-text messages or non-native runtimes, pass a connector. The connector owns parsing and serialization: + +```ts +const client = createClient({ + connector: { + async connect(events) { + const ws = new WebSocket('ws://localhost:3000/ws') + ws.addEventListener('message', event => events.message(String(event.data))) + ws.addEventListener('close', event => events.close({ code: event.code, reason: event.reason, wasClean: event.wasClean })) + ws.addEventListener('error', event => events.error(event)) + + await new Promise((resolve) => { + ws.addEventListener('open', () => resolve(), { once: true }) + }) + + return { + send: next => ws.send(next), + close: (code, reason) => ws.close(code, reason), + } + }, + }, +}) + +await client.connect() +client.send('hello') +``` + +## When to use + +- You need connection lifecycle, peer registry, reconnect, and broadcast primitives. +- You want to keep message shape controlled by the application or a higher-level adapter. +- You are building an Eventa, JSON-RPC, extension, or custom protocol adapter on top. + +## When not to use + +- You need a full application protocol out of the box. +- You want Socket.IO-compatible clients or packet formats. +- You only need one direct native `WebSocket` without reconnect or peer management. + +## License + +[MIT](../../LICENSE) diff --git a/packages/better-ws/package.json b/packages/better-ws/package.json new file mode 100644 index 000000000..97771be63 --- /dev/null +++ b/packages/better-ws/package.json @@ -0,0 +1,55 @@ +{ + "name": "@proj-airi/better-ws", + "type": "module", + "version": "0.10.2", + "private": true, + "description": "Transport-agnostic reliable WebSocket runtime primitives", + "author": { + "name": "Moeru AI Project AIRI Team", + "email": "airi@moeru.ai", + "url": "https://github.com/moeru-ai" + }, + "license": "MIT", + "repository": { + "type": "git", + "url": "https://github.com/moeru-ai/airi.git", + "directory": "packages/better-ws" + }, + "exports": { + ".": { + "types": "./dist/index.d.mts", + "default": "./dist/index.mjs" + }, + "./server": { + "types": "./dist/server.d.mts", + "default": "./dist/server.mjs" + }, + "./client/crossws": { + "types": "./dist/client/crossws.d.mts", + "default": "./dist/client/crossws.mjs" + }, + "./server/h3": { + "types": "./dist/server/h3.d.mts", + "default": "./dist/server/h3.mjs" + } + }, + "main": "./dist/index.mjs", + "types": "./dist/index.d.mts", + "files": [ + "README.md", + "dist", + "package.json" + ], + "scripts": { + "build": "tsdown", + "dev": "pnpm run build", + "test": "vitest", + "typecheck": "tsc --noEmit" + }, + "dependencies": { + "@moeru/eventa": "catalog:", + "crossws": "catalog:", + "h3": "catalog:", + "srvx": "catalog:" + } +} diff --git a/packages/better-ws/src/client/crossws/index.test.ts b/packages/better-ws/src/client/crossws/index.test.ts new file mode 100644 index 000000000..8e3258e84 --- /dev/null +++ b/packages/better-ws/src/client/crossws/index.test.ts @@ -0,0 +1,118 @@ +import { describe, expect, it, vi } from 'vitest' + +import { createCrossWsConnector } from '.' +import { createClient } from '../..' + +const { MockWebSocket } = vi.hoisted(() => { + class MockWebSocket { + static readonly CONNECTING = 0 + static readonly OPEN = 1 + static readonly CLOSING = 2 + static readonly CLOSED = 3 + static readonly instances: MockWebSocket[] = [] + + readonly sent: string[] = [] + readyState = MockWebSocket.CONNECTING + onclose?: (event: { code?: number, reason?: string, wasClean?: boolean }) => void + onerror?: (event: { error?: Error } | unknown) => void + onmessage?: (event: { data: string | ArrayBuffer }) => void + onopen?: () => void + + constructor(readonly url: string | URL, readonly protocols?: string | string[]) { + MockWebSocket.instances.push(this) + } + + send(message: string) { + this.sent.push(message) + } + + close() { + this.readyState = MockWebSocket.CLOSED + this.onclose?.({ code: 1000, reason: 'closed', wasClean: true }) + } + + ping = vi.fn() + pong = vi.fn() + } + + return { MockWebSocket } +}) + +function lastSocket() { + const socket = MockWebSocket.instances.at(-1) + if (!socket) { + throw new Error('Expected a mock websocket instance.') + } + + return socket +} + +describe('createCrossWsConnector', () => { + it('connects and forwards text messages through better-ws', async () => { + const client = createClient({ + connector: createCrossWsConnector({ + url: 'ws://localhost:6121/ws', + wsConstructor: MockWebSocket, + }), + reconnect: false, + }) + const messages: string[] = [] + client.onMessage(({ message }) => { + messages.push(message) + }) + + const connecting = client.connect() + const socket = lastSocket() + socket.readyState = MockWebSocket.OPEN + socket.onopen?.() + await connecting + + expect(client.state).toBe('ready') + + client.send('hello') + socket.onmessage?.({ data: 'world' }) + + expect(socket.sent).toEqual(['hello']) + expect(messages).toEqual(['world']) + }) + + it('reports non-text messages as connector errors', async () => { + const onFailed = vi.fn() + const client = createClient({ + connector: createCrossWsConnector({ + url: 'ws://localhost:6121/ws', + wsConstructor: MockWebSocket, + }), + reconnect: { + retries: 0, + onFailed, + }, + }) + + const connecting = client.connect() + const socket = lastSocket() + socket.readyState = MockWebSocket.OPEN + socket.onopen?.() + await connecting + + socket.onmessage?.({ data: new ArrayBuffer(1) }) + + expect(onFailed).toHaveBeenCalledWith(expect.any(TypeError)) + }) + + it('rejects when the socket closes before opening', async () => { + const client = createClient({ + connector: createCrossWsConnector({ + url: 'ws://localhost:6121/ws', + wsConstructor: MockWebSocket, + }), + reconnect: false, + }) + + const connecting = client.connect() + const socket = lastSocket() + socket.onclose?.({ code: 1006, reason: 'aborted', wasClean: false }) + + await expect(connecting).rejects.toThrow('closed before opening') + }) +}) diff --git a/packages/better-ws/src/client/crossws/index.ts b/packages/better-ws/src/client/crossws/index.ts new file mode 100644 index 000000000..56af89972 --- /dev/null +++ b/packages/better-ws/src/client/crossws/index.ts @@ -0,0 +1,132 @@ +import type { ClientConnector } from '../..' + +import NativeWebSocket from 'crossws/websocket' + +interface CrossWsMessageEvent { + data: unknown +} + +interface CrossWsCloseEvent { + code?: number + reason?: string + wasClean?: boolean +} + +interface CrossWsErrorEvent { + error?: unknown +} + +interface CrossWsSocket { + onclose?: ((event: CrossWsCloseEvent) => void) | null + onerror?: ((event: CrossWsErrorEvent | unknown) => void) | null + onmessage?: ((event: CrossWsMessageEvent) => void) | null + onopen?: ((event: unknown) => void) | null + close: (code?: number, reason?: string) => void + send: (message: string) => boolean | number | void + ping?: () => boolean | number | void + pong?: () => boolean | number | void +} + +/** + * WebSocket constructor accepted by the CrossWS client connector. + * + * CrossWS exports a DOM-compatible constructor in browser-like runtimes and a + * Node-backed constructor in Node. Tests may inject a narrower fake as long as + * it exposes the event handler properties and text send/close operations used + * by the connector. + */ +export interface CrossWsConstructor { + new(url: string | URL, protocols?: string | string[]): CrossWsSocket +} + +/** + * Options for the CrossWS-backed text client connector. + */ +export interface CrossWsConnectorOptions { + /** URL passed to the CrossWS socket constructor. */ + url: string | URL + /** Optional subprotocols passed to the CrossWS socket constructor. */ + protocols?: string | string[] + /** Runtime socket constructor. Defaults to `crossws/websocket`. */ + wsConstructor?: CrossWsConstructor +} + +/** Creates a CrossWS-backed text connector for better-ws clients. */ +export function createCrossWsConnector(options: CrossWsConnectorOptions): ClientConnector { + return { + connect(events) { + const WsConstructor = options.wsConstructor ?? (NativeWebSocket as unknown as CrossWsConstructor) + const ws = new WsConstructor(options.url, options.protocols) + + return new Promise((resolve, reject) => { + let opened = false + let failedBeforeOpen = false + + ws.onopen = () => { + opened = true + resolve({ + send: message => ws.send(message), + close: (code, reason) => ws.close(code, reason), + ping: typeof ws.ping === 'function' ? () => ws.ping!() : undefined, + pong: typeof ws.pong === 'function' ? () => ws.pong!() : undefined, + }) + } + + ws.onmessage = (event) => { + if (typeof event.data === 'string') { + events.message(event.data) + return + } + + events.error(new TypeError('The CrossWS connector only supports text messages.')) + } + + ws.onerror = (event) => { + const error = errorFromEvent(event) + if (!opened) { + failedBeforeOpen = true + reject(error) + return + } + + events.error(error) + } + + ws.onclose = (event) => { + if (failedBeforeOpen) { + return + } + + if (!opened) { + reject(createCloseBeforeOpenError(event)) + return + } + + events.close({ + code: event.code, + reason: event.reason, + wasClean: event.wasClean, + }) + } + }) + }, + } +} + +function errorFromEvent(event: CrossWsErrorEvent | unknown): Error { + if (event instanceof Error) { + return event + } + + if (typeof event === 'object' && event !== null && 'error' in event && event.error instanceof Error) { + return event.error + } + + return new Error('CrossWS connection error.') +} + +function createCloseBeforeOpenError(event: CrossWsCloseEvent): Error { + const reason = event.reason ? ` ${event.reason}` : '' + const code = typeof event.code === 'number' ? ` with code ${event.code}` : '' + return new Error(`CrossWS connection closed before opening${code}.${reason}`) +} diff --git a/packages/better-ws/src/client/index.ts b/packages/better-ws/src/client/index.ts new file mode 100644 index 000000000..0fa2e6028 --- /dev/null +++ b/packages/better-ws/src/client/index.ts @@ -0,0 +1,917 @@ +import type { WsCloseDetails, WsSendResult, WsState } from '../shared' + +import { createContext, defineEventa } from '@moeru/eventa' + +import { createEventWaitFor, normalizeSendResult } from '../shared' + +const clientStateChangeEvent = defineEventa('better-ws:client:state-change') + +/** + * Low-level connection adapter used by {@link Client}. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface ClientConnection { + /** Sends one caller-owned message through the active connection. */ + send: (message: TMessage) => boolean | number | void + /** Sends a native ping frame when the runtime adapter exposes one. */ + ping?: () => boolean | number | void + /** Sends a native pong frame when the runtime adapter exposes one. */ + pong?: () => boolean | number | void + /** Closes the active connection. */ + close?: (code?: number, reason?: string) => void +} + +/** + * Event sink passed to client connectors. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface ClientEvents { + /** Delivers one adapter-decoded message to the client runtime. */ + message: (message: TMessage) => void + /** Reports that the underlying transport closed. */ + close: (details?: WsCloseDetails) => void + /** + * Reports a fatal transport error for the active connection. + * + * Adapters may emit both error and close for the same failure. The client + * keeps those paths isolated by connection epoch so stale or duplicate + * follow-up events do not schedule additional reconnects. + */ + error: (error: unknown) => void +} + +/** + * Creates runtime-specific client connections. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface ClientConnector { + /** Opens a new connection and wires adapter events into the provided event sink. */ + connect: (events: ClientEvents) => Promise> | ClientConnection +} + +/** + * Timer handle used by reconnect scheduling. + */ +export interface ScheduledTask { + /** Cancels the scheduled task if it has not run yet. */ + cancel: () => void +} + +/** + * Reconnect policy for {@link createClient}. + */ +export interface ReconnectOptions { + /** + * Maximum reconnect attempts, or a predicate that decides whether the + * current attempt should run. + * + * @default Infinity + */ + retries?: number | ((attempt: number, error: unknown) => boolean) + /** + * Delay in milliseconds, or a resolver for the current attempt. + * + * Defaults to exponential backoff capped at 30 seconds: + * `Math.min(1000 * 2 ** (attempt - 1), 30000)`. + */ + delay?: number | ((attempt: number, error: unknown) => number) + /** Called when the retry policy stops reconnecting. */ + onFailed?: (error: unknown) => void + /** Whether prepare failures should schedule reconnect attempts. @default true */ + retryOnPrepareError?: boolean + /** Randomizes reconnect delay by this factor in both directions. Clamped to the 0..1 range. @default 0 */ + reconnectRandomFactor?: number + /** Minimum open duration before reconnect attempts reset. @default 0 */ + reconnectMinConnectedDuration?: number +} + +/** + * Application-level or adapter-native heartbeat policy for {@link createClient}. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface HeartbeatOptions { + /** + * Heartbeat transport mode. + * + * `auto` uses `connection.ping()` when the adapter exposes it, otherwise it + * falls back to message heartbeat when `message` is provided. `native` + * requires adapter support for `connection.ping()`. `message` sends the + * configured `message` value. + * + * @default 'auto' + */ + mode?: 'auto' | 'native' | 'message' + /** + * Delay in milliseconds between heartbeat checks while the client is ready. + * + * @default 30000 + */ + interval?: number + /** + * Maximum time in milliseconds to wait for read liveness after a heartbeat. + * + * @default 10000 + */ + timeout?: number + /** Message value, or message factory, used by message heartbeat mode. */ + message?: TMessage | (() => TMessage) + /** + * Predicate that marks an inbound message as the strict heartbeat response. + * + * When omitted, any inbound message clears the pending heartbeat timeout. + */ + isResponse?: (message: TMessage) => boolean +} + +export interface ClientMessageContext { + /** Client that received the message. */ + client: Client + /** Incoming caller-owned message. */ + message: TMessage +} + +export interface ClientStateChange { + /** Previous client state. */ + previousState: WsState + /** Current client state. */ + state: WsState +} + +/** + * Controls how long a prepare procedure waits for an incoming message. + */ +export interface WaitForOptions { + /** Maximum time in milliseconds to wait before rejecting. When omitted, no timeout is scheduled. */ + timeout?: number + /** External cancellation signal that aborts this wait independently from the owning prepare procedure. */ + signal?: AbortSignal +} + +/** + * Context passed to a client prepare procedure after the transport opens. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface PrepareContext { + /** Signal aborted when the client closes, the active connect becomes stale, or prepare is cancelled. */ + signal: AbortSignal + /** Reconnect attempt index for this connection. The first connection uses `0`. */ + attempt: number + /** Whether this prepare call belongs to a reconnect attempt. */ + reconnecting: boolean + /** Sends a bootstrap message while the client is `open`, `preparing`, or `ready`. */ + send: (message: TMessage) => WsSendResult + /** + * Waits for the first future message that matches the predicate. + * + * Rejects when its timeout expires or when the prepare/client signal aborts. + * Matching messages are still dispatched to normal `onMessage` handlers. + * Async predicates may overlap when multiple messages arrive before earlier + * predicate promises settle. + */ + waitFor: ( + predicate: (message: TMessage) => boolean | Promise, + options?: WaitForOptions, + ) => Promise +} + +/** + * Shared options for all client connection adapters. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface ClientBaseOptions { + /** Reconnect policy. Reconnect is enabled by default; pass `false` to disable. @default true */ + reconnect?: boolean | ReconnectOptions + /** Heartbeat policy. Heartbeats are disabled unless this option is provided. @default false */ + heartbeat?: false | HeartbeatOptions + /** + * Optional bootstrap procedure that must finish before the client becomes `ready`. + * + * If prepare fails with reconnect disabled, the client transitions to + * `failed`. If reconnect is enabled, the current `connect()` rejects and a + * retry is scheduled through the reconnect policy. + */ + prepare?: (context: PrepareContext) => Promise | void + /** Injectable scheduler for tests or custom timer runtimes. */ + schedule?: (delay: number, run: () => void) => ScheduledTask +} + +/** + * Options for a client backed by a caller-provided runtime connector. + * + * @param TMessage - Message shape owned by the caller. + */ +export interface ClientConnectorOptions extends ClientBaseOptions { + /** Adapter that opens concrete runtime connections. */ + connector: ClientConnector +} + +/** + * Options for the built-in native text socket adapter. + */ +export interface ClientUrlOptions extends ClientBaseOptions { + /** URL passed to the socket constructor. */ + url: string | URL + /** Optional subprotocols passed to the socket constructor. */ + protocols?: string | string[] + /** WebSocket constructor. Defaults to `globalThis.WebSocket` when available. */ + wsConstructor?: typeof WebSocket +} + +export type ClientOptions = ClientConnectorOptions | ClientUrlOptions + +/** + * Options that control local send gating. + */ +export interface ClientSendOptions { + /** + * Whether `send` requires the client to be `ready`. + * + * When `true` or omitted, sends are allowed only in `ready`. When `false`, + * sends may run in `open`, `preparing`, or `ready`, which is intended for + * connection preparation and protocol bootstrap messages. + * + * @default true + */ + requireReady?: boolean +} + +export interface Client { + /** Current connection lifecycle state. */ + readonly state: WsState + /** + * Opens the client connection. + * + * Calling `connect()` again replaces any active or pending connection. This + * is intentional restart behavior, not idempotent ensure-connected behavior. + */ + connect: () => Promise + /** Sends one message over the active connection. */ + send: (message: TMessage, options?: ClientSendOptions) => WsSendResult + /** Closes the client and suppresses reconnect scheduling. */ + close: (code?: number, reason?: string) => void + /** Registers an incoming message handler. */ + onMessage: (handler: (context: ClientMessageContext) => void | Promise) => () => void + /** Registers a state change handler. */ + onStateChange: (handler: (change: ClientStateChange) => void) => () => void +} + +function defaultSchedule(delay: number, run: () => void): ScheduledTask { + const handle = setTimeout(run, delay) + return { + cancel: () => clearTimeout(handle), + } +} + +function createConnectionClosedError() { + return new Error('Connection closed') +} + +/** + * Normalizes reconnect policy into the runtime shape used by close and + * prepare-failure paths. + * + * Before: + * - `undefined` + * - `true` + * - `{ delay: 100 }` + * + * After: + * - full reconnect policy with default infinite retries and exponential backoff + */ +function normalizeReconnectOptions(reconnect: ClientBaseOptions['reconnect']): false | Required { + if (reconnect === false) { + return false + } + + const options = reconnect === true || typeof reconnect === 'undefined' ? {} : reconnect + return { + retries: options.retries ?? Number.POSITIVE_INFINITY, + delay: options.delay ?? ((attempt: number) => Math.min(1000 * 2 ** (attempt - 1), 30_000)), + onFailed: options.onFailed ?? (() => {}), + retryOnPrepareError: options.retryOnPrepareError ?? true, + reconnectRandomFactor: normalizeReconnectRandomFactor(options.reconnectRandomFactor), + reconnectMinConnectedDuration: options.reconnectMinConnectedDuration ?? 0, + } +} + +function shouldRetry(retries: Required['retries'], attempt: number, error: unknown): boolean { + return typeof retries === 'number' + ? attempt <= retries + : retries(attempt, error) +} + +function resolveReconnectDelay(options: Required, attempt: number, error: unknown): number { + return typeof options.delay === 'number' ? options.delay : options.delay(attempt, error) +} + +function normalizeReconnectRandomFactor(randomFactor: number | undefined): number { + if (typeof randomFactor !== 'number' || Number.isNaN(randomFactor)) { + return 0 + } + + return Math.min(1, Math.max(0, randomFactor)) +} + +function applyReconnectRandomFactor(delay: number, randomFactor: number): number { + if (randomFactor <= 0 || delay <= 0) { + return delay + } + + const factor = 1 + ((Math.random() * 2 - 1) * randomFactor) + return Math.max(1, Math.round(delay * factor)) +} + +type NormalizedHeartbeatOptions = Required, 'mode' | 'interval' | 'timeout'>> + & Pick, 'message' | 'isResponse'> + +function normalizeHeartbeatOptions( + heartbeat: ClientBaseOptions['heartbeat'], +): false | NormalizedHeartbeatOptions { + if (!heartbeat) { + return false + } + + return { + mode: heartbeat.mode ?? 'auto', + interval: heartbeat.interval ?? 30_000, + timeout: heartbeat.timeout ?? 10_000, + message: heartbeat.message, + isResponse: heartbeat.isResponse, + } +} + +function createHeartbeatTimeoutError(timeout: number) { + return new Error(`Heartbeat timed out after ${timeout}ms.`) +} + +/** + * Creates a client that owns reconnect, state tracking, and handler dispatch. + * + * Pass `url` for the built-in text socket adapter. Pass `connector` when the + * runtime is not the browser global socket, or when messages need custom + * serialization before they reach the client runtime. + */ +export function createClient(options: ClientUrlOptions): Client +export function createClient(options: ClientConnectorOptions): Client +export function createClient(options: ClientOptions): Client | Client { + if ('connector' in options) { + return createClientWithConnector(options, options.connector) + } + + return createClientWithConnector(options, createSocketConnector(options)) +} + +function createClientWithConnector( + options: ClientBaseOptions, + connector: ClientConnector, +): Client { + let state: WsState = 'idle' + let connection: ClientConnection | undefined + let reconnectAttempt = 0 + let manuallyClosed = false + + const reconnectOptions = normalizeReconnectOptions(options.reconnect) + const heartbeatOptions = normalizeHeartbeatOptions(options.heartbeat) + + let lastCloseError: unknown = createConnectionClosedError() + let reconnectTask: ScheduledTask | undefined + let connectedAt: number | undefined + let heartbeatIntervalTask: ScheduledTask | undefined + let heartbeatTimeoutTask: ScheduledTask | undefined + let prepareController: AbortController | undefined + + const waiters = new Set<(message: TMessage) => void>() + // NOTICE: + // `connect()` can resolve after `close()` or a newer reconnect has already + // changed the active connection. Track a monotonic connection epoch so stale + // async completions close their own connection instead of replacing current + // client state. + let connectionEpoch = 0 + const events = createContext() + const clientMessageEvent = defineEventa>('better-ws:client:message') + + const client: Client = { + get state() { + return state + }, + async connect() { + await connectInternal(false) + }, + send(message, sendOptions) { + const requiredState = sendOptions?.requireReady === false ? 'open' : 'ready' + const canSend = requiredState === 'ready' + ? state === 'ready' + : state === 'open' || state === 'preparing' || state === 'ready' + + if (!connection || !canSend) { + return { ok: false, reason: 'closed' } + } + + return normalizeSendResult(() => connection?.send(message)) + }, + close(code, reason) { + manuallyClosed = true + connectionEpoch += 1 + stopHeartbeat() + + prepareController?.abort() + prepareController = undefined + + reconnectTask?.cancel() + reconnectTask = undefined + + transition('closing') + + connection?.close?.(code, reason) + connection = undefined + + connectedAt = undefined + + transition('closed') + }, + onMessage(handler) { + return events.on(clientMessageEvent, event => handler(event.body!)) + }, + onStateChange(handler) { + return events.on(clientStateChangeEvent, event => handler(event.body!)) + }, + } + + async function connectInternal(automaticReconnect: boolean): Promise { + manuallyClosed = false + const currentConnectionEpoch = ++connectionEpoch + stopHeartbeat() + + prepareController?.abort() + prepareController = undefined + + connection?.close?.() + connection = undefined + + reconnectTask?.cancel() + reconnectTask = undefined + + connectedAt = undefined + + transition(automaticReconnect ? 'reconnecting' : 'connecting') + + let nextConnection: ClientConnection + try { + nextConnection = await connector.connect({ + message: message => dispatchMessage(currentConnectionEpoch, message), + close: details => handleClose(currentConnectionEpoch, details), + error: error => dispatchError(currentConnectionEpoch, error), + }) + } + catch (error) { + if (currentConnectionEpoch === connectionEpoch) { + stopHeartbeat() + if (reconnectOptions && !manuallyClosed) { + lastCloseError = error + scheduleReconnect(error) + if (automaticReconnect) { + return + } + } + + else { + transition('closed') + } + } + if (automaticReconnect) { + return + } + + throw error + } + + if ((currentConnectionEpoch !== connectionEpoch || manuallyClosed)) { + (prepareController as AbortController | undefined)?.abort() + prepareController = undefined + + nextConnection.close?.() + return + } + + connection = nextConnection + connectedAt = Date.now() + lastCloseError = createConnectionClosedError() + + transition('open') + + let currentPrepareController: AbortController | undefined + try { + if (options.prepare) { + transition('preparing') + + currentPrepareController = new AbortController() + prepareController = currentPrepareController + + await options.prepare({ + signal: currentPrepareController.signal, + attempt: reconnectAttempt, + reconnecting: reconnectAttempt > 0, + send: message => client.send(message, { requireReady: false }), + waitFor: createWaitForMessage(currentPrepareController.signal, currentConnectionEpoch), + }) + + currentPrepareController.abort() + if (prepareController === currentPrepareController) { + prepareController = undefined + } + } + } + catch (error) { + currentPrepareController?.abort() + if (prepareController === currentPrepareController) { + prepareController = undefined + } + if (currentConnectionEpoch !== connectionEpoch || manuallyClosed) { + if (connection === nextConnection) { + nextConnection.close?.() + connection = undefined + } + + return + } + + connectionEpoch += 1 + stopHeartbeat() + + nextConnection.close?.() + connection = undefined + + resetReconnectAttemptAfterStableConnection() + + lastCloseError = error + if (reconnectOptions && reconnectOptions.retryOnPrepareError) { + scheduleReconnect(error) + } + else { + transition('failed') + } + if (automaticReconnect) { + return + } + + throw error + } + + if (currentConnectionEpoch !== connectionEpoch || manuallyClosed) { + currentPrepareController?.abort() + if (prepareController === currentPrepareController) { + prepareController = undefined + } + if (connection === nextConnection) { + stopHeartbeat() + nextConnection.close?.() + connection = undefined + } + + return + } + + if (!reconnectOptions || reconnectOptions.reconnectMinConnectedDuration <= 0) { + reconnectAttempt = 0 + } + + transition('ready') + scheduleHeartbeat(currentConnectionEpoch) + } + + function transition(nextState: WsState) { + if (state === nextState) { + return + } + + const previousState = state + state = nextState + events.emit(clientStateChangeEvent, { previousState, state }) + } + + function dispatchMessage(connectionMessageEpoch: number, message: TMessage) { + if (connectionMessageEpoch !== connectionEpoch) { + return + } + + refreshHeartbeat(connectionMessageEpoch, message) + + events.emit(clientMessageEvent, { client, message }) + for (const waiter of waiters) { + waiter(message) + } + } + + function dispatchError(connectionErrorEpoch: number, error: unknown) { + if (connectionErrorEpoch !== connectionEpoch) { + return + } + + lastCloseError = error + const erroredConnection = connection + erroredConnection?.close?.() + handleClose(connectionErrorEpoch, { reason: 'error' }) + } + + function handleClose(connectionCloseEpoch: number, _details?: WsCloseDetails) { + if (connectionCloseEpoch !== connectionEpoch) { + return + } + + connectionEpoch += 1 + stopHeartbeat() + + prepareController?.abort() + prepareController = undefined + connection = undefined + + resetReconnectAttemptAfterStableConnection() + if (manuallyClosed) { + transition('closed') + return + } + + if (!reconnectOptions) { + transition('closed') + return + } + + scheduleReconnect(lastCloseError) + } + + function scheduleReconnect(error: unknown) { + stopHeartbeat() + + const nextAttempt = reconnectAttempt + 1 + if (!reconnectOptions) { + failReconnect(error) + return + } + + let retryAllowed: boolean + try { + retryAllowed = shouldRetry(reconnectOptions.retries, nextAttempt, error) + } + catch (policyError) { + failReconnect(policyError) + + return + } + + if (!retryAllowed) { + failReconnect(error) + + return + } + + let delay: number + try { + delay = applyReconnectRandomFactor( + resolveReconnectDelay(reconnectOptions, nextAttempt, error), + reconnectOptions.reconnectRandomFactor, + ) + } + catch (policyError) { + failReconnect(policyError) + return + } + + reconnectAttempt = nextAttempt + + transition('reconnecting') + + const schedule = options.schedule ?? defaultSchedule + reconnectTask = schedule(delay, () => { + void connectInternal(true) + }) + } + + function resetReconnectAttemptAfterStableConnection() { + if (!reconnectOptions || reconnectOptions.reconnectMinConnectedDuration <= 0 || connectedAt === undefined) { + connectedAt = undefined + return + } + + const connectedFor = Date.now() - connectedAt + connectedAt = undefined + if (connectedFor >= reconnectOptions.reconnectMinConnectedDuration) { + reconnectAttempt = 0 + } + } + + function scheduleHeartbeat(heartbeatConnectionEpoch: number) { + if (!heartbeatOptions || heartbeatConnectionEpoch !== connectionEpoch || state !== 'ready') { + return + } + + heartbeatIntervalTask?.cancel() + + const schedule = options.schedule ?? defaultSchedule + heartbeatIntervalTask = schedule(heartbeatOptions.interval, () => { + heartbeatIntervalTask = undefined + runHeartbeat(heartbeatConnectionEpoch) + }) + } + + function runHeartbeat(heartbeatConnectionEpoch: number) { + if (!heartbeatOptions || heartbeatConnectionEpoch !== connectionEpoch || state !== 'ready' || manuallyClosed) { + return + } + + const heartbeatSent = sendHeartbeat() + if (!heartbeatSent) { + if (heartbeatOptions.mode === 'native') { + failHeartbeat(heartbeatConnectionEpoch, new Error('Native heartbeat requires connection.ping().')) + } + else if (heartbeatOptions.mode === 'message' && heartbeatOptions.message === undefined) { + failHeartbeat(heartbeatConnectionEpoch, new Error('Message heartbeat requires heartbeat.message.')) + } + return + } + + heartbeatTimeoutTask?.cancel() + const schedule = options.schedule ?? defaultSchedule + heartbeatTimeoutTask = schedule(heartbeatOptions.timeout, () => { + heartbeatTimeoutTask = undefined + failHeartbeat(heartbeatConnectionEpoch, createHeartbeatTimeoutError(heartbeatOptions.timeout)) + }) + } + + function sendHeartbeat(): boolean { + if (!heartbeatOptions) { + return false + } + + if ((heartbeatOptions.mode === 'auto' || heartbeatOptions.mode === 'native') && connection?.ping) { + // Since we asserted that connection?.ping() defined, + // then here we explicitly ! the invoke. + const result = normalizeSendResult(() => connection?.ping!()) + return result.ok + } + + if ((heartbeatOptions.mode === 'auto' || heartbeatOptions.mode === 'message') && heartbeatOptions.message !== undefined) { + const heartbeatMessage = typeof heartbeatOptions.message === 'function' + ? (heartbeatOptions.message as () => TMessage)() + : heartbeatOptions.message + + return client.send(heartbeatMessage, { requireReady: false }).ok + } + + return false + } + + function refreshHeartbeat(heartbeatConnectionEpoch: number, message: TMessage) { + if (!heartbeatOptions || heartbeatConnectionEpoch !== connectionEpoch) { + return + } + + const isStrictResponse = heartbeatOptions.isResponse?.(message) + if (isStrictResponse ?? true) { + heartbeatTimeoutTask?.cancel() + heartbeatTimeoutTask = undefined + } + else if (heartbeatTimeoutTask) { + return + } + + if (state === 'ready') { + scheduleHeartbeat(heartbeatConnectionEpoch) + } + } + + function failHeartbeat(heartbeatConnectionEpoch: number, error: unknown) { + if (heartbeatConnectionEpoch !== connectionEpoch || manuallyClosed) { + return + } + + lastCloseError = error + const timedOutConnection = connection + connectionEpoch += 1 + stopHeartbeat() + prepareController?.abort() + prepareController = undefined + connection = undefined + resetReconnectAttemptAfterStableConnection() + timedOutConnection?.close?.() + + if (manuallyClosed) { + transition('closed') + return + } + + if (!reconnectOptions) { + transition('closed') + return + } + + scheduleReconnect(error) + } + + function stopHeartbeat() { + heartbeatIntervalTask?.cancel() + heartbeatIntervalTask = undefined + heartbeatTimeoutTask?.cancel() + heartbeatTimeoutTask = undefined + } + + function failReconnect(error: unknown) { + transition('failed') + if (!reconnectOptions) { + return + } + + try { + reconnectOptions.onFailed(error) + } + catch { + // onFailed is a notification hook; state has already moved to failed. + } + } + + function createWaitForMessage( + activePrepareSignal: AbortSignal, + prepareConnectionEpoch: number, + ) { + return ( + predicate: (message: TMessage) => boolean | Promise, + waitOptions: WaitForOptions = {}, + ): Promise => { + const wait = createEventWaitFor({ + match: predicate, + timeout: waitOptions.timeout, + signals: [activePrepareSignal, waitOptions.signal], + isActive: () => prepareConnectionEpoch === connectionEpoch, + abortMessage: 'Wait for message aborted.', + timeoutMessage: 'Timed out waiting for message.', + }) + + waiters.add(wait.emit) + + void wait.promise.finally(() => { + waiters.delete(wait.emit) + }).catch(() => {}) + + return wait.promise + } + } + + return client +} + +function createSocketConnector(options: ClientUrlOptions): ClientConnector { + return { + connect(events) { + const WsConstructor = options.wsConstructor ?? globalThis.WebSocket + if (!WsConstructor) { + throw new Error('No WebSocket constructor is available. Pass `wsConstructor` or use a connector.') + } + + const ws = new WsConstructor(options.url, options.protocols) + return new Promise>((resolve, reject) => { + let opened = false + let failedBeforeOpen = false + ws.onopen = () => { + opened = true + resolve({ + send: message => ws.send(message), + close: (code, reason) => ws.close(code, reason), + }) + } + ws.onmessage = (event) => { + if (typeof event.data === 'string') { + events.message(event.data) + return + } + + events.error(new TypeError('The built-in WebSocket connector only supports text messages.')) + } + ws.onerror = (event) => { + if (!opened) { + failedBeforeOpen = true + reject(new Error('WebSocket connection failed before opening.')) + return + } + + events.error(event) + } + ws.onclose = (event) => { + if (failedBeforeOpen) { + return + } + + events.close({ + code: event.code, + reason: event.reason, + wasClean: event.wasClean, + }) + } + }) + }, + } +} diff --git a/packages/better-ws/src/index.test.ts b/packages/better-ws/src/index.test.ts new file mode 100644 index 000000000..49db2cc7d --- /dev/null +++ b/packages/better-ws/src/index.test.ts @@ -0,0 +1,2489 @@ +import type { Message as CrossWsMessage, Peer as CrossWsPeer } from 'crossws' + +import { describe, expect, it, vi } from 'vitest' + +import { createServer, toCrossWsHooks } from './server' + +import * as betterWs from './index' + +function errorText(error: unknown) { + if (typeof error === 'object' && error !== null && 'message' in error && typeof error.message === 'string') { + return error.message + } + + return String(error) +} + +class FakeWebSocket extends EventTarget implements WebSocket { + static readonly CONNECTING = 0 + static readonly OPEN = 1 + static readonly CLOSING = 2 + static readonly CLOSED = 3 + static readonly instances: FakeWebSocket[] = [] + + readonly CONNECTING = 0 + readonly OPEN = 1 + readonly CLOSING = 2 + readonly CLOSED = 3 + binaryType: BinaryType = 'blob' + readonly bufferedAmount = 0 + readonly extensions = '' + onclose: ((this: WebSocket, ev: CloseEvent) => unknown) | null = null + onerror: ((this: WebSocket, ev: Event) => unknown) | null = null + onmessage: ((this: WebSocket, ev: MessageEvent) => unknown) | null = null + onopen: ((this: WebSocket, ev: Event) => unknown) | null = null + readonly protocol = '' + readonly readyState = FakeWebSocket.CONNECTING + + readonly sent: string[] = [] + + readonly url: string + + constructor(url: string | URL) { + super() + this.url = String(url) + FakeWebSocket.instances.push(this) + } + + send(message: string | ArrayBufferLike | Blob | ArrayBufferView) { + if (typeof message === 'string') { + this.sent.push(message) + } + } + + close() {} + + open() { + this.onopen?.(new Event('open')) + } + + error() { + this.onerror?.(new Event('error')) + } + + closeEvent() { + this.onclose?.(new CloseEvent('close', { code: 1006, reason: 'open failed', wasClean: false })) + } + + receive(message: string) { + this.onmessage?.(new MessageEvent('message', { data: message })) + } +} + +function createFakeSocketClient() { + const client = betterWs.createClient({ + url: 'ws://localhost/ws', + wsConstructor: FakeWebSocket, + }) + const socket = () => { + const instance = FakeWebSocket.instances.at(-1) + if (!instance) { + throw new Error('FakeWebSocket was not constructed.') + } + return instance + } + return { + client, + get socket() { + return socket() + }, + } +} + +describe('better-ws package exports', () => { + it('keeps server APIs behind the server subpath', () => { + expect('createClient' in betterWs).toBe(true) + expect('createServer' in betterWs).toBe(false) + }) +}) + +describe('better-ws server runtime', () => { + it('exposes peers through a peer manager object', () => { + const server = createServer() + const sent: string[] = [] + + const peer = server.peers.accept({ + id: 'peer-1', + send: message => sent.push(message), + }).peer + + peer.send('hello') + + expect(server.peers.has('peer-1')).toBe(true) + expect(server.peers.get('peer-1')).toBe(peer) + expect(server.peers.list()).toEqual([peer]) + expect([...server.peers.entries()]).toEqual([['peer-1', peer]]) + expect(sent).toEqual(['hello']) + }) + + it('keeps server-level accept and remove as peer manager shortcuts', () => { + const server = createServer() + + const peer = server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }) + server.remove('peer-1') + + expect(peer.id).toBe('peer-1') + expect(server.peers.has('peer-1')).toBe(false) + }) + + it('keeps replacement accepted during adapter close', () => { + const server = createServer() + const replacementSend = vi.fn(() => true) + const first = server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + close: () => { + server.accept({ + id: 'peer-1', + send: replacementSend, + }) + }, + }) + + first.close() + server.peers.get('peer-1')?.send('replacement') + + expect(server.peers.has('peer-1')).toBe(true) + expect(replacementSend).toHaveBeenCalledExactlyOnceWith('replacement') + }) + + it('continues closing peers after one peer close throws during closeAll', () => { + const server = createServer() + const firstClose = vi.fn(() => { + throw new Error('first close failed') + }) + const secondClose = vi.fn() + + server.accept({ + id: 'first', + send: vi.fn(() => true), + close: firstClose, + }) + server.accept({ + id: 'second', + send: vi.fn(() => true), + close: secondClose, + }) + + expect(() => server.close()).toThrow('first close failed') + expect(firstClose).toHaveBeenCalledOnce() + expect(secondClose).toHaveBeenCalledOnce() + expect(server.peers.size).toBe(0) + }) + + it('cleans empty group records after peers leave or are removed', () => { + const server = createServer() + const peer = server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }) + + peer.join('room:a') + peer.leave('room:a') + peer.join('room:b') + server.remove('peer-1') + + expect(server.to('room:a').send('stale-room')).toEqual([]) + expect(server.to('room:b').send('stale-room')).toEqual([]) + }) + + it('registers peers, dispatches raw messages, and disposes peer handlers', () => { + const server = createServer() + const sent: string[] = [] + const received: Array<{ peerId: string, message: string }> = [] + const unsubscribe = server.onMessage((context) => { + received.push({ peerId: context.peer.id, message: context.message }) + }) + + const peer = server.accept({ + id: 'peer-1', + send: (message) => { + sent.push(message) + return true + }, + }) + + peer.receive('hello') + peer.send('reply') + unsubscribe() + peer.receive('ignored') + + expect(server.peers.size).toBe(1) + expect(received).toEqual([{ peerId: 'peer-1', message: 'hello' }]) + expect(sent).toEqual(['reply']) + }) + + it('broadcasts to all peers and sends to named groups only', () => { + const server = createServer() + const firstSent: string[] = [] + const secondSent: string[] = [] + + const first = server.accept({ id: 'first', send: message => firstSent.push(message) }) + server.accept({ + id: 'second', + send: (message) => { + secondSent.push(message) + }, + }) + + first.join('room:a') + + const broadcast = server.broadcast('global') + const room = server.to('room:a').send('room-only') + + expect(broadcast).toEqual([ + { peerId: 'first', ok: true }, + { peerId: 'second', ok: true }, + ]) + expect(room).toEqual([{ peerId: 'first', ok: true }]) + expect(firstSent).toEqual(['global', 'room-only']) + expect(secondSent).toEqual(['global']) + }) + + it('adapts CrossWS hooks into server peers and raw messages', () => { + const server = createServer() + const received: string[] = [] + const sent: string[] = [] + server.onMessage(({ message }) => { + received.push(message) + }) + + const hooks = toCrossWsHooks(server) + const peer = { + id: 'crossws-peer', + send: (message: unknown) => { + sent.push(String(message)) + }, + close: vi.fn(), + } + + // NOTICE: + // CrossWS peers and messages are runtime-owned objects with a wider shape + // than better-ws needs here. The fake only models the fields used by the + // adapter, so the cast stays local to this adapter-boundary test. + // Remove this when CrossWS provides a small public testing fixture type. + hooks.open?.(peer as unknown as CrossWsPeer) + hooks.message?.(peer as unknown as CrossWsPeer, { text: () => 'hello' } as unknown as CrossWsMessage) + server.peers.get('crossws-peer')?.send('reply') + + expect(received).toEqual(['hello']) + expect(sent).toEqual(['reply']) + }) + + it('removes CrossWS peers without closing an already closed raw connection', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const peer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // CrossWS close hooks receive runtime-owned peers. This fake keeps the test + // focused on better-ws registry cleanup instead of depending on CrossWS + // internals. Remove this when CrossWS exposes a narrow fake peer helper. + hooks.open?.(peer as unknown as CrossWsPeer) + await hooks.close?.(peer as unknown as CrossWsPeer, { code: 1000, reason: 'done' }) + + expect(server.peers.has('crossws-peer')).toBe(false) + expect(peer.close).not.toHaveBeenCalled() + }) + + it('passes CrossWS close details to peer close handlers', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const closed: Array<{ peerId: string, code?: number, reason?: string }> = [] + server.onPeerClose(({ peerId, details }) => { + closed.push({ + peerId, + code: details?.code, + reason: details?.reason, + }) + }) + const peer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + hooks.open?.(peer as unknown as CrossWsPeer) + await hooks.close?.(peer as unknown as CrossWsPeer, { code: 1001, reason: 'runtime close' }) + + expect(closed).toEqual([{ peerId: 'crossws-peer', code: 1001, reason: 'runtime close' }]) + }) + + it('replaces an existing peer when the adapter reuses a peer id', () => { + const server = createServer() + const firstSent: string[] = [] + const secondSent: string[] = [] + + const first = server.accept({ + id: 'same-peer', + send: message => firstSent.push(message), + }) + first.join('room:a') + + server.accept({ + id: 'same-peer', + send: message => secondSent.push(message), + }) + + const result = server.to('room:a').send('stale-room') + server.peers.get('same-peer')?.send('direct') + + expect(result).toEqual([]) + expect(firstSent).toEqual([]) + expect(secondSent).toEqual(['direct']) + }) + + it('returns previous peer snapshot when accepting the same id', () => { + const server = createServer() + + const first = server.peers.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }, { + state: { token: 'first-token' }, + }).peer + first.join('ready') + + const { peer: second, previous } = server.peers.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }) + + expect(previous).toEqual({ + id: 'peer-1', + state: { token: 'first-token' }, + groups: ['ready'], + lastSeenAt: expect.any(Number), + reason: 'replaced', + }) + expect(second.state).toEqual({ token: 'first-token' }) + expect(second.isIn('ready')).toBe(false) + }) + + it('keeps stale peer handles inert after same-id replacement', () => { + const server = createServer() + const firstSend = vi.fn(() => true) + const firstClose = vi.fn() + const secondSend = vi.fn(() => true) + const secondClose = vi.fn() + const received: string[] = [] + server.onMessage(({ message }) => { + received.push(message) + }) + + const first = server.peers.accept({ + id: 'peer-1', + send: firstSend, + close: firstClose, + }).peer + first.join('room') + + const second = server.peers.accept({ + id: 'peer-1', + send: secondSend, + close: secondClose, + }).peer + second.join('room') + + expect(first.send('stale')).toEqual({ ok: false, reason: 'closed' }) + first.receive('stale') + first.close() + first.join('stale-room') + first.leave('room') + + expect(firstSend).not.toHaveBeenCalled() + expect(secondClose).not.toHaveBeenCalled() + expect(received).toEqual([]) + expect(server.peers.get('peer-1')).toBe(second) + expect(server.to('room').send('fresh')).toEqual([{ peerId: 'peer-1', ok: true }]) + }) + + it('keeps stale peer handles inert after replacement', () => { + const server = createServer() + const firstSend = vi.fn(() => true) + const firstClose = vi.fn() + const secondSend = vi.fn(() => true) + const secondClose = vi.fn() + const received: Array<{ peerId: string, message: string }> = [] + server.onMessage(({ peer, message }) => { + received.push({ peerId: peer.id, message }) + }) + + const firstPeer = server.accept({ + id: 'peer-1', + send: firstSend, + close: firstClose, + }) + firstPeer.join('room') + + const secondPeer = server.accept({ + id: 'peer-1', + send: secondSend, + close: secondClose, + }) + secondPeer.join('room') + + const staleSend = firstPeer.send('stale-send') + firstPeer.receive('stale-receive') + firstPeer.close() + firstPeer.join('stale-room') + firstPeer.leave('room') + + expect(staleSend).toEqual({ ok: false, reason: 'closed' }) + expect(firstSend).not.toHaveBeenCalled() + expect(secondClose).not.toHaveBeenCalled() + expect(received).toEqual([]) + expect(server.peers.get('peer-1')).toBe(secondPeer) + expect(secondPeer.isIn('room')).toBe(true) + expect(server.to('room').send('current-room')).toEqual([{ peerId: 'peer-1', ok: true }]) + expect(server.to('stale-room').send('stale-room')).toEqual([]) + expect(secondSend).toHaveBeenCalledExactlyOnceWith('current-room') + }) + + it('rebinds repeated CrossWS opens with the same id to the newest raw peer', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // CrossWS peers are runtime-owned objects. These fakes model only the + // adapter fields better-ws consumes, keeping the replacement behavior under + // test without coupling to CrossWS internals. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + await hooks.open?.(secondRawPeer as unknown as CrossWsPeer) + server.peers.get('crossws-peer')?.send('reply') + + expect(firstRawPeer.send).not.toHaveBeenCalled() + expect(secondRawPeer.send).toHaveBeenCalledExactlyOnceWith('reply', { compress: undefined }) + }) + + it('ignores stale CrossWS close events after same-id reopen', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // CrossWS raw peer identity determines which runtime connection owns a + // close event. These fakes keep the test focused on adapter cache + // isolation for repeated ids. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + await hooks.open?.(secondRawPeer as unknown as CrossWsPeer) + await hooks.close?.(firstRawPeer as unknown as CrossWsPeer, { code: 1000, reason: 'stale' }) + server.peers.get('crossws-peer')?.send('reply') + + expect(server.peers.has('crossws-peer')).toBe(true) + expect(firstRawPeer.send).not.toHaveBeenCalled() + expect(secondRawPeer.send).toHaveBeenCalledExactlyOnceWith('reply', { compress: undefined }) + }) + + it('ignores late CrossWS messages from a stale raw peer after same-id reopen', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const received: Array<{ peerId: string, message: string }> = [] + server.onMessage(({ peer, message }) => { + received.push({ peerId: peer.id, message }) + }) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // CrossWS may deliver late events from a replaced raw connection. The fake + // peers share an id but have different object identity, which is the + // adapter-level ownership boundary under test. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + await hooks.open?.(secondRawPeer as unknown as CrossWsPeer) + hooks.message?.(firstRawPeer as unknown as CrossWsPeer, { text: () => 'stale-message' } as unknown as CrossWsMessage) + + expect(received).toEqual([]) + expect(server.peers.has('crossws-peer')).toBe(true) + }) + + it('keeps replaced CrossWS raw peers stale after the replacement peer closes', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const received: string[] = [] + server.onMessage(({ message }) => { + received.push(message) + }) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + await hooks.open?.(secondRawPeer as unknown as CrossWsPeer) + await hooks.close?.(secondRawPeer as unknown as CrossWsPeer, { code: 1000, reason: 'new closed' }) + hooks.message?.(firstRawPeer as unknown as CrossWsPeer, { text: () => 'stale-after-close' } as unknown as CrossWsMessage) + + expect(received).toEqual([]) + expect(server.peers.has('crossws-peer')).toBe(false) + }) + + it('does not expose the current peer to stale CrossWS close hooks after same-id reopen', async () => { + const server = createServer() + const close = vi.fn() + const hooks = toCrossWsHooks(server, { close }) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // A stale close belongs to the replaced raw connection, not the current + // same-id server peer. The fake peers keep that object identity distinction + // visible at the CrossWS adapter boundary. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + await hooks.open?.(secondRawPeer as unknown as CrossWsPeer) + await hooks.close?.(firstRawPeer as unknown as CrossWsPeer, { code: 1000, reason: 'stale' }) + + expect(close).toHaveBeenCalledOnce() + expect(close.mock.calls[0]?.[0].peer).toBeUndefined() + expect(server.peers.has('crossws-peer')).toBe(true) + }) + + it('ignores stale CrossWS close after server-side same-id replacement', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // This models a replacement that happens through the server peer manager + // rather than a second CrossWS open. A late raw close from the adapter cache + // must not remove the current same-id server peer. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + server.accept({ + id: 'crossws-peer', + send: message => secondRawPeer.send(message), + close: (code, reason) => secondRawPeer.close(code, reason), + }) + await hooks.close?.(firstRawPeer as unknown as CrossWsPeer, { code: 1000, reason: 'stale' }) + server.peers.get('crossws-peer')?.send('reply') + + expect(server.peers.has('crossws-peer')).toBe(true) + expect(firstRawPeer.send).not.toHaveBeenCalled() + expect(secondRawPeer.send).toHaveBeenCalledExactlyOnceWith('reply') + }) + + it('ignores stale CrossWS messages after server-side same-id replacement', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const received: string[] = [] + server.onMessage(({ message }) => { + received.push(message) + }) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // A late message from a raw peer replaced through the server peer manager + // must stay stale. Re-accepting it would remove the current same-id peer + // and route messages through the wrong raw connection. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + server.accept({ + id: 'crossws-peer', + send: message => secondRawPeer.send(message), + close: (code, reason) => secondRawPeer.close(code, reason), + }) + hooks.message?.(firstRawPeer as unknown as CrossWsPeer, { text: () => 'stale-message' } as unknown as CrossWsMessage) + server.peers.get('crossws-peer')?.send('reply') + + expect(received).toEqual([]) + expect(firstRawPeer.send).not.toHaveBeenCalled() + expect(secondRawPeer.send).toHaveBeenCalledExactlyOnceWith('reply') + expect(server.peers.has('crossws-peer')).toBe(true) + }) + + it('keeps server-side replaced CrossWS raw peers stale after the replacement peer closes', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const received: string[] = [] + server.onMessage(({ message }) => { + received.push(message) + }) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + const replacement = server.accept({ + id: 'crossws-peer', + send: message => secondRawPeer.send(message), + close: (code, reason) => secondRawPeer.close(code, reason), + }) + replacement.close() + hooks.message?.(firstRawPeer as unknown as CrossWsPeer, { text: () => 'stale-after-server-close' } as unknown as CrossWsMessage) + + expect(received).toEqual([]) + expect(server.peers.has('crossws-peer')).toBe(false) + }) + + it('refreshes the CrossWS peer cache after server-side peer close', () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const received: string[] = [] + server.onMessage(({ message }) => { + received.push(message) + }) + const rawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // CrossWS messages and peers are runtime-owned objects with a wider public + // shape than better-ws needs. These fakes model only text reads and peer IO + // used by this adapter-boundary test. + // Remove this when CrossWS exposes small public testing fixtures. + hooks.open?.(rawPeer as unknown as CrossWsPeer) + server.peers.get('crossws-peer')?.close() + hooks.message?.(rawPeer as unknown as CrossWsPeer, { text: () => 'fresh-message' } as unknown as CrossWsMessage) + server.peers.get('crossws-peer')?.send('fresh-send') + + expect(received).toEqual(['fresh-message']) + expect(rawPeer.close).toHaveBeenCalledOnce() + expect(rawPeer.send).toHaveBeenCalledExactlyOnceWith('fresh-send', { compress: undefined }) + expect(server.peers.has('crossws-peer')).toBe(true) + }) + + it('refreshes the CrossWS peer cache after server close', async () => { + const server = createServer() + const hooks = toCrossWsHooks(server) + const firstRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + const secondRawPeer = { + id: 'crossws-peer', + send: vi.fn(), + close: vi.fn(), + } + + // NOTICE: + // CrossWS peers are runtime-owned objects. These fakes model only the + // adapter fields better-ws consumes, keeping cache refresh behavior under + // test without depending on CrossWS internals. + // Remove this when CrossWS exposes a narrow fake peer helper. + await hooks.open?.(firstRawPeer as unknown as CrossWsPeer) + server.close() + await hooks.open?.(secondRawPeer as unknown as CrossWsPeer) + server.peers.get('crossws-peer')?.send('reply') + + expect(firstRawPeer.close).toHaveBeenCalledOnce() + expect(firstRawPeer.send).not.toHaveBeenCalled() + expect(secondRawPeer.send).toHaveBeenCalledExactlyOnceWith('reply', { compress: undefined }) + expect(server.peers.has('crossws-peer')).toBe(true) + }) +}) + +describe('better-ws client runtime', () => { + it('applies reconnectRandomFactor to scheduled reconnect delay', async () => { + vi.useFakeTimers() + vi.spyOn(Math, 'random').mockReturnValue(1) + FakeWebSocket.instances.length = 0 + const client = betterWs.createClient({ + url: 'ws://localhost/ws', + wsConstructor: FakeWebSocket, + reconnect: { + retries: 1, + delay: 1000, + reconnectRandomFactor: 0.5, + }, + }) + + void client.connect() + const first = FakeWebSocket.instances.at(-1)! + first.open() + first.closeEvent() + + await vi.advanceTimersByTimeAsync(1499) + expect(FakeWebSocket.instances).toHaveLength(1) + await vi.advanceTimersByTimeAsync(1) + expect(FakeWebSocket.instances).toHaveLength(2) + + vi.useRealTimers() + vi.restoreAllMocks() + }) + + it('does not reset reconnect attempt before reconnectMinConnectedDuration', async () => { + vi.useFakeTimers() + FakeWebSocket.instances.length = 0 + const client = betterWs.createClient({ + url: 'ws://localhost/ws', + wsConstructor: FakeWebSocket, + reconnect: { + retries: 2, + delay: attempt => attempt * 100, + reconnectMinConnectedDuration: 1000, + }, + }) + + void client.connect() + FakeWebSocket.instances.at(-1)!.open() + FakeWebSocket.instances.at(-1)!.closeEvent() + await vi.advanceTimersByTimeAsync(100) + FakeWebSocket.instances.at(-1)!.open() + FakeWebSocket.instances.at(-1)!.closeEvent() + await vi.advanceTimersByTimeAsync(199) + expect(FakeWebSocket.instances).toHaveLength(2) + await vi.advanceTimersByTimeAsync(1) + expect(FakeWebSocket.instances).toHaveLength(3) + + vi.useRealTimers() + }) + it('clamps reconnectRandomFactor so positive delays do not collapse to zero', async () => { + vi.spyOn(Math, 'random').mockReturnValue(0) + const reconnectDelays: number[] = [] + const client = betterWs.createClient({ + url: 'ws://localhost/ws', + wsConstructor: FakeWebSocket, + reconnect: { + retries: 1, + delay: 1000, + reconnectRandomFactor: 2, + }, + schedule: (delay, run) => { + reconnectDelays.push(delay) + return { cancel: vi.fn(), run } + }, + }) + + void client.connect() + FakeWebSocket.instances.at(-1)!.open() + FakeWebSocket.instances.at(-1)!.closeEvent() + + expect(reconnectDelays).toEqual([1]) + vi.restoreAllMocks() + }) + + it('resets reconnect attempt after reconnectMinConnectedDuration', async () => { + vi.useFakeTimers() + FakeWebSocket.instances.length = 0 + const client = betterWs.createClient({ + url: 'ws://localhost/ws', + wsConstructor: FakeWebSocket, + reconnect: { + retries: 2, + delay: attempt => attempt * 100, + reconnectMinConnectedDuration: 1000, + }, + }) + + void client.connect() + FakeWebSocket.instances.at(-1)!.open() + FakeWebSocket.instances.at(-1)!.closeEvent() + await vi.advanceTimersByTimeAsync(100) + FakeWebSocket.instances.at(-1)!.open() + await vi.advanceTimersByTimeAsync(1000) + FakeWebSocket.instances.at(-1)!.closeEvent() + await vi.advanceTimersByTimeAsync(99) + expect(FakeWebSocket.instances).toHaveLength(2) + await vi.advanceTimersByTimeAsync(1) + expect(FakeWebSocket.instances).toHaveLength(3) + + vi.useRealTimers() + }) + + it('resets reconnect attempt when prepare fails after reconnectMinConnectedDuration', async () => { + vi.useFakeTimers() + const reconnectDelays: number[] = [] + const scheduled: Array<() => void> = [] + const client = betterWs.createClient({ + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: vi.fn(), + }), + }, + reconnect: { + retries: 3, + delay: attempt => attempt * 100, + reconnectMinConnectedDuration: 1000, + }, + schedule: (delay, run) => { + reconnectDelays.push(delay) + scheduled.push(run) + return { cancel: vi.fn(), run } + }, + prepare: async ({ attempt }) => { + if (attempt === 0) { + throw new Error('initial prepare failed') + } + + await new Promise(resolve => setTimeout(resolve, 1000)) + throw new Error('retry prepare failed') + }, + }) + + await expect(client.connect()).rejects.toThrow('initial prepare failed') + expect(reconnectDelays).toEqual([100]) + + scheduled[0]?.() + await vi.advanceTimersByTimeAsync(1000) + + expect(reconnectDelays).toEqual([100, 100]) + vi.useRealTimers() + }) + + it('creates a text client from a native WebSocket constructor', async () => { + const fake = createFakeSocketClient() + const { client } = fake + const received: string[] = [] + client.onMessage(({ message }) => { + received.push(message) + }) + + const pendingConnect = client.connect() + const ws = fake.socket + ws.open() + await pendingConnect + client.send('hello') + ws.receive('from-server') + + expect(client.state).toBe('ready') + expect(ws.url).toBe('ws://localhost/ws') + expect(ws.sent).toEqual(['hello']) + expect(received).toEqual(['from-server']) + }) + + it('moves a url client to ready when no prepare procedure is provided', async () => { + const fake = createFakeSocketClient() + const { client } = fake + + const states: string[] = [] + client.onStateChange(({ state }) => states.push(state)) + + const connecting = client.connect() + fake.socket.open() + await connecting + + expect(client.state).toBe('ready') + expect(states).toEqual(['connecting', 'open', 'ready']) + }) + + it('rejects a url client open error without scheduling reconnect when reconnect is disabled', async () => { + const reconnectDelays: number[] = [] + const client = betterWs.createClient({ + url: 'ws://localhost/ws', + wsConstructor: FakeWebSocket, + reconnect: false, + schedule: (delay, run) => { + reconnectDelays.push(delay) + return { cancel: vi.fn(), run } + }, + }) + + const connecting = client.connect() + const socket = FakeWebSocket.instances.at(-1) + socket?.error() + socket?.closeEvent() + + await expect(connecting).rejects.toThrow('WebSocket connection failed before opening.') + + expect(reconnectDelays).toEqual([]) + expect(client.state).toBe('closed') + }) + + it('schedules reconnect after an initial connector open failure', async () => { + const reconnectDelays: number[] = [] + const scheduled: Array<() => void> = [] + let attempt = 0 + const client = betterWs.createClient({ + connector: { + connect() { + attempt += 1 + if (attempt === 1) { + throw new Error('server unavailable') + } + + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + reconnect: { retries: 1, delay: attempt => attempt }, + schedule: (delay, run) => { + reconnectDelays.push(delay) + scheduled.push(run) + return { cancel: vi.fn(), run } + }, + }) + + await expect(client.connect()).rejects.toThrow('server unavailable') + + expect(client.state).toBe('reconnecting') + expect(reconnectDelays).toEqual([1]) + + scheduled[0]?.() + await Promise.resolve() + + expect(client.state).toBe('ready') + }) + + it('enters ready only after prepare resolves', async () => { + let serverMessage: ((message: string) => void) | undefined + const sent: string[] = [] + const states: string[] = [] + const client = betterWs.createClient({ + connector: { + connect(events) { + serverMessage = events.message + return { + send: (message) => { + sent.push(message) + return true + }, + close: vi.fn(), + } + }, + }, + async prepare(ctx) { + ctx.send('auth') + await ctx.waitFor(message => message === 'authenticated', { timeout: 100 }) + ctx.send('announce') + await ctx.waitFor(message => message === 'announced', { timeout: 100 }) + }, + }) + client.onStateChange(({ state }) => states.push(state)) + + const connecting = client.connect() + await Promise.resolve() + serverMessage?.('authenticated') + await Promise.resolve() + serverMessage?.('announced') + await connecting + + expect(client.state).toBe('ready') + expect(sent).toEqual(['auth', 'announce']) + expect(states).toEqual(['connecting', 'open', 'preparing', 'ready']) + }) + + it('rejects prepare when waitFor times out', async () => { + const closed = vi.fn() + const client = betterWs.createClient({ + reconnect: false, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: closed, + }), + }, + async prepare(ctx) { + await ctx.waitFor(message => message === 'never', { timeout: 1 }) + }, + }) + + await expect(client.connect()).rejects.toThrow('Timed out waiting for message.') + + expect(client.state).toBe('failed') + expect(closed).toHaveBeenCalledOnce() + }) + + it('aborts prepare when the client closes while waiting for a message', async () => { + let prepareError: unknown + const closed = vi.fn() + const client = betterWs.createClient({ + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: closed, + }), + }, + async prepare(ctx) { + try { + await ctx.waitFor(message => message === 'ready', { timeout: 100 }) + } + catch (error) { + prepareError = error + throw error + } + }, + }) + + const connecting = client.connect() + await Promise.resolve() + client.close() + await connecting + + expect(client.state).toBe('closed') + expect(closed).toHaveBeenCalledOnce() + expect(prepareError).toBeInstanceOf(Error) + }) + + it('does not enter ready when the transport closes during prepare', async () => { + let closeTransport: (() => void) | undefined + let resolvePrepare: (() => void) | undefined + const states: string[] = [] + const client = betterWs.createClient({ + reconnect: false, + connector: { + connect: (events) => { + closeTransport = () => events.close({ code: 1006, reason: 'lost' }) + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + async prepare() { + await new Promise((resolve) => { + resolvePrepare = resolve + }) + }, + }) + client.onStateChange(({ state }) => states.push(state)) + + const connecting = client.connect() + await Promise.resolve() + closeTransport?.() + resolvePrepare?.() + await connecting + + expect(client.state).toBe('closed') + expect(states).not.toContain('ready') + }) + + it('aborts stale prepare when a newer connect starts', async () => { + const closed = [vi.fn(), vi.fn()] + const prepareErrors: unknown[] = [] + const prepareSignals: AbortSignal[] = [] + let connectCount = 0 + let serverMessage: ((message: string) => void) | undefined + const client = betterWs.createClient({ + connector: { + connect(events) { + const connectionIndex = connectCount++ + serverMessage = events.message + return { + send: vi.fn(() => true), + close: closed[connectionIndex], + } + }, + }, + async prepare(ctx) { + prepareSignals.push(ctx.signal) + try { + await ctx.waitFor(message => message === `ready:${prepareSignals.length}`, { timeout: 100 }) + } + catch (error) { + prepareErrors.push(error) + throw error + } + }, + }) + + const firstConnect = client.connect() + await Promise.resolve() + const secondConnect = client.connect() + await Promise.resolve() + serverMessage?.('ready:2') + await Promise.all([firstConnect, secondConnect]) + + expect(prepareSignals[0]?.aborted).toBe(true) + expect(prepareErrors[0]).toBeInstanceOf(Error) + expect(client.state).toBe('ready') + expect(closed[0]).toHaveBeenCalledOnce() + }) + + it('keeps stale prepare waitFor bound to its aborted prepare context', async () => { + const closed = [vi.fn(), vi.fn()] + let connectCount = 0 + let serverMessage: ((message: string) => void) | undefined + let firstWaitFor: ((message: string) => Promise) | undefined + let firstWaitError: unknown + let firstWaitResult: string | undefined + const client = betterWs.createClient({ + connector: { + connect(events) { + const connectionIndex = connectCount++ + serverMessage = events.message + return { + send: vi.fn(() => true), + close: closed[connectionIndex], + } + }, + }, + async prepare(ctx) { + if (!firstWaitFor) { + firstWaitFor = (expected: string) => ctx.waitFor(message => message === expected, { timeout: 100 }) + return + } + + await ctx.waitFor(message => message === 'second-ready', { timeout: 100 }) + }, + }) + + const firstConnect = client.connect() + await Promise.resolve() + const secondConnect = client.connect() + await Promise.resolve() + const staleWait = firstWaitFor?.('second-ready') + .then((message) => { + firstWaitResult = message + }) + .catch((error: unknown) => { + firstWaitError = error + }) + serverMessage?.('second-ready') + await secondConnect + await staleWait + await firstConnect + + expect(client.state).toBe('ready') + expect(firstWaitResult).toBeUndefined() + expect(firstWaitError).toBeInstanceOf(Error) + expect((firstWaitError as Error).message).toBe('Wait for message aborted.') + expect(closed[0]).toHaveBeenCalledOnce() + }) + + it('does not stay preparing when prepare fails and reconnect is enabled', async () => { + const reconnects: number[] = [] + const closed = vi.fn() + const client = betterWs.createClient({ + reconnect: { retries: 1, delay: attempt => attempt }, + schedule: (delay, run) => { + reconnects.push(delay) + return { cancel: vi.fn(), run } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: closed, + }), + }, + async prepare() { + throw new Error('prepare failed') + }, + }) + + await expect(client.connect()).rejects.toThrow('prepare failed') + + expect(client.state).toBe('reconnecting') + expect(reconnects).toEqual([1]) + expect(closed).toHaveBeenCalledOnce() + }) + + it('passes reconnect attempt metadata to prepare after a scheduled retry', async () => { + let scheduledRun: (() => void) | undefined + const prepareAttempts: Array<{ attempt: number, reconnecting: boolean }> = [] + const client = betterWs.createClient({ + reconnect: { retries: 1, delay: attempt => attempt }, + schedule: (_delay, run) => { + scheduledRun = run + return { cancel: vi.fn() } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: vi.fn(), + }), + }, + async prepare(ctx) { + prepareAttempts.push({ attempt: ctx.attempt, reconnecting: ctx.reconnecting }) + if (prepareAttempts.length === 1) { + throw new Error('prepare failed') + } + }, + }) + + await expect(client.connect()).rejects.toThrow('prepare failed') + scheduledRun?.() + await Promise.resolve() + await Promise.resolve() + + expect(client.state).toBe('ready') + expect(prepareAttempts).toEqual([ + { attempt: 0, reconnecting: false }, + { attempt: 1, reconnecting: true }, + ]) + }) + + it('does not schedule duplicate reconnect when failed prepare connection later closes', async () => { + const reconnects: number[] = [] + const closed = vi.fn() + let closeTransport: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { retries: 2, delay: attempt => attempt }, + schedule: (delay, run) => { + reconnects.push(delay) + return { cancel: vi.fn(), run } + }, + connector: { + connect: (events) => { + closeTransport = () => events.close({ code: 1006, reason: 'late close' }) + return { + send: vi.fn(() => true), + close: closed, + } + }, + }, + async prepare() { + throw new Error('prepare failed') + }, + }) + + await expect(client.connect()).rejects.toThrow('prepare failed') + closeTransport?.() + + expect(client.state).toBe('reconnecting') + expect(reconnects).toEqual([1]) + expect(closed).toHaveBeenCalledOnce() + }) + + it('blocks normal send before ready', async () => { + const sent: string[] = [] + let resolveConnection: ((connection: { send: (message: string) => void, close: () => void }) => void) | undefined + const client = betterWs.createClient({ + connector: { + connect: () => new Promise((resolve) => { + resolveConnection = resolve + }), + }, + }) + + const connecting = client.connect() + + expect(client.send('before-open')).toEqual({ ok: false, reason: 'closed' }) + + resolveConnection?.({ + send: message => sent.push(message), + close: vi.fn(), + }) + await connecting + + expect(client.send('ready')).toEqual({ ok: true }) + expect(sent).toEqual(['ready']) + }) + + it('allows send during open before ready when requireReady is false', async () => { + const sent: string[] = [] + let resolveConnection: ((connection: { send: (message: string) => void, close: () => void }) => void) | undefined + const client = betterWs.createClient({ + connector: { + connect: () => new Promise((resolve) => { + resolveConnection = resolve + }), + }, + }) + + client.onStateChange(({ state }) => { + if (state === 'open') { + expect(client.send('before-ready', { requireReady: false })).toEqual({ ok: true }) + } + }) + + const connecting = client.connect() + resolveConnection?.({ + send: message => sent.push(message), + close: vi.fn(), + }) + await connecting + + expect(client.state).toBe('ready') + expect(sent).toEqual(['before-ready']) + }) + + it('connects through an adapter, dispatches messages, and reports send results', async () => { + const sent: string[] = [] + let adapterMessage: ((message: string) => void) | undefined + const client = betterWs.createClient({ + connector: { + connect: async ({ message }) => { + adapterMessage = message + return { + send: (nextMessage) => { + sent.push(nextMessage) + return true + }, + close: vi.fn(), + } + }, + }, + }) + const received: string[] = [] + client.onMessage(({ message }) => { + received.push(message) + }) + + await client.connect() + const result = client.send('hello') + adapterMessage?.('from-server') + + expect(client.state).toBe('ready') + expect(result).toEqual({ ok: true }) + expect(sent).toEqual(['hello']) + expect(received).toEqual(['from-server']) + }) + + it('sends message heartbeat and keeps the connection alive after a response', async () => { + const sent: string[] = [] + let serverMessage: ((message: string) => void) | undefined + const cancelHeartbeatTimeout = vi.fn() + const scheduled: Array<{ delay: number, run: () => void }> = [] + const client = betterWs.createClient({ + heartbeat: { + mode: 'message', + interval: 10, + timeout: 20, + message: 'ping', + isResponse: message => message === 'pong', + }, + schedule: (_delay, run) => { + scheduled.push({ delay: _delay, run }) + return { cancel: scheduled.length === 2 ? cancelHeartbeatTimeout : vi.fn() } + }, + connector: { + connect: (events) => { + serverMessage = events.message + return { + send: (message) => { + sent.push(message) + return true + }, + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + scheduled[0]?.run() + serverMessage?.('pong') + + expect(sent).toEqual(['ping']) + expect(cancelHeartbeatTimeout).toHaveBeenCalledOnce() + expect(client.state).toBe('ready') + }) + + it('does not let strict non-response messages defer heartbeat timeout', async () => { + const scheduled: Array<{ delay: number, run: () => void, cancel: ReturnType }> = [] + const reconnectErrors: unknown[] = [] + let serverMessage: ((message: string) => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: 1, + delay: (_attempt, error) => { + reconnectErrors.push(error) + return 10 + }, + }, + heartbeat: { + mode: 'message', + interval: 1, + timeout: 5, + message: 'ping', + isResponse: message => message === 'pong', + }, + schedule: (delay, run) => { + const task = { delay, run, cancel: vi.fn() } + scheduled.push(task) + return task + }, + connector: { + connect: (events) => { + serverMessage = events.message + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + scheduled[0]?.run() + serverMessage?.('not-pong') + + expect(scheduled).toHaveLength(2) + expect(scheduled[1]?.delay).toBe(5) + expect(scheduled[1]?.cancel).not.toHaveBeenCalled() + + scheduled[1]?.run() + + expect(client.state).toBe('reconnecting') + expect(reconnectErrors).toHaveLength(1) + expect(reconnectErrors[0]).toBeInstanceOf(Error) + expect((reconnectErrors[0] as Error).message).toBe('Heartbeat timed out after 5ms.') + }) + + it('clears pending heartbeat timeout on any inbound message when no response predicate is provided', async () => { + const cancelHeartbeatTimeout = vi.fn() + const scheduled: Array<{ delay: number, run: () => void }> = [] + let serverMessage: ((message: string) => void) | undefined + const client = betterWs.createClient({ + heartbeat: { + mode: 'message', + interval: 1, + timeout: 5, + message: 'ping', + }, + schedule: (delay, run) => { + scheduled.push({ delay, run }) + return { cancel: scheduled.length === 2 ? cancelHeartbeatTimeout : vi.fn() } + }, + connector: { + connect: (events) => { + serverMessage = events.message + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + scheduled[0]?.run() + serverMessage?.('any-message') + + expect(cancelHeartbeatTimeout).toHaveBeenCalledOnce() + expect(scheduled).toHaveLength(3) + expect(scheduled[2]?.delay).toBe(1) + }) + + it('reconnects when heartbeat response times out', async () => { + const scheduled: Array<{ label: string, delay: number, run: () => void }> = [] + const closed = vi.fn() + const client = betterWs.createClient({ + reconnect: { retries: 1, delay: 5 }, + heartbeat: { + mode: 'message', + interval: 1, + timeout: 1, + message: 'ping', + isResponse: message => message === 'pong', + }, + schedule: (delay, run) => { + const label = scheduled.length === 0 + ? 'heartbeat interval' + : scheduled.length === 1 + ? 'heartbeat timeout' + : 'reconnect delay' + scheduled.push({ label, delay, run }) + return { cancel: vi.fn() } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: closed, + }), + }, + }) + + await client.connect() + scheduled.find(task => task.label === 'heartbeat interval')?.run() + scheduled.find(task => task.label === 'heartbeat timeout')?.run() + + expect(closed).toHaveBeenCalledOnce() + expect(client.state).toBe('reconnecting') + expect(scheduled).toMatchObject([ + { label: 'heartbeat interval', delay: 1 }, + { label: 'heartbeat timeout', delay: 1 }, + { label: 'reconnect delay', delay: 5 }, + ]) + }) + + it('schedules one reconnect when heartbeat timeout close emits synchronously', async () => { + const reconnects: number[] = [] + const scheduled: Array<{ label: string, run: () => void }> = [] + const client = betterWs.createClient({ + reconnect: { retries: 2, delay: attempt => attempt }, + heartbeat: { + mode: 'message', + interval: 1, + timeout: 1, + message: 'ping', + isResponse: message => message === 'pong', + }, + schedule: (delay, run) => { + if (scheduled.length < 2) { + scheduled.push({ + label: scheduled.length === 0 ? 'heartbeat interval' : 'heartbeat timeout', + run, + }) + } + else { + reconnects.push(delay) + } + return { cancel: vi.fn() } + }, + connector: { + connect: ({ close }) => ({ + send: vi.fn(() => true), + close: () => close({ code: 1006, reason: 'heartbeat timeout' }), + }), + }, + }) + + await client.connect() + scheduled.find(task => task.label === 'heartbeat interval')?.run() + scheduled.find(task => task.label === 'heartbeat timeout')?.run() + + expect(client.state).toBe('reconnecting') + expect(reconnects).toEqual([1]) + }) + + it('uses native ping for automatic heartbeat when the connection exposes ping', async () => { + const ping = vi.fn(() => true) + let scheduled: (() => void) | undefined + const client = betterWs.createClient({ + heartbeat: { + mode: 'auto', + interval: 10, + }, + schedule: (_delay, run) => { + scheduled = run + return { cancel: vi.fn() } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + ping, + close: vi.fn(), + }), + }, + }) + + await client.connect() + scheduled?.() + + expect(ping).toHaveBeenCalledOnce() + expect(client.state).toBe('ready') + }) + + it('uses message heartbeat in auto mode when native ping is unavailable and message exists', async () => { + const sent: string[] = [] + let scheduled: (() => void) | undefined + const client = betterWs.createClient({ + heartbeat: { + mode: 'auto', + interval: 10, + message: 'ping', + }, + schedule: (_delay, run) => { + scheduled = run + return { cancel: vi.fn() } + }, + connector: { + connect: () => ({ + send: (message) => { + sent.push(message) + return true + }, + close: vi.fn(), + }), + }, + }) + + await client.connect() + scheduled?.() + + expect(sent).toEqual(['ping']) + expect(client.state).toBe('ready') + }) + + it('fails coherently when native heartbeat has no ping support', async () => { + const reconnectErrors: unknown[] = [] + let scheduled: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: 1, + delay: (_attempt, error) => { + reconnectErrors.push(error) + return 10 + }, + }, + heartbeat: { + mode: 'native', + interval: 1, + }, + schedule: (_delay, run) => { + scheduled = run + return { cancel: vi.fn() } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: vi.fn(), + }), + }, + }) + + await client.connect() + scheduled?.() + + expect(client.state).toBe('reconnecting') + expect(reconnectErrors).toHaveLength(1) + expect((reconnectErrors[0] as Error).message).toBe('Native heartbeat requires connection.ping().') + }) + + it('fails coherently when message heartbeat has no configured message', async () => { + const reconnectErrors: unknown[] = [] + let scheduled: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: 1, + delay: (_attempt, error) => { + reconnectErrors.push(error) + return 10 + }, + }, + heartbeat: { + mode: 'message', + interval: 1, + }, + schedule: (_delay, run) => { + scheduled = run + return { cancel: vi.fn() } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: vi.fn(), + }), + }, + }) + + await client.connect() + scheduled?.() + + expect(client.state).toBe('reconnecting') + expect(reconnectErrors).toHaveLength(1) + expect((reconnectErrors[0] as Error).message).toBe('Message heartbeat requires heartbeat.message.') + }) + + it('cancels heartbeat tasks on manual close', async () => { + const cancelHeartbeatInterval = vi.fn() + const cancelHeartbeatTimeout = vi.fn() + let scheduledCount = 0 + const closed = vi.fn() + const client = betterWs.createClient({ + heartbeat: { + mode: 'message', + interval: 1, + timeout: 1, + message: 'ping', + isResponse: message => message === 'pong', + }, + schedule: (_delay, run) => { + scheduledCount += 1 + return { + cancel: scheduledCount === 1 ? cancelHeartbeatInterval : cancelHeartbeatTimeout, + run, + } + }, + connector: { + connect: () => ({ + send: vi.fn(() => true), + close: closed, + }), + }, + }) + + await client.connect() + client.close() + + expect(closed).toHaveBeenCalledOnce() + expect(cancelHeartbeatInterval).toHaveBeenCalledOnce() + expect(cancelHeartbeatTimeout).not.toHaveBeenCalled() + expect(client.state).toBe('closed') + }) + + it('ignores stale heartbeat timeouts after a newer connection becomes active', async () => { + const scheduled: Array<{ delay: number, run: () => void }> = [] + const closed = [vi.fn(), vi.fn()] + let connectCount = 0 + const client = betterWs.createClient({ + reconnect: { retries: 1, delay: 5 }, + heartbeat: { + mode: 'message', + interval: 1, + timeout: 1, + message: 'ping', + isResponse: message => message === 'pong', + }, + schedule: (delay, run) => { + scheduled.push({ delay, run }) + return { cancel: vi.fn() } + }, + connector: { + connect: () => { + const connectionIndex = connectCount++ + return { + send: vi.fn(() => true), + close: closed[connectionIndex], + } + }, + }, + }) + + await client.connect() + scheduled[0]?.run() + await client.connect() + scheduled[1]?.run() + + expect(closed[0]).toHaveBeenCalledOnce() + expect(closed[1]).not.toHaveBeenCalled() + expect(client.state).toBe('ready') + expect(scheduled).toHaveLength(3) + }) + + it('schedules reconnect after unexpected close', async () => { + const reconnects: number[] = [] + let closeHandler: ((details?: { code?: number, reason?: string }) => void) | undefined + const client = betterWs.createClient({ + reconnect: { retries: 2, delay: attempt => attempt * 10 }, + schedule: (delay, run) => { + reconnects.push(delay) + return { cancel: vi.fn(), run } + }, + connector: { + connect: async ({ close }) => { + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + closeHandler?.({ code: 1006, reason: 'network' }) + + expect(client.state).toBe('reconnecting') + expect(reconnects).toEqual([10]) + }) + + it('closes the active connection when a post-open adapter error schedules reconnect', async () => { + const reconnects: number[] = [] + let emitError: ((error: unknown) => void) | undefined + const close = vi.fn() + const client = betterWs.createClient({ + connector: { + connect(events) { + emitError = events.error + return { + send: vi.fn(() => true), + close, + } + }, + }, + reconnect: { retries: 1, delay: attempt => attempt }, + schedule: (delay, run) => { + reconnects.push(delay) + return { cancel: vi.fn(), run } + }, + }) + + await client.connect() + emitError?.(new Error('invalid payload')) + + expect(close).toHaveBeenCalledOnce() + expect(client.state).toBe('reconnecting') + expect(reconnects).toEqual([1]) + }) + + it('reconnects by default after an unexpected close', async () => { + const reconnectDelays: number[] = [] + let closeHandler: (() => void) | undefined + let connects = 0 + const client = betterWs.createClient({ + schedule: (delay, run) => { + reconnectDelays.push(delay) + return { cancel: vi.fn(), run } + }, + connector: { + connect: ({ close }) => { + connects += 1 + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + closeHandler?.() + + expect(client.state).toBe('reconnecting') + expect(connects).toBe(1) + expect(reconnectDelays).toEqual([1000]) + }) + + it('calls onFailed when retry predicate stops reconnecting', async () => { + const onFailed = vi.fn() + let closeHandler: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: attempt => attempt < 1, + delay: 1, + onFailed, + }, + connector: { + connect: ({ close }) => { + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + closeHandler?.() + await Promise.resolve() + + expect(client.state).toBe('failed') + expect(onFailed).toHaveBeenCalledOnce() + }) + + it('continues reconnect policy when an automatic reconnect fails to open', async () => { + const reconnectDelays: number[] = [] + const scheduledRuns: Array<() => void> = [] + let closeHandler: (() => void) | undefined + let connects = 0 + const client = betterWs.createClient({ + reconnect: { retries: 2, delay: attempt => attempt }, + schedule: (delay, run) => { + reconnectDelays.push(delay) + scheduledRuns.push(run) + return { cancel: vi.fn(), run } + }, + connector: { + connect: ({ close }) => { + connects += 1 + if (connects === 2) { + throw new Error('automatic reconnect failed') + } + + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + closeHandler?.() + scheduledRuns[0]?.() + await Promise.resolve() + await Promise.resolve() + + expect(connects).toBe(2) + expect(client.state).toBe('reconnecting') + expect(reconnectDelays).toEqual([1, 2]) + }) + + it('keeps automatic reconnect dialing in reconnecting state', async () => { + const states: string[] = [] + const scheduledRuns: Array<() => void> = [] + let closeHandler: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { retries: 1, delay: 1 }, + schedule: (_delay, run) => { + scheduledRuns.push(run) + return { cancel: vi.fn(), run } + }, + connector: { + connect: ({ close }) => { + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + client.onStateChange(({ state }) => states.push(state)) + + await client.connect() + + expect(states).toContain('connecting') + + states.length = 0 + closeHandler?.() + scheduledRuns[0]?.() + await Promise.resolve() + await Promise.resolve() + + expect(states).not.toContain('connecting') + expect(states).toEqual(['reconnecting', 'open', 'ready']) + expect(client.state).toBe('ready') + }) + + it('resets close error after a successful reconnect', async () => { + const firstError = new Error('first connection failed') + const reconnectErrors: unknown[] = [] + const scheduledRuns: Array<() => void> = [] + const closeHandlers: Array<() => void> = [] + const errorHandlers: Array<(error: unknown) => void> = [] + const client = betterWs.createClient({ + reconnect: { + retries: 2, + delay: (_attempt, error) => { + reconnectErrors.push(error) + return 1 + }, + }, + schedule: (_delay, run) => { + scheduledRuns.push(run) + return { cancel: vi.fn(), run } + }, + connector: { + connect: ({ close, error }) => { + closeHandlers.push(close) + errorHandlers.push(error) + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + errorHandlers[0]?.(firstError) + scheduledRuns[0]?.() + await Promise.resolve() + await Promise.resolve() + closeHandlers[1]?.() + + expect(client.state).toBe('reconnecting') + expect(reconnectErrors).toHaveLength(2) + expect(reconnectErrors[0]).toBe(firstError) + expect(reconnectErrors[1]).toBeInstanceOf(Error) + expect((reconnectErrors[1] as Error).message).toBe('Connection closed') + expect(reconnectErrors[1]).not.toBe(firstError) + }) + + it('fails coherently when retry predicate throws', async () => { + const policyError = new Error('retry predicate failed') + const onFailed = vi.fn() + let closeHandler: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: () => { + throw policyError + }, + delay: 1, + onFailed, + }, + connector: { + connect: ({ close }) => { + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + closeHandler?.() + + expect(client.state).toBe('failed') + expect(onFailed).toHaveBeenCalledOnce() + expect(onFailed).toHaveBeenCalledWith(policyError) + }) + + it('fails coherently when reconnect delay resolver throws', async () => { + const policyError = new Error('delay resolver failed') + const onFailed = vi.fn() + let closeHandler: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: 1, + delay: () => { + throw policyError + }, + onFailed, + }, + connector: { + connect: ({ close }) => { + closeHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + await client.connect() + closeHandler?.() + + expect(client.state).toBe('failed') + expect(onFailed).toHaveBeenCalledOnce() + expect(onFailed).toHaveBeenCalledWith(policyError) + }) + + it('ignores stale error events when resolving the current reconnect delay', async () => { + const staleError = new Error('stale connection failed') + const reconnectErrors: unknown[] = [] + const connectionErrors: Array<(error: unknown) => void> = [] + let currentCloseHandler: (() => void) | undefined + const client = betterWs.createClient({ + reconnect: { + retries: 1, + delay: (_attempt, error) => { + reconnectErrors.push(error) + return 1 + }, + }, + schedule: (_delay, run) => ({ cancel: vi.fn(), run }), + connector: { + connect: ({ close, error }) => { + connectionErrors.push(error) + currentCloseHandler = close + return { + send: vi.fn(() => true), + close: vi.fn(), + } + }, + }, + }) + + // ROOT CAUSE: + // + // Stale error events used to assign `lastCloseError` before the epoch + // guard in close handling could reject them. A later current close then + // scheduled reconnect using the stale connection's error. + await client.connect() + await client.connect() + connectionErrors[0]?.(staleError) + currentCloseHandler?.() + + expect(client.state).toBe('reconnecting') + expect(reconnectErrors).toHaveLength(1) + expect(reconnectErrors[0]).toBeInstanceOf(Error) + expect((reconnectErrors[0] as Error).message).toBe('Connection closed') + }) + + it('schedules one reconnect when prepare failure close emits synchronously', async () => { + const reconnects: number[] = [] + const client = betterWs.createClient({ + reconnect: { retries: 2, delay: attempt => attempt }, + schedule: (delay, run) => { + reconnects.push(delay) + return { cancel: vi.fn(), run } + }, + connector: { + connect: ({ close }) => ({ + send: vi.fn(() => true), + close: () => close({ code: 1006, reason: 'prepare close' }), + }), + }, + async prepare() { + throw new Error('prepare failed') + }, + }) + + // ROOT CAUSE: + // + // Prepare failure used to call the adapter close hook before invalidating + // the connection epoch. If that close emitted synchronously, both the close + // path and prepare catch scheduled reconnect. + await expect(client.connect()).rejects.toThrow('prepare failed') + + expect(client.state).toBe('reconnecting') + expect(reconnects).toEqual([1]) + }) + + it('returns to closed when a connector fails to open with reconnect disabled', async () => { + const reconnectDelays: number[] = [] + const client = betterWs.createClient({ + reconnect: false, + schedule: (delay, run) => { + reconnectDelays.push(delay) + return { cancel: vi.fn(), run } + }, + connector: { + connect: async () => { + throw new Error('connect failed') + }, + }, + }) + + await expect(client.connect()).rejects.toThrow('connect failed') + + expect(client.state).toBe('closed') + expect(reconnectDelays).toEqual([]) + }) + + it('keeps a client closed when a pending connect resolves after manual close', async () => { + let resolveConnection: ((connection: { send: (message: string) => boolean, close: () => void }) => void) | undefined + const connection = { + send: vi.fn(() => true), + close: vi.fn(), + } + const client = betterWs.createClient({ + connector: { + connect: () => new Promise((resolve) => { + resolveConnection = resolve + }), + }, + }) + + const pendingConnect = client.connect() + client.close() + resolveConnection?.(connection) + await pendingConnect + + expect(client.state).toBe('closed') + expect(client.send('late')).toEqual({ ok: false, reason: 'closed' }) + expect(connection.close).toHaveBeenCalledOnce() + }) +}) + +describe('better-ws server control messages', () => { + it('routes heartbeat control messages to onPing and onPong without business onMessage', () => { + type Message = { type: 'ping' } | { type: 'pong' } | { type: 'data', value: string } + const server = createServer({ + heartbeat: { + message: () => ({ type: 'ping' }), + isPing: message => message.type === 'ping', + isPong: message => message.type === 'pong', + }, + }) + const messages: Message[] = [] + const pings: string[] = [] + const pongs: string[] = [] + server.onMessage(({ message }) => { + messages.push(message) + }) + server.onPing(({ peer }) => { + pings.push(peer.id) + }) + server.onPong(({ peer }) => { + pongs.push(peer.id) + }) + + const peer = server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }).peer + peer.receive({ type: 'ping' }) + peer.receive({ type: 'pong' }) + peer.receive({ type: 'data', value: 'hello' }) + + expect(pings).toEqual(['peer-1']) + expect(pongs).toEqual(['peer-1']) + expect(messages).toEqual([{ type: 'data', value: 'hello' }]) + }) +}) + +describe('better-ws server lifecycle procedures', () => { + it('passes peers and previous snapshot to peer open handlers', () => { + const server = createServer() + const opens: Array<{ id: string, previousState?: { ready: boolean }, peerCount: number }> = [] + + server.onPeerOpen(({ peer, previous, peers }) => { + opens.push({ + id: peer.id, + previousState: previous?.state, + peerCount: peers.list().length, + }) + }) + + server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }, { state: { ready: false } }) + server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }) + + expect(opens).toEqual([ + { id: 'peer-1', previousState: undefined, peerCount: 1 }, + { id: 'peer-1', previousState: { ready: false }, peerCount: 1 }, + ]) + }) + + it('emits peer close only when a peer is actually removed', () => { + const server = createServer() + const closed: string[] = [] + server.onPeerClose(({ peerId }) => { + closed.push(peerId) + }) + + server.peers.remove('missing') + const peer = server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }).peer + server.peers.remove(peer.id) + server.peers.remove(peer.id) + + expect(closed).toEqual(['peer-1']) + }) + + it('emits peer close for direct close, manager close, liveness, and replacement', () => { + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 10, + }, + heartbeat: { + timeout: 10, + }, + }) + const closed: Array<{ peerId: string, code?: number, reason?: string }> = [] + server.onPeerClose(({ peerId, details }) => { + closed.push({ peerId, code: details?.code, reason: details?.reason }) + }) + + const direct = server.peers.accept({ id: 'direct', send: vi.fn(() => true), close: vi.fn() }).peer + direct.close(4000, 'direct close') + + server.peers.accept({ id: 'manager', send: vi.fn(() => true), close: vi.fn() }) + server.peers.close('manager', 4001, 'manager close') + + server.peers.accept({ id: 'liveness', send: vi.fn(() => true), close: vi.fn() }) + server.checkLiveness(Date.now() + 10) + + server.peers.accept({ id: 'replace', send: vi.fn(() => true), close: vi.fn() }) + server.peers.accept({ id: 'replace', send: vi.fn(() => true), close: vi.fn() }) + + expect(closed).toEqual([ + { peerId: 'direct', code: 4000, reason: 'direct close' }, + { peerId: 'manager', code: 4001, reason: 'manager close' }, + { peerId: 'liveness', code: undefined, reason: undefined }, + { peerId: 'replace', code: undefined, reason: undefined }, + ]) + }) + + it('rejects waitFor on timeout and external abort', async () => { + const timeoutServer = createServer() + const timeoutFailures: string[] = [] + timeoutServer.onPeerOpen(async (event) => { + try { + await event.procedure(async (ctx) => { + await ctx.waitFor(message => message === 'never', { timeout: 1 }) + }) + } + catch (error) { + timeoutFailures.push(errorText(error)) + } + }) + timeoutServer.peers.accept({ id: 'timeout', send: vi.fn(() => true) }) + + await vi.waitFor(() => { + expect(timeoutFailures).toEqual(['Procedure waitFor timed out.']) + }) + + const abortServer = createServer() + const controller = new AbortController() + const abortFailures: string[] = [] + abortServer.onPeerOpen(async (event) => { + try { + await event.procedure(async (ctx) => { + await ctx.waitFor(message => message === 'never', { signal: controller.signal }) + }) + } + catch (error) { + abortFailures.push(errorText(error)) + } + }) + abortServer.peers.accept({ id: 'abort', send: vi.fn(() => true) }) + controller.abort() + + await vi.waitFor(() => { + expect(abortFailures).toEqual(['Procedure aborted.']) + }) + }) + + it('does not let a replaced peer procedure receive the next peer messages', async () => { + const server = createServer() + const ready: string[] = [] + const failures: string[] = [] + + server.onPeerOpen(async (event) => { + try { + await event.procedure(async (ctx) => { + const message = await ctx.waitFor(message => message === 'auth-ok', { timeout: 5 }) + ready.push(`${ctx.peer.id}:${message}`) + }) + } + catch (error) { + failures.push(errorText(error)) + } + }) + + server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }) + const next = server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }).peer + next.receive('auth-ok') + + await vi.waitFor(() => { + expect(ready).toEqual(['peer-1:auth-ok']) + expect(failures).toEqual(['Procedure aborted.']) + }) + }) + + it('cleans up waitFor when the predicate throws', async () => { + const server = createServer() + const failures: string[] = [] + + server.onPeerOpen(async (event) => { + try { + await event.procedure(async (ctx) => { + await ctx.waitFor(() => { + throw new Error('predicate failed') + }) + }) + } + catch (error) { + failures.push(errorText(error)) + } + }) + + const peer = server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }).peer + peer.receive('first') + peer.receive('second') + + await vi.waitFor(() => { + expect(failures).toEqual(['predicate failed']) + }) + }) + + it('runs peer open procedure with waitFor and cleanup', async () => { + const server = createServer() + const ready: string[] = [] + + server.onPeerOpen(async (event) => { + await event.procedure(async (ctx) => { + const message = await ctx.waitFor(message => message === 'auth-ok', { timeout: 50 }) + ready.push(message) + }) + }) + + const peer = server.peers.accept({ id: 'peer-1', send: vi.fn(() => true) }).peer + peer.receive('auth-ok') + + await vi.waitFor(() => { + expect(ready).toEqual(['auth-ok']) + }) + }) +}) + +describe('better-ws integrated runtime', () => { + it('connects a client connector to a server peer and exchanges messages both ways', async () => { + const server = createServer() + const clientReceived: string[] = [] + const serverReceived: Array<{ peerId: string, message: string }> = [] + let serverPeerId: string | undefined + + server.onMessage(({ peer, message }) => { + serverReceived.push({ peerId: peer.id, message }) + peer.send(`ack:${message}`) + }) + + const client = betterWs.createClient({ + connector: { + connect: ({ message }) => { + const peer = server.accept({ + id: 'client-1', + send: (serverMessage) => { + message(serverMessage) + return true + }, + }) + serverPeerId = peer.id + + return { + send: clientMessage => peer.receive(clientMessage), + close: () => server.remove(peer.id), + } + }, + }, + }) + client.onMessage(({ message }) => { + clientReceived.push(message) + }) + + await client.connect() + const sendResult = client.send('hello') + + expect(sendResult).toEqual({ ok: true }) + expect(serverPeerId).toBe('client-1') + expect(server.peers.has('client-1')).toBe(true) + expect(serverReceived).toEqual([{ peerId: 'client-1', message: 'hello' }]) + expect(clientReceived).toEqual(['ack:hello']) + + client.close() + + expect(server.peers.has('client-1')).toBe(false) + }) +}) diff --git a/packages/better-ws/src/index.ts b/packages/better-ws/src/index.ts new file mode 100644 index 000000000..765c2df3e --- /dev/null +++ b/packages/better-ws/src/index.ts @@ -0,0 +1,2 @@ +export * from './client' +export type { WsCloseDetails, WsSendResult, WsState } from './shared' diff --git a/packages/better-ws/src/server/h3/index.test.ts b/packages/better-ws/src/server/h3/index.test.ts new file mode 100644 index 000000000..b185a268f --- /dev/null +++ b/packages/better-ws/src/server/h3/index.test.ts @@ -0,0 +1,55 @@ +import type { Hooks } from 'crossws' + +import { describe, expect, it, vi } from 'vitest' + +vi.mock('h3', () => ({ + defineWebSocketHandler: vi.fn(hooks => ({ kind: 'h3-handler', hooks })), +})) + +vi.mock('crossws/server', () => ({ + plugin: vi.fn(options => ({ kind: 'crossws-plugin', options })), +})) + +describe('better-ws H3 adapter', () => { + it('converts a better-ws server to an H3 websocket handler', async () => { + const { createServer } = await import('..') + const { toH3Handler } = await import('.') + + const server = createServer() + const handler = toH3Handler(server) + + expect(handler).toMatchObject({ + kind: 'h3-handler', + hooks: { + open: expect.any(Function), + message: expect.any(Function), + }, + }) + }) + + it('creates the CrossWS plugin resolver used by H3 serve', async () => { + const { createH3CrossWsPlugin } = await import('.') + const app = { + fetch: vi.fn(async () => Object.assign(new Response(null), { + crossws: { open: vi.fn() }, + })), + } + + const plugin = createH3CrossWsPlugin(app) + // NOTICE: + // `createH3CrossWsPlugin` deliberately returns the public Srvx plugin + // type, while this mock exposes its resolver for assertion. The cast stays + // inside the test so production code does not depend on mock-only shape. + // Remove this if Srvx/CrossWS exposes an inspectable plugin test helper. + const mockedPlugin = plugin as unknown as { + options: { + resolve: (request: Request) => Promise | undefined> + } + } + const resolved = await mockedPlugin.options.resolve(new Request('http://localhost/ws')) + + expect(plugin).toMatchObject({ kind: 'crossws-plugin' }) + expect(app.fetch).toHaveBeenCalledOnce() + expect(resolved).toEqual({ open: expect.any(Function) }) + }) +}) diff --git a/packages/better-ws/src/server/h3/index.ts b/packages/better-ws/src/server/h3/index.ts new file mode 100644 index 000000000..79f10187f --- /dev/null +++ b/packages/better-ws/src/server/h3/index.ts @@ -0,0 +1,42 @@ +import type { Hooks } from 'crossws' +import type { EventHandler } from 'h3' +import type { ServerPlugin, ServerRequest } from 'srvx' + +import type { WsCrossWsHandlerOptions, WsServer } from '..' + +import { plugin as crossWsPlugin } from 'crossws/server' +import { defineWebSocketHandler } from 'h3' + +import { toCrossWsHooks } from '..' + +export interface H3CrossWsResponse extends Response { + crossws?: Partial +} + +export interface H3CrossWsApp { + fetch: (request: ServerRequest) => Promise +} + +/** + * Converts a better-ws server into an H3 websocket route handler. The H3 + * application still needs the CrossWS plugin installed so upgrade requests can + * resolve the route-scoped hooks produced by this adapter. + */ +export function toH3Handler( + server: WsServer, + options?: WsCrossWsHandlerOptions, +): EventHandler { + return defineWebSocketHandler(toCrossWsHooks(server, options)) +} + +/** + * Creates the CrossWS plugin resolver used by H3 `serve(...)`. + */ +export function createH3CrossWsPlugin(app: H3CrossWsApp): ServerPlugin { + return crossWsPlugin({ + resolve: async (request) => { + const response = await app.fetch(request) + return response.crossws! + }, + }) +} diff --git a/packages/better-ws/src/server/index.ts b/packages/better-ws/src/server/index.ts new file mode 100644 index 000000000..850814854 --- /dev/null +++ b/packages/better-ws/src/server/index.ts @@ -0,0 +1,621 @@ +import type { Message as CrossWsMessage, Peer as CrossWsPeer, Hooks } from 'crossws' + +import type { WsCloseDetails, WsSendResult } from '../shared' +import type { + PeerHealthRecord, + Peer as WsPeer, + PeerAdapter as WsPeerAdapter, + PeerManager as WsPeerManager, + PreviousPeer as WsPreviousPeer, +} from './peers' + +import { createEventWaitFor } from '../shared' +import { createPeers } from './peers' + +export type { + Peer as WsPeer, + PeerAdapter as WsPeerAdapter, + PeerManager as WsPeerManager, + PreviousPeer as WsPreviousPeer, +} from './peers' + +export interface WsServerMessageContext { + /** Server that received the message. */ + server: WsServer + /** Peer that sent the message. */ + peer: WsPeer + /** Incoming caller-owned message. */ + message: TMessage +} + +export interface WsGroup { + /** Sends one message to every peer currently in the group. */ + send: (message: TMessage) => Array +} + +export interface ServerHeartbeatOptions { + /** Reserved app-driven heartbeat transport policy; checkLiveness does not send pings. @default auto */ + mode?: 'auto' | 'native' | 'message' + /** Reserved scheduler hint in milliseconds; callers must still invoke checkLiveness themselves. */ + interval?: number + /** Keepalive timeout and peer liveness fallback in milliseconds. @default 60000 */ + timeout?: number + /** Reserved protocol-neutral heartbeat message; checkLiveness never sends it automatically. */ + message?: TMessage | (() => TMessage) + /** Detects inbound ping control messages that should not enter business handlers. */ + isPing?: (message: TMessage) => boolean + /** Detects inbound pong control messages that should not enter business handlers. */ + isPong?: (message: TMessage) => boolean + /** Legacy response predicate treated as a pong detector when isPong is absent. */ + isResponse?: (message: TMessage) => boolean +} + +export interface WsPeerHealthChange { + /** Peer whose liveness state changed. */ + peer: WsPeer + /** Whether the peer is now considered healthy. */ + healthy: boolean + /** Milliseconds of inbound silence recorded at the time of this change. */ + silentFor: number +} + +export interface ProcedureWaitForOptions { + /** Milliseconds before waitFor rejects. */ + timeout?: number + /** Optional caller-owned abort signal. */ + signal?: AbortSignal +} + +export interface ProcedureContext { + /** Peer that owns this procedure. */ + readonly peer: WsPeer + /** Active peer manager at procedure execution time. */ + readonly peers: WsPeerManager + /** Signal aborted when the procedure finishes or the peer closes. */ + readonly signal: AbortSignal + /** Sends one message to the procedure peer. */ + send: (message: TMessage) => WsSendResult + /** Waits for the next message from the procedure peer matching a predicate. */ + waitFor: ( + predicate: (message: TMessage) => boolean | Promise, + options?: ProcedureWaitForOptions, + ) => Promise +} + +export interface PeerOpenEvent { + /** Server that accepted the peer. */ + readonly server: WsServer + /** Active peer manager after the peer has been accepted. */ + readonly peers: WsPeerManager + /** Accepted peer. */ + readonly peer: WsPeer + /** Snapshot of the same-id peer that was replaced, when present. */ + readonly previous?: WsPreviousPeer + /** Runs a scoped lifecycle procedure for this peer. */ + procedure: (run: (ctx: ProcedureContext) => Promise | T) => Promise +} + +export interface PeerCloseEvent { + /** Server that removed the peer. */ + readonly server: WsServer + /** Active peer manager after the peer has been removed. */ + readonly peers: WsPeerManager + /** Removed peer id. */ + readonly peerId: string + /** Runtime close details when an adapter provides them. */ + readonly details?: WsCloseDetails +} + +export interface PeerControlMessageEvent { + /** Server that received the control message. */ + readonly server: WsServer + /** Active peer manager at control-message dispatch time. */ + readonly peers: WsPeerManager + /** Peer that sent the control message. */ + readonly peer: WsPeer + /** Incoming control message. */ + readonly message: TMessage +} + +export interface WsServerPeerOptions { + /** Milliseconds of inbound silence before a peer is marked unhealthy. */ + unhealthyTimeout?: number + /** Milliseconds of inbound silence before a peer is closed and removed. */ + closeTimeout?: number +} + +export interface WsServerOptions { + /** Peer manager health policy. */ + peers?: WsServerPeerOptions + /** Enables neutral server keepalive signal handling. @default false */ + heartbeat?: false | ServerHeartbeatOptions +} + +export interface WsServer { + /** Active peers managed by the server peer manager. */ + readonly peers: WsPeerManager + /** Accepts a runtime peer adapter into the server registry. */ + accept: (adapter: WsPeerAdapter, options?: { state?: TState | ((previous?: WsPreviousPeer) => TState | undefined) }) => WsPeer + /** Removes one peer from the registry without closing the underlying connection. */ + remove: (peerId: string, details?: WsCloseDetails) => void + /** Registers an incoming message handler. */ + onMessage: (handler: (context: WsServerMessageContext) => void | Promise) => () => void + /** Registers a handler for accepted peers. */ + onPeerOpen: (handler: (event: PeerOpenEvent) => void | Promise) => () => void + /** Registers a handler for removed peers. */ + onPeerClose: (handler: (event: PeerCloseEvent) => void | Promise) => () => void + /** Registers a handler for inbound ping control messages. */ + onPing: (handler: (event: PeerControlMessageEvent) => void | Promise) => () => void + /** Registers a handler for inbound pong control messages. */ + onPong: (handler: (event: PeerControlMessageEvent) => void | Promise) => () => void + /** Registers a handler for server-side peer health transitions. */ + onPeerHealthChange: (handler: (event: WsPeerHealthChange) => void | Promise) => () => void + /** Advances check-based inbound liveness tracking for active peers; it does not send heartbeat messages. */ + checkLiveness: (now?: number) => void + /** Sends one message to all active peers. */ + broadcast: (message: TMessage) => Array + /** Selects a named group for group sends. */ + to: (group: string) => WsGroup + /** Removes all peers and clears runtime state. */ + close: () => void +} + +export interface WsCrossWsHandlerOptions { + /** Reads one caller-owned message from a CrossWS message. @default message.text() */ + readMessage?: (message: CrossWsMessage, peer: CrossWsPeer) => TMessage + /** Resolves the better-ws peer id from a CrossWS peer. @default peer.id */ + peerId?: (peer: CrossWsPeer) => string + /** Creates initial better-ws peer state for a CrossWS peer. */ + state?: (peer: CrossWsPeer) => TState | undefined + /** CrossWS send compression option for messages sent through better-ws peers. */ + compress?: boolean + /** Optional lifecycle hook after a CrossWS peer is accepted. */ + open?: (context: { peer: WsPeer, rawPeer: CrossWsPeer }) => void | Promise + /** Optional lifecycle hook after a CrossWS peer is removed. */ + close?: (context: { peer?: WsPeer, rawPeer: CrossWsPeer, details?: WsCloseDetails }) => void | Promise + /** Optional lifecycle hook for CrossWS errors. */ + error?: (context: { peer?: WsPeer, rawPeer: CrossWsPeer, error: unknown }) => void | Promise +} + +interface CrossWsPeerEntry { + peer: WsPeer + rawPeer: CrossWsPeer +} + +/** + * Creates a runtime-agnostic websocket server peer registry. + * + * Use when: + * - A concrete server adapter needs shared peer management, message handlers, groups, and broadcast semantics + * - Application protocols want to keep full control of message shape and serialization + * + * Expects: + * - Runtime adapters call `accept(...)` on open, `peer.receive(...)` on message, and `peer.close(...)` on close + * + * Returns: + * - A server runtime that tracks peers and dispatches caller-owned messages + */ +export function createServer( + options: WsServerOptions = {}, +): WsServer { + let server: WsServer + const messageHandlers = new Set<(context: WsServerMessageContext) => void | Promise>() + const peerOpenHandlers = new Set<(event: PeerOpenEvent) => void | Promise>() + const peerCloseHandlers = new Set<(event: PeerCloseEvent) => void | Promise>() + const pingHandlers = new Set<(event: PeerControlMessageEvent) => void | Promise>() + const pongHandlers = new Set<(event: PeerControlMessageEvent) => void | Promise>() + const healthChangeHandlers = new Set<(event: WsPeerHealthChange) => void | Promise>() + const procedureControllers = new WeakMap, Set>() + const heartbeat = options.heartbeat === false ? undefined : options.heartbeat + + function createProcedure(peer: WsPeer) { + return async function procedure( + run: (ctx: ProcedureContext) => Promise | T, + ): Promise { + const controller = new AbortController() + + const listeners = new Set<() => void>() + let controllers = procedureControllers.get(peer) + if (!controllers) { + controllers = new Set() + procedureControllers.set(peer, controllers) + } + + controllers.add(controller) + + const ctx: ProcedureContext = { + peer, + peers: server.peers, + signal: controller.signal, + send: message => peer.send(message), + waitFor(predicate, waitOptions = {}) { + const wait = createEventWaitFor, TMessage>({ + match: async ({ peer: fromPeer, message }) => fromPeer === peer && await predicate(message), + select: ({ message }) => message, + timeout: waitOptions.timeout, + signals: [controller.signal, waitOptions.signal], + abortMessage: 'Procedure aborted.', + timeoutMessage: 'Procedure waitFor timed out.', + }) + + const unsubscribe = server.onMessage(wait.emit) + listeners.add(unsubscribe) + + void wait.promise.finally(() => { + unsubscribe() + listeners.delete(unsubscribe) + }).catch(() => {}) + + return wait.promise + }, + } + + try { + return await run(ctx) + } + finally { + controller.abort() + for (const unsubscribe of listeners) { + unsubscribe() + } + + listeners.clear() + + controllers.delete(controller) + if (controllers.size === 0) { + procedureControllers.delete(peer) + } + } + } + } + + function emitPeerOpen(peer: WsPeer, previous?: WsPreviousPeer) { + const event: PeerOpenEvent = { + server, + peers: server.peers, + peer, + previous, + procedure: createProcedure(peer), + } + + for (const handler of peerOpenHandlers) { + void handler(event) + } + } + + function abortPeerProcedures(peer: WsPeer) { + for (const controller of procedureControllers.get(peer) ?? []) { + controller.abort() + } + } + + function emitPeerClose(peerId: string, details?: WsCloseDetails) { + const event: PeerCloseEvent = { + server, + peers: server.peers, + peerId, + details, + } + + for (const handler of peerCloseHandlers) { + void handler(event) + } + } + + function emitControlMessage(peer: WsPeer, message: TMessage, handlers: Set<(event: PeerControlMessageEvent) => void | Promise>) { + const event: PeerControlMessageEvent = { + server, + peers: server.peers, + peer, + message, + } + + for (const handler of handlers) { + void handler(event) + } + } + + function emitHealthChange(peer: WsPeer, health: PeerHealthRecord, silentFor: number) { + const event: WsPeerHealthChange = { + peer, + healthy: health.healthy, + silentFor, + } + + for (const handler of healthChangeHandlers) { + void handler(event) + } + } + + const rawPeers = createPeers({ + onMessage(peer, message) { + const isPing = heartbeat?.isPing?.(message) ?? false + const isPong = heartbeat?.isPong?.(message) ?? heartbeat?.isResponse?.(message) ?? false + + if (isPing || isPong) { + emitControlMessage(peer, message, isPing ? pingHandlers : pongHandlers) + + return + } + + for (const handler of messageHandlers) { + void handler({ server, peer, message }) + } + }, + onSeen(peer, health, wasHealthy) { + if (!wasHealthy) { + emitHealthChange(peer, health, 0) + } + }, + onRemove(peer, details) { + abortPeerProcedures(peer) + emitPeerClose(peer.id, details) + }, + }) + + const peers: WsPeerManager = { + get size() { + return rawPeers.size + }, + get: (peerId) => { + return rawPeers.get(peerId) + }, + has: (peerId) => { + return rawPeers.has(peerId) + }, + list: () => { + return rawPeers.list() + }, + entries: () => { + return rawPeers.entries() + }, + accept(adapter, options) { + const result = rawPeers.accept(adapter, options) + emitPeerOpen(result.peer, result.previous) + return result + }, + remove(peerId, details) { + rawPeers.remove(peerId, details) + }, + close(peerId, code, reason) { + rawPeers.close(peerId, code, reason) + }, + closeAll() { + rawPeers.closeAll() + }, + to: (group) => { + return rawPeers.to(group) + }, + broadcast: (message) => { + return rawPeers.broadcast(message) + }, + markSeen: (peer, now) => { + return rawPeers.markSeen(peer, now) + }, + markUnhealthy: (peer, now) => { + return rawPeers.markUnhealthy(peer, now) + }, + healthOf: (peerId) => { + return rawPeers.healthOf(peerId) + }, + } + + server = { + peers, + accept(adapter, options) { + return peers.accept(adapter, options).peer + }, + remove(peerId, details) { + peers.remove(peerId, details) + }, + onMessage(handler) { + messageHandlers.add(handler) + return () => messageHandlers.delete(handler) + }, + onPeerOpen(handler) { + peerOpenHandlers.add(handler) + return () => peerOpenHandlers.delete(handler) + }, + onPeerClose(handler) { + peerCloseHandlers.add(handler) + return () => peerCloseHandlers.delete(handler) + }, + onPing(handler) { + pingHandlers.add(handler) + return () => pingHandlers.delete(handler) + }, + onPong(handler) { + pongHandlers.add(handler) + return () => pongHandlers.delete(handler) + }, + onPeerHealthChange(handler) { + healthChangeHandlers.add(handler) + return () => healthChangeHandlers.delete(handler) + }, + checkLiveness(now = Date.now()) { + if (!heartbeat && !options.peers) { + return + } + + const unhealthyTimeout = options.peers?.unhealthyTimeout ?? heartbeat?.timeout ?? 60_000 + const closeTimeout = options.peers?.closeTimeout ?? unhealthyTimeout * 2 + + const activePeers = [...rawPeers.entries()] + + for (const [id, peer] of activePeers) { + const health = rawPeers.healthOf(id) + if (!health) { + continue + } + + const silentFor = now - health.lastSeenAt + if (silentFor >= closeTimeout) { + if (health.healthy) { + rawPeers.markUnhealthy(peer, now) + emitHealthChange(peer, { healthy: false, lastSeenAt: health.lastSeenAt, unhealthyAt: now }, silentFor) + } + + peer.close() + continue + } + + if (silentFor >= unhealthyTimeout && health.healthy) { + rawPeers.markUnhealthy(peer, now) + emitHealthChange(peer, { healthy: false, lastSeenAt: health.lastSeenAt, unhealthyAt: now }, silentFor) + } + } + }, + broadcast(message) { + return rawPeers.broadcast(message) + }, + to(group) { + return rawPeers.to(group) + }, + close() { + try { + peers.closeAll() + } + finally { + messageHandlers.clear() + peerOpenHandlers.clear() + peerCloseHandlers.clear() + pingHandlers.clear() + pongHandlers.clear() + healthChangeHandlers.clear() + } + }, + } + + return server +} + +/** + * Creates CrossWS hooks backed by a {@link WsServer} peer registry. + * + * Use when: + * - A CrossWS-compatible runtime should feed raw messages into better-ws server primitives + * - Message shape should stay controlled by the caller through `readMessage` + * + * Expects: + * - CrossWS calls `open`, `message`, and `close` with stable peer ids + * - The default message reader is only used for text protocols + * + * Returns: + * - CrossWS hooks that can be passed to CrossWS adapters or H3 websocket handlers + */ +export function toCrossWsHooks( + server: WsServer, + options: WsCrossWsHandlerOptions = {}, +): Partial { + const peers = new Map>() + const retiredRawPeers = new WeakSet() + const peerId = options.peerId ?? ((peer: CrossWsPeer) => peer.id) + const readMessage = options.readMessage ?? ((message: CrossWsMessage) => message.text() as TMessage) + + server.onPeerOpen(({ peer, previous }) => { + if (!previous) { + return + } + + const existingEntry = peers.get(peer.id) + if (existingEntry && existingEntry.peer !== peer) { + retiredRawPeers.add(existingEntry.rawPeer) + peers.delete(peer.id) + } + }) + + function accept(rawPeer: CrossWsPeer, acceptOptions: { replaceCurrent?: boolean } = {}) { + retiredRawPeers.delete(rawPeer) + const id = peerId(rawPeer) + + const existingEntry = peers.get(id) + const currentServerPeer = server.peers.get(id) + + if (existingEntry && currentServerPeer === existingEntry.peer && !acceptOptions.replaceCurrent) { + return existingEntry.peer + } + + if (existingEntry && currentServerPeer === existingEntry.peer) { + retiredRawPeers.add(existingEntry.rawPeer) + server.remove(id) + peers.delete(id) + } + else if (existingEntry) { + retiredRawPeers.add(existingEntry.rawPeer) + peers.delete(id) + } + + const peer = server.accept({ + id, + send: message => rawPeer.send(message, { compress: options.compress }), + close: (code, reason) => rawPeer.close(code, reason), + }, { + state: options.state?.(rawPeer), + }) + + peers.set(id, { peer, rawPeer }) + + return peer + } + + return { + async open(rawPeer) { + const peer = accept(rawPeer, { replaceCurrent: true }) + await options.open?.({ peer, rawPeer }) + }, + message(rawPeer, message) { + if (retiredRawPeers.has(rawPeer)) { + return + } + + const id = peerId(rawPeer) + const entry = peers.get(id) + if (!entry) { + if (server.peers.has(id)) { + return + } + + accept(rawPeer).receive(readMessage(message, rawPeer)) + return + } + + if (entry.rawPeer !== rawPeer) { + return + } + + if (server.peers.get(entry.peer.id) !== entry.peer) { + peers.delete(id) + if (!server.peers.has(id)) { + accept(rawPeer).receive(readMessage(message, rawPeer)) + } + + return + } + + entry.peer.receive(readMessage(message, rawPeer)) + }, + async close(rawPeer, details) { + if (retiredRawPeers.has(rawPeer)) { + await options.close?.({ peer: undefined, rawPeer, details }) + return + } + + const id = peerId(rawPeer) + const entry = peers.get(id) + + const currentPeer = entry?.rawPeer === rawPeer && server.peers.get(id) === entry.peer ? entry.peer : undefined + if (entry?.rawPeer === rawPeer) { + peers.delete(id) + if (currentPeer) { + server.remove(id, details) + } + } + + await options.close?.({ peer: currentPeer, rawPeer, details }) + }, + async error(rawPeer, error) { + const entry = peers.get(peerId(rawPeer)) + await options.error?.({ peer: entry?.rawPeer === rawPeer ? entry.peer : undefined, rawPeer, error }) + }, + } +} diff --git a/packages/better-ws/src/server/liveness.test.ts b/packages/better-ws/src/server/liveness.test.ts new file mode 100644 index 000000000..ffff90f29 --- /dev/null +++ b/packages/better-ws/src/server/liveness.test.ts @@ -0,0 +1,190 @@ +import { describe, expect, it, vi } from 'vitest' + +import { createServer } from '.' + +describe('better-ws server liveness', () => { + it('marks a peer unhealthy and then removes it after peer health timeouts', () => { + const health: Array<{ peerId: string, healthy: boolean }> = [] + const closed = vi.fn() + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 20, + }, + heartbeat: { + interval: 1, + timeout: 10, + message: () => 'ping', + isResponse: message => message === 'pong', + }, + }) + server.onPeerHealthChange((event) => { + health.push({ peerId: event.peer.id, healthy: event.healthy }) + }) + + const peer = server.peers.accept({ + id: 'peer-1', + send: vi.fn(() => true), + close: closed, + }).peer + const now = Date.now() + + server.checkLiveness(now + 11) + server.checkLiveness(now + 21) + + expect(peer.id).toBe('peer-1') + expect(health).toEqual([{ peerId: 'peer-1', healthy: false }]) + expect(closed).toHaveBeenCalledOnce() + expect(server.peers.has('peer-1')).toBe(false) + }) + + it('marks an unhealthy peer healthy after inbound traffic', () => { + const health: boolean[] = [] + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 30, + }, + heartbeat: { + timeout: 10, + }, + }) + const peer = server.peers.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }).peer + server.onPeerHealthChange((event) => { + health.push(event.healthy) + }) + + server.checkLiveness(Date.now() + 11) + peer.receive('hello') + + expect(health).toEqual([false, true]) + }) + + it('checks peer liveness when peer timeouts are configured and heartbeat transport is disabled', () => { + const closed = vi.fn() + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 20, + }, + heartbeat: false, + }) + server.peers.accept({ + id: 'peer-1', + send: vi.fn(() => true), + close: closed, + }) + + server.checkLiveness(Date.now() + 20) + + expect(closed).toHaveBeenCalledOnce() + expect(server.peers.has('peer-1')).toBe(false) + }) + + it('does not check peer liveness when heartbeat is disabled', () => { + const closed = vi.fn() + const server = createServer() + const peer = server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + close: closed, + }) + + server.checkLiveness(Date.now() + 120_000) + + expect(peer.id).toBe('peer-1') + expect(closed).not.toHaveBeenCalled() + expect(server.peers.has('peer-1')).toBe(true) + }) + + it('removes replaced peer health while keeping group membership on the new peer', () => { + const health: boolean[] = [] + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 30, + }, + heartbeat: { + timeout: 10, + }, + }) + const firstPeer = server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }) + firstPeer.join('room') + server.onPeerHealthChange((event) => { + health.push(event.healthy) + }) + + server.checkLiveness(Date.now() + 11) + const secondPeer = server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + }) + secondPeer.join('room') + secondPeer.receive('hello') + + expect(health).toEqual([false]) + expect(secondPeer.isIn('room')).toBe(true) + expect(server.to('room').send('hello')).toEqual([ + { peerId: 'peer-1', ok: true }, + ]) + }) + + it('emits silent duration when marking a peer unhealthy before close timeout', () => { + const health: Array<{ healthy: boolean, silentFor: number }> = [] + const closed = vi.fn() + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 30, + }, + heartbeat: { + timeout: 10, + }, + }) + server.onPeerHealthChange((event) => { + health.push({ + healthy: event.healthy, + silentFor: event.silentFor, + }) + }) + server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + close: closed, + }) + + server.checkLiveness(Date.now() + 11) + + expect(health).toEqual([{ healthy: false, silentFor: expect.any(Number) }]) + expect(closed).not.toHaveBeenCalled() + expect(server.peers.has('peer-1')).toBe(true) + }) + + it('removes the peer even when the underlying close operation throws', () => { + const server = createServer({ + peers: { + unhealthyTimeout: 10, + closeTimeout: 10, + }, + heartbeat: { + timeout: 10, + }, + }) + server.accept({ + id: 'peer-1', + send: vi.fn(() => true), + close: () => { + throw new Error('raw close failed') + }, + }) + + expect(() => server.checkLiveness(Date.now() + 10)).toThrow('raw close failed') + expect(server.peers.has('peer-1')).toBe(false) + }) +}) diff --git a/packages/better-ws/src/server/peers.ts b/packages/better-ws/src/server/peers.ts new file mode 100644 index 000000000..76fd3721e --- /dev/null +++ b/packages/better-ws/src/server/peers.ts @@ -0,0 +1,341 @@ +import type { WsCloseDetails, WsSendResult } from '../shared' + +import { normalizeSendResult } from '../shared' + +export interface PeerAdapter { + /** Optional caller-owned stable peer id. */ + id?: string + /** Sends one caller-owned message to the underlying connection. */ + send: (message: TMessage) => boolean | number | void + /** Requests closing the underlying connection. */ + close?: (code?: number, reason?: string) => void +} + +export interface PreviousPeer { + /** Stable peer id that was replaced. */ + id: string + /** Caller-owned state snapshot from the replaced peer. */ + state: TState | undefined + /** Groups the replaced peer belonged to. */ + groups: string[] + /** Last inbound activity timestamp known for the replaced peer. */ + lastSeenAt?: number + /** Why the snapshot exists. */ + reason: 'replaced' +} + +export interface Peer { + /** Stable id for this accepted peer. */ + readonly id: string + /** Caller-owned mutable state associated with this peer. */ + state: TState | undefined + /** Sends one message to this peer. */ + send: (message: TMessage) => WsSendResult + /** Feeds one incoming adapter message into server handlers. */ + receive: (message: TMessage) => void + /** Closes the peer and removes it from the server registry. */ + close: (code?: number, reason?: string) => void + /** Adds this peer to a named group. */ + join: (group: string) => void + /** Removes this peer from a named group. */ + leave: (group: string) => void + /** Checks group membership. */ + isIn: (group: string) => boolean +} + +export interface PeerHealthRecord { + /** Whether the peer is currently considered healthy by server liveness policy. */ + healthy: boolean + /** Last inbound activity timestamp known for this peer. */ + lastSeenAt: number + /** Timestamp when the peer was first marked unhealthy. */ + unhealthyAt?: number +} + +type PeerStateFactory = (previous?: PreviousPeer) => TState | undefined + +export interface PeerManager { + readonly size: number + get: (peerId: string) => Peer | undefined + has: (peerId: string) => boolean + list: () => Array> + entries: () => IterableIterator<[string, Peer]> + accept: ( + adapter: PeerAdapter, + options?: { + state?: TState | ((previous?: PreviousPeer) => TState | undefined) + }, + ) => { peer: Peer, previous?: PreviousPeer } + remove: (peerId: string, details?: WsCloseDetails) => void + close: (peerId: string, code?: number, reason?: string) => void + closeAll: () => void + to: (group: string) => { + send: (message: TMessage) => Array + } + broadcast: (message: TMessage) => Array + markSeen: (peer: Peer, now?: number) => void + markUnhealthy: (peer: Peer, now?: number) => void + healthOf: (peerId: string) => Readonly | undefined +} + +let peerIdCounter = 0 + +/** + * Creates the server-owned peer manager that tracks active peers, group + * membership, stale peer handles, and per-peer liveness records. + */ +export function createPeers(input: { + onMessage: (peer: Peer, message: TMessage) => void + onSeen?: (peer: Peer, health: PeerHealthRecord, wasHealthy: boolean) => void + onRemove?: (peer: Peer, details?: WsCloseDetails) => void +}): PeerManager { + const peers = new Map>() + const membershipsByPeer = new Map>() + const peersByGroup = new Map>() + const healthByPeer = new Map() + + function isCurrentPeer(peer: Peer) { + return peers.get(peer.id) === peer + } + + function snapshot(peerId: string): PreviousPeer | undefined { + const previous = peers.get(peerId) + if (!previous) { + return undefined + } + + return { + id: peerId, + state: previous.state, + groups: [...(membershipsByPeer.get(peerId) ?? [])], + lastSeenAt: healthByPeer.get(peerId)?.lastSeenAt, + reason: 'replaced', + } + } + + function removePeerFromGroup(peerId: string, group: string) { + const groupPeers = peersByGroup.get(group) + if (!groupPeers) { + return + } + + groupPeers.delete(peerId) + if (groupPeers.size === 0) { + peersByGroup.delete(group) + } + } + + function remove(peerId: string, details?: WsCloseDetails) { + const peer = peers.get(peerId) + if (!peer) { + return + } + + peers.delete(peerId) + healthByPeer.delete(peerId) + + const memberships = membershipsByPeer.get(peerId) + if (!memberships) { + return + } + + for (const group of memberships) { + removePeerFromGroup(peerId, group) + } + memberships.clear() + membershipsByPeer.delete(peerId) + input.onRemove?.(peer, details) + } + + function removeCurrent(peer: Peer, details?: WsCloseDetails) { + if (!isCurrentPeer(peer)) { + return + } + + remove(peer.id, details) + } + + function markSeen(peer: Peer, now = Date.now()) { + if (!isCurrentPeer(peer)) { + return + } + + const health = healthByPeer.get(peer.id) + if (!health) { + return + } + + const wasHealthy = health.healthy + health.healthy = true + health.lastSeenAt = now + delete health.unhealthyAt + input.onSeen?.(peer, health, wasHealthy) + } + + function markUnhealthy(peer: Peer, now = Date.now()) { + if (!isCurrentPeer(peer)) { + return + } + + const health = healthByPeer.get(peer.id) + if (!health) { + return + } + + health.healthy = false + health.unhealthyAt = now + } + + const manager: PeerManager = { + get size() { + return peers.size + }, + get(peerId) { + return peers.get(peerId) + }, + has(peerId) { + return peers.has(peerId) + }, + list() { + return [...peers.values()] + }, + entries() { + return peers.entries() + }, + accept(adapter, options) { + const id = adapter.id ?? `peer-${++peerIdCounter}` + const previous = snapshot(id) + if (previous) { + remove(id) + } + + const memberships = new Set() + membershipsByPeer.set(id, memberships) + const nextState = options?.state + + const peer: Peer = { + id, + // REVIEW: + // Default previous.state inheritance is convenient for weak-network remote + // plugin reconnects, but it may be wrong for protocols that require forced + // state reinitialization on token rotation, identity switch, or multi-device + // takeover. If those cases appear, replace this default with an explicit + // state(previous, adapter) policy. + state: typeof nextState === 'function' + ? (nextState as PeerStateFactory)(previous) + : nextState ?? previous?.state, + send(message) { + if (!isCurrentPeer(peer)) { + return { ok: false, reason: 'closed' } + } + + return normalizeSendResult(() => adapter.send(message)) + }, + receive(message) { + if (!isCurrentPeer(peer)) { + return + } + + markSeen(peer) + input.onMessage(peer, message) + }, + close(code, reason) { + if (!isCurrentPeer(peer)) { + return + } + + try { + adapter.close?.(code, reason) + } + finally { + removeCurrent(peer, { code, reason }) + } + }, + join(group) { + if (!isCurrentPeer(peer)) { + return + } + + memberships.add(group) + let groupPeers = peersByGroup.get(group) + if (!groupPeers) { + groupPeers = new Set() + peersByGroup.set(group, groupPeers) + } + groupPeers.add(id) + }, + leave(group) { + if (!isCurrentPeer(peer)) { + return + } + + memberships.delete(group) + removePeerFromGroup(id, group) + }, + isIn(group) { + if (!isCurrentPeer(peer)) { + return false + } + + return memberships.has(group) + }, + } + + peers.set(id, peer) + healthByPeer.set(id, { + healthy: true, + lastSeenAt: Date.now(), + }) + + return previous ? { peer, previous } : { peer } + }, + remove, + close(peerId, code, reason) { + peers.get(peerId)?.close(code, reason) + }, + closeAll() { + let firstError: unknown + let nextPeer = peers.values().next().value + while (nextPeer) { + try { + nextPeer.close() + } + catch (error) { + firstError ??= error + } + nextPeer = peers.values().next().value + } + + if (firstError) { + throw firstError + } + }, + to(group) { + return { + send(message) { + return [...(peersByGroup.get(group) ?? new Set())] + .map(peerId => peers.get(peerId)) + .filter(peer => typeof peer !== 'undefined') + .map(peer => ({ + peerId: peer.id, + ...peer.send(message), + })) + }, + } + }, + broadcast(message) { + return [...peers.values()].map(peer => ({ + peerId: peer.id, + ...peer.send(message), + })) + }, + markSeen, + markUnhealthy, + healthOf(peerId) { + const health = healthByPeer.get(peerId) + return health ? { ...health } : undefined + }, + } + + return manager +} diff --git a/packages/better-ws/src/shared/index.ts b/packages/better-ws/src/shared/index.ts new file mode 100644 index 000000000..a9a17088b --- /dev/null +++ b/packages/better-ws/src/shared/index.ts @@ -0,0 +1,80 @@ +/** + * Client connection lifecycle state. + * + * `preparing` means the transport is open while caller-owned bootstrap work can + * still run. `ready` means normal application sends are allowed. `failed` means + * the client reached a terminal failure instead of a clean close. + */ +export type WsState + = | 'idle' + | 'connecting' + | 'open' + | 'preparing' + | 'ready' + | 'reconnecting' + | 'closing' + | 'closed' + | 'failed' + +/** + * Result returned by best-effort send operations. + * + * Use when: + * - Callers need a stable result instead of runtime-specific WebSocket return values + * - Adapters may report backpressure, closed sockets, or thrown send errors differently + * + * Expects: + * - `ok: false` means the message was not accepted by the local runtime + * + * Returns: + * - A transport-neutral send result for client, peer, and broadcast calls + */ +export interface WsSendResult { + /** Whether the local runtime accepted the outgoing message. */ + ok: boolean + /** Optional stable reason for rejected local sends. */ + reason?: 'closed' | 'backpressure' | 'error' + /** Original error when the adapter threw during send. */ + error?: unknown +} + +/** + * Describes a websocket close notification without exposing a concrete runtime type. + */ +export interface WsCloseDetails { + /** WebSocket close code when one is available. */ + code?: number + /** WebSocket close reason when one is available. */ + reason?: string + /** Whether the runtime considered the close clean. */ + wasClean?: boolean +} + +/** + * Converts runtime-specific send return values into a stable public result. + * + * Use when: + * - Client and server adapters expose different send result semantics + * - Public APIs need to avoid leaking concrete runtime return values + * + * Expects: + * - `false` from an adapter means local backpressure or rejected send + * + * Returns: + * - A normalized send result. + */ +export function normalizeSendResult(run: () => boolean | number | void): WsSendResult { + try { + const result = run() + if (result === false) { + return { ok: false, reason: 'backpressure' } + } + + return { ok: true } + } + catch (error) { + return { ok: false, reason: 'error', error } + } +} + +export * from './utils' diff --git a/packages/better-ws/src/shared/utils/event-wait-for.test.ts b/packages/better-ws/src/shared/utils/event-wait-for.test.ts new file mode 100644 index 000000000..c14b0f253 --- /dev/null +++ b/packages/better-ws/src/shared/utils/event-wait-for.test.ts @@ -0,0 +1,52 @@ +import { describe, expect, it, vi } from 'vitest' + +import { createEventWaitFor } from './event-wait-for' + +describe('createEventWaitFor', () => { + it('resolves with the selected value when a future event matches', async () => { + const wait = createEventWaitFor<{ type: string, value: number }, number>({ + match: event => event.type === 'ready', + select: event => event.value, + }) + + wait.emit({ type: 'ignore', value: 1 }) + wait.emit({ type: 'ready', value: 2 }) + + await expect(wait.promise).resolves.toBe(2) + }) + + it('rejects when the timeout expires before a matching event arrives', async () => { + vi.useFakeTimers() + try { + const wait = createEventWaitFor({ + match: message => message === 'ready', + timeout: 100, + timeoutMessage: 'Timed out waiting for ready message.', + }) + + wait.emit('ignore') + vi.advanceTimersByTime(100) + + await expect(wait.promise).rejects.toThrow('Timed out waiting for ready message.') + } + finally { + vi.useRealTimers() + } + }) + + it('rejects when an abort signal fires and ignores later events', async () => { + const controller = new AbortController() + const selected = vi.fn((message: string) => message) + const wait = createEventWaitFor({ + select: selected, + signals: [controller.signal], + abortMessage: 'Wait aborted.', + }) + + controller.abort() + wait.emit('late') + + await expect(wait.promise).rejects.toThrow('Wait aborted.') + expect(selected).not.toHaveBeenCalled() + }) +}) diff --git a/packages/better-ws/src/shared/utils/event-wait-for.ts b/packages/better-ws/src/shared/utils/event-wait-for.ts new file mode 100644 index 000000000..8d4c63766 --- /dev/null +++ b/packages/better-ws/src/shared/utils/event-wait-for.ts @@ -0,0 +1,185 @@ +import { createContext, defineEventa } from '@moeru/eventa' + +/** + * Options for a one-shot event wait controller. + * + * @param TEvent - Event value fed to the wait controller. + * @param TResult - Value resolved from the matching event. + */ +export interface EventWaitForOptions { + /** Predicate that decides whether an event should resolve the wait. @default () => true */ + match?: (event: TEvent) => boolean | Promise + /** Projects the matched event into the resolved value. @default event => event */ + select?: (event: TEvent) => TResult + /** Milliseconds before the wait rejects. When omitted, no timeout is scheduled. */ + timeout?: number + /** Abort signals that reject the wait when any signal aborts. */ + signals?: Array + /** Runtime guard checked before handling and resolving events. @default () => true */ + isActive?: () => boolean + /** Error message used when the wait aborts or becomes inactive. @default 'Wait aborted.' */ + abortMessage?: string + /** Error message used when timeout expires. @default 'Timed out waiting for event.' */ + timeoutMessage?: string +} + +/** + * One-shot wait controller for callback-driven event sources. + * + * @param TEvent - Event value fed to the wait controller. + * @param TResult - Value resolved from the matching event. + */ +export interface EventWaitFor { + /** Promise resolved by the first matching event, or rejected by abort/timeout/predicate errors. */ + promise: Promise + /** Emits one future event into the wait controller. */ + emit: (event: TEvent) => void + /** Rejects the wait and releases timers/listeners when the caller abandons it. */ + dispose: (reason?: unknown) => void +} + +function createWaitError(message: string, reason?: unknown) { + if (reason instanceof Error) { + return reason + } + + return new Error(reason ? `${message}: ${String(reason)}` : message) +} + +/** + * Creates a promise plus emit pair for waiting on callback-driven events. + * + * Use when: + * - An event source pushes values through callbacks, sets, or Eventa handlers + * - The caller needs timeout, abort, predicate, and cleanup semantics around one future event + * + * Expects: + * - Callers register {@link EventWaitFor.emit} with their event source and unregister it when the promise settles + * - Matching starts only for future events emitted through the returned controller + * + * Returns: + * - A one-shot wait controller that settles once and ignores later events + */ +export function createEventWaitFor( + options: EventWaitForOptions = {}, +): EventWaitFor { + const { + match = () => true, + select = (event: TEvent) => event as unknown as TResult, + abortMessage = 'Wait aborted.', + timeoutMessage = 'Timed out waiting for event.', + signals = [], + } = options + + const normalizedSignals = signals.filter((signal): signal is AbortSignal => signal != null) + + const events = createContext() + const waitEvent = defineEventa('better-ws:wait-for-event') + + let settled = false + + let timeoutHandle: ReturnType | undefined + let unsubscribe = () => {} + + let resolvePromise!: (value: TResult) => void + let rejectPromise!: (error: unknown) => void + const promise = new Promise((resolve, reject) => { + resolvePromise = resolve + rejectPromise = reject + }) + + const cleanup = () => { + if (timeoutHandle) { + clearTimeout(timeoutHandle) + timeoutHandle = undefined + } + + for (const signal of normalizedSignals) { + signal.removeEventListener('abort', abort) + } + + unsubscribe() + unsubscribe = () => {} + } + + const settle = (run: () => void) => { + if (settled) { + return + } + + settled = true + cleanup() + run() + } + + function rejectWith(message: string, reason?: unknown) { + settle(() => rejectPromise(createWaitError(message, reason))) + } + + function abort() { + rejectWith(abortMessage) + } + + function assertActive() { + if (options.isActive?.() === false) { + rejectWith(abortMessage) + return false + } + + return true + } + + unsubscribe = events.on(waitEvent, ({ body }) => { + if (settled || !assertActive()) { + return + } + + // NOTICE: + // Eventa payload wrappers expose an optional body field, but better-ws waits + // allow undefined as a valid caller-owned event value. Use the property value + // directly instead of treating missing truthiness as absent payload. + const event = body as TEvent + + try { + const matched = match(event) + if (typeof matched === 'boolean') { + if (matched && assertActive()) { + settle(() => resolvePromise(select(event))) + } + return + } + + void matched + .then((asyncMatched) => { + if (asyncMatched && assertActive()) { + settle(() => resolvePromise(select(event))) + } + }) + .catch(error => rejectWith(abortMessage, error)) + } + catch (error) { + rejectWith(abortMessage, error) + } + }) + + if (!assertActive() || normalizedSignals.some(signal => signal.aborted)) { + rejectWith(abortMessage) + } + else { + for (const signal of normalizedSignals) { + signal.addEventListener('abort', abort, { once: true }) + } + + if (options.timeout !== undefined) { + timeoutHandle = setTimeout(() => { + rejectWith(timeoutMessage) + }, options.timeout) + } + } + + return { + promise, + emit: event => events.emit(waitEvent, event), + dispose: reason => rejectWith(abortMessage, reason), + } +} diff --git a/packages/better-ws/src/shared/utils/index.ts b/packages/better-ws/src/shared/utils/index.ts new file mode 100644 index 000000000..cbe4dbefa --- /dev/null +++ b/packages/better-ws/src/shared/utils/index.ts @@ -0,0 +1 @@ +export * from './event-wait-for' diff --git a/packages/better-ws/tsconfig.json b/packages/better-ws/tsconfig.json new file mode 100644 index 000000000..c58de42e2 --- /dev/null +++ b/packages/better-ws/tsconfig.json @@ -0,0 +1,22 @@ +{ + "compilerOptions": { + "target": "ESNext", + "lib": [ + "ESNext", + "DOM" + ], + "module": "ESNext", + "moduleResolution": "bundler", + "types": [ + "node" + ], + "esModuleInterop": true, + "forceConsistentCasingInFileNames": true, + "isolatedModules": true, + "verbatimModuleSyntax": true, + "skipLibCheck": true + }, + "include": [ + "src/**/*.ts" + ] +} diff --git a/packages/better-ws/tsdown.config.ts b/packages/better-ws/tsdown.config.ts new file mode 100644 index 000000000..e3839146f --- /dev/null +++ b/packages/better-ws/tsdown.config.ts @@ -0,0 +1,11 @@ +import { defineConfig } from 'tsdown' + +export default defineConfig({ + entry: { + 'index': 'src/index.ts', + 'client/crossws': 'src/client/crossws/index.ts', + 'server': 'src/server/index.ts', + 'server/h3': 'src/server/h3/index.ts', + }, + dts: true, +}) diff --git a/packages/better-ws/vitest.config.ts b/packages/better-ws/vitest.config.ts new file mode 100644 index 000000000..647f39364 --- /dev/null +++ b/packages/better-ws/vitest.config.ts @@ -0,0 +1,7 @@ +import { defineConfig } from 'vitest/config' + +export default defineConfig({ + test: { + include: ['src/**/*.test.ts'], + }, +}) diff --git a/packages/server-runtime/package.json b/packages/server-runtime/package.json index 0974965ae..3f6ae210e 100644 --- a/packages/server-runtime/package.json +++ b/packages/server-runtime/package.json @@ -41,11 +41,13 @@ "dependencies": { "@guiiai/logg": "catalog:", "@moeru/std": "catalog:", + "@proj-airi/better-ws": "workspace:^", "@proj-airi/server-shared": "workspace:^", "crossws": "catalog:", "h3": "catalog:", "nanoid": "catalog:", "srvx": "catalog:", - "superjson": "catalog:" + "superjson": "catalog:", + "valibot": "catalog:" } } diff --git a/packages/server-runtime/src/index.test.ts b/packages/server-runtime/src/index.test.ts index 8c91b220b..d2945a461 100644 --- a/packages/server-runtime/src/index.test.ts +++ b/packages/server-runtime/src/index.test.ts @@ -1,8 +1,12 @@ import type { WebSocketBaseEvent, WebSocketEvents } from '@proj-airi/server-shared/types' +import type { ConsumerStickyAssignment } from './server-ws/airi/consumers' + import { describe, expect, it } from 'vitest' -import { heartbeatFrameFrom, resolveEventDelivery, selectConsumerPeerId } from './index' +import { heartbeatFrameFrom } from './server-ws/airi/codec' +import { selectConsumerPeerId } from './server-ws/airi/consumers' +import { resolveEventDelivery } from './server-ws/airi/routing' function createInputTextEvent( overrides: Partial> = {}, @@ -133,7 +137,7 @@ describe('selectConsumerPeerId', () => { }) it('keeps sticky delivery on the same consumer when available', () => { - const stickyAssignments = new Map() + const stickyAssignments = new Map() const firstSelectedPeerId = selectConsumerPeerId({ eventType: 'input:text', diff --git a/packages/server-runtime/src/index.ts b/packages/server-runtime/src/index.ts index 398a5dc0c..dd17504a4 100644 --- a/packages/server-runtime/src/index.ts +++ b/packages/server-runtime/src/index.ts @@ -1,3 +1,5 @@ +import type { WsCloseDetails } from '@proj-airi/better-ws' +import type { WsPeer } from '@proj-airi/better-ws/server' import type { DeliveryConfig, ExtensionIdentity, @@ -6,12 +8,12 @@ import type { WebSocketBaseEvent, WebSocketEvent, } from '@proj-airi/server-shared/types' +import type { Message as CrossWsMessage, Peer as CrossWsPeer } from 'crossws' import type { RouteMiddleware, RoutingPolicy, } from './middlewares' -import type { ServerWsConsumerSelectionCandidate, ServerWsStickyAssignment } from './server-ws/core' import type { AuthenticatedPeer, Peer, RegisteredExtensionModule } from './types' import { Buffer } from 'node:buffer' @@ -19,6 +21,8 @@ import { timingSafeEqual } from 'node:crypto' import { availableLogLevelStrings, Format, LogLevelString, logLevelStringToLogLevelMap, useLogg } from '@guiiai/logg' import { errorMessageFrom } from '@moeru/std' +import { createServer as createWsServer } from '@proj-airi/better-ws/server' +import { toH3Handler } from '@proj-airi/better-ws/server/h3' import { createInvalidJsonServerErrorMessage, ServerErrorMessages, @@ -27,7 +31,7 @@ import { MessageHeartbeat, MessageHeartbeatKind, } from '@proj-airi/server-shared/types' -import { defineWebSocketHandler, H3 } from 'h3' +import { H3 } from 'h3' import { nanoid } from 'nanoid' import { optionOrEnv } from './config' @@ -38,114 +42,53 @@ import { matchesDestinations, } from './middlewares' import { - createEventMetadata, - createGateway, - createResponses, - forEachEventMiddlewares, heartbeatFrameFrom, - isAiriWebSocketEventFormatError, + isInvalidEventError, parseEvent, - resolveEventDelivery, stringifyEvent, -} from './server-ws/airi' +} from './server-ws/airi/codec' import { createConsumerOrchestrator, - createServerWsPeerStore, isConsumerDeliveryMode, normalizeConsumerMode, normalizeConsumerPriority, - resolveServerWsHealthCheckIntervalMs, - selectConsumerPeerId as selectServerWsConsumerPeerId, +} from './server-ws/airi/consumers' +import { + resolveHealthCheckIntervalMs, serverWsDefaultHeartbeatTtlMs, - serverWsHealthCheckMissesDead, - serverWsHealthCheckMissesUnhealthy, -} from './server-ws/core' - -export { - heartbeatFrameFrom, +} from './server-ws/airi/liveness' +import { + createEventMetadata, + createResponses, +} from './server-ws/airi/responses' +import { + forEachEventMiddlewares, resolveEventDelivery, +} from './server-ws/airi/routing' + +interface AiriWsMessage { + text: () => string } -/** - * Candidate peer metadata used for consumer selection. - */ -export type ConsumerSelectionCandidate = ServerWsConsumerSelectionCandidate - -function normalizeRootConsumerGroup(mode: DeliveryConfig['mode'], group?: string) { - if (mode === 'consumer') { - return 'default' - } - - return group || 'default' +interface AiriWsPeerState { + rawPeer: CrossWsPeer } -/** - * Selects a concrete consumer peer for consumer-style delivery modes. - * - * Use when: - * - Existing server-runtime callers need the package-root consumer selector - * - Sticky and round-robin state should remain stored in the original root API shape - * - * Expects: - * - Candidates already describe authenticated and health state - * - * Returns: - * - The selected peer id, or `undefined` when no eligible consumer is available - */ -export function selectConsumerPeerId(options: { - eventType: string - fromPeerId: string - delivery?: DeliveryConfig - candidates: ConsumerSelectionCandidate[] - roundRobinCursor?: Map - stickyAssignments?: Map -}) { - if (!options.delivery || !isConsumerDeliveryMode(options.delivery.mode)) { - return selectServerWsConsumerPeerId({ - eventType: options.eventType, - fromPeerId: options.fromPeerId, - delivery: options.delivery, - candidates: options.candidates, - roundRobinCursor: options.roundRobinCursor, - }) +function airiPeerFromRaw(rawPeer: CrossWsPeer): Peer { + // CrossWS peers expose the connection fields AIRI historically used directly + // (`id`, `send`, `close`, `remoteAddress`, and `request`). Keep the cast in + // this adapter boundary so protocol code below still depends on the AIRI peer + // contract instead of the concrete transport type. + return rawPeer as Peer +} + +function rawPeerFrom(wsPeer: WsPeer): Peer | undefined { + const rawPeer = wsPeer.state?.rawPeer + if (!rawPeer) { + return undefined } - const normalizedGroup = normalizeRootConsumerGroup(options.delivery.mode, options.delivery.group) - const legacyRegistryKey = `${options.eventType}::${normalizedGroup}` - const coreRegistryKey = JSON.stringify([options.eventType, normalizedGroup]) - const roundRobinCursor = options.roundRobinCursor - ? new Map([[coreRegistryKey, options.roundRobinCursor.get(legacyRegistryKey) ?? 0]]) - : undefined - - const stickyAssignments = new Map() - if (options.delivery.selection === 'sticky' && options.delivery.stickyKey && options.stickyAssignments) { - const legacyStickyKey = `${legacyRegistryKey}::${options.delivery.stickyKey}` - const stickyPeerId = options.stickyAssignments.get(legacyStickyKey) - if (stickyPeerId) { - stickyAssignments.set(JSON.stringify([options.eventType, normalizedGroup, options.delivery.stickyKey]), { - event: options.eventType, - group: normalizedGroup, - peerId: stickyPeerId, - }) - } - } - - const selectedPeerId = selectServerWsConsumerPeerId({ - ...options, - roundRobinCursor, - stickyAssignments, - }) - - const nextCursor = roundRobinCursor?.get(coreRegistryKey) - if (typeof nextCursor === 'number') { - options.roundRobinCursor?.set(legacyRegistryKey, nextCursor) - } - - if (options.delivery.selection === 'sticky' && options.delivery.stickyKey && selectedPeerId) { - options.stickyAssignments?.set(`${legacyRegistryKey}::${options.delivery.stickyKey}`, selectedPeerId) - } - - return selectedPeerId + return airiPeerFromRaw(rawPeer) } /** @@ -284,8 +227,10 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => }) // === Registries & Orchestrators === - const peerStore = createServerWsPeerStore() - const peers = peerStore.peers + // TODO: Move protocol-neutral peer registry, consumer selection, and heartbeat + // primitives into `@proj-airi/better-ws/server` so server-runtime only owns + // AIRI authentication, registry sync, route policy, and extension events. + const peers = new Map() const peersByModule = new Map>() const consumers = createConsumerOrchestrator() const heartbeatTtlMs = options?.heartbeat?.readTimeout ?? serverWsDefaultHeartbeatTtlMs @@ -296,12 +241,12 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => ...(options?.routing?.middleware ?? []), ] - const healthCheckIntervalMs = resolveServerWsHealthCheckIntervalMs(heartbeatTtlMs) + const healthCheckIntervalMs = resolveHealthCheckIntervalMs(heartbeatTtlMs) let disposed = false // === Health Check & Peer Liveness === function broadcastPeerHealthy(peerInfo: AuthenticatedPeer, parentId?: string) { - if (!peerInfo.name || !peerInfo.identity) { + if (!peerInfo.authenticated || !peerInfo.name || !peerInfo.identity) { return } @@ -312,6 +257,24 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => }) } + function broadcastPeerUnhealthy(peerInfo: AuthenticatedPeer, reason: string) { + if (peerInfo.name && peerInfo.identity) { + broadcastToAuthenticated({ + type: 'registry:modules:health:unhealthy', + data: { name: peerInfo.name, index: peerInfo.index, identity: peerInfo.identity, reason }, + metadata: createEventMetadata(instanceId), + }) + } + + for (const module of peerInfo.extensionModules?.values() ?? []) { + broadcastToAuthenticated({ + type: 'registry:modules:health:unhealthy', + data: { name: module.name, identity: module.identity, reason }, + metadata: createEventMetadata(instanceId), + }) + } + } + function markPeerAlive(peerInfo: AuthenticatedPeer, options?: { parentId?: string, logMessage?: string }) { peerInfo.lastHeartbeatAt = Date.now() peerInfo.missedHeartbeats = 0 @@ -333,61 +296,6 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => consumers.clear() } - const healthCheckInterval = setInterval(() => { - const now = Date.now() - for (const [id, peerInfo] of peers.entries()) { - if (!peerInfo.lastHeartbeatAt) { - continue - } - - const elapsed = now - peerInfo.lastHeartbeatAt - if (elapsed > healthCheckIntervalMs) { - peerInfo.missedHeartbeats = (peerInfo.missedHeartbeats ?? 0) + 1 - } - else { - peerInfo.missedHeartbeats = 0 - } - - if (peerInfo.missedHeartbeats >= serverWsHealthCheckMissesDead) { - // 10 consecutive misses — completely dead, drop the peer - logger.withFields({ peer: id, peerName: peerInfo.name, missedHeartbeats: peerInfo.missedHeartbeats }).debug('heartbeat expired after max misses, dropping peer') - try { - peerInfo.peer.close?.() - } - catch (error) { - logger.withFields({ peer: id, peerName: peerInfo.name }).withError(error as Error).debug('failed to close expired peer') - } - - peers.delete(id) - unregisterModulePeer(peerInfo, 'heartbeat expired') - } - else if (peerInfo.missedHeartbeats >= serverWsHealthCheckMissesUnhealthy && peerInfo.healthy !== false) { - // 5 consecutive misses — mark unhealthy - peerInfo.healthy = false - logger.withFields({ peer: id, peerName: peerInfo.name, missedHeartbeats: peerInfo.missedHeartbeats }).debug('heartbeat late, marking unhealthy') - - if (peerInfo.name && peerInfo.identity) { - broadcastToAuthenticated({ - type: 'registry:modules:health:unhealthy', - data: { name: peerInfo.name, index: peerInfo.index, identity: peerInfo.identity, reason: 'heartbeat late' }, - metadata: createEventMetadata(instanceId), - }) - } - - for (const module of peerInfo.extensionModules?.values() ?? []) { - broadcastToAuthenticated({ - type: 'registry:modules:health:unhealthy', - data: { name: module.name, identity: module.identity, reason: 'heartbeat late' }, - metadata: createEventMetadata(instanceId), - }) - } - } - } - }, healthCheckIntervalMs) - if (typeof healthCheckInterval === 'object') { - healthCheckInterval.unref?.() - } - // === Module Registry & Consumer Management === function registerExtensionModulePeer(p: AuthenticatedPeer, module: RegisteredExtensionModule) { p.extensionModules ??= new Map() @@ -609,501 +517,600 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => } } - // === WebSocket Gateway Handler === - // Handles peer lifecycle: open, message, error, close - const websocketGateway = createGateway({ - handler: { - open: (peer) => { - if (authToken) { - peers.set(peer.id, { peer, authenticated: false, name: '', lastHeartbeatAt: Date.now() }) - } - else { - send(peer, RESPONSES.authenticated()) - peers.set(peer.id, { peer, authenticated: true, name: '', lastHeartbeatAt: Date.now() }) - sendRegistrySync(peer) + // === WebSocket Server Handlers === + // Handles AIRI peer lifecycle: open, message, error, close. + const wsServer = createWsServer({ + peers: { + unhealthyTimeout: heartbeatTtlMs, + closeTimeout: heartbeatTtlMs * 2, + }, + heartbeat: { + interval: healthCheckIntervalMs, + timeout: heartbeatTtlMs, + }, + }) + wsServer.onPeerOpen(({ peer: wsPeer }) => { + const peer = rawPeerFrom(wsPeer) + if (!peer) + return + + if (authToken) { + peers.set(peer.id, { peer, authenticated: false, name: '', lastHeartbeatAt: Date.now() }) + } + else { + send(peer, RESPONSES.authenticated()) + peers.set(peer.id, { peer, authenticated: true, name: '', lastHeartbeatAt: Date.now() }) + sendRegistrySync(peer) + } + + logger.withFields({ peer: peer.id, activePeers: peers.size }).log('connected') + }) + + wsServer.onMessage(({ peer: wsPeer, message }) => { + const peer = rawPeerFrom(wsPeer) + if (!peer) + return + + const authenticatedPeer = peers.get(peer.id) + let event: WebSocketEvent + + try { + const text = message.text() + const controlFrame = heartbeatFrameFrom(text) + + // Some websocket runtimes surface control frames as plain text messages instead of + // exposing them through dedicated ping/pong hooks. Treat those payloads as transport + // liveness only so they do not leak into the application event protocol. + if (controlFrame) { + if (authenticatedPeer) { + markPeerAlive(authenticatedPeer, { logMessage: 'ping/pong recovered, marking healthy' }) } - logger.withFields({ peer: peer.id, activePeers: peers.size }).log('connected') - }, - message: (peer, message) => { - const authenticatedPeer = peers.get(peer.id) - let event: WebSocketEvent + return + } - try { - const text = message.text() - const controlFrame = heartbeatFrameFrom(text) + event = parseEvent(text) + } + catch (err) { + if (isInvalidEventError(err)) { + send(peer, RESPONSES.error(ServerErrorMessages.invalidEventFormat)) + return + } - // Some websocket runtimes surface control frames as plain text messages instead of - // exposing them through dedicated ping/pong hooks. Treat those payloads as transport - // liveness only so they do not leak into the application event protocol. - if (controlFrame) { - if (authenticatedPeer) { - markPeerAlive(authenticatedPeer, { logMessage: 'ping/pong recovered, marking healthy' }) - } + const errorMessage = errorMessageFrom(err) ?? 'Unknown JSON parsing error' + send(peer, RESPONSES.error(createInvalidJsonServerErrorMessage(errorMessage))) - return - } + return + } - event = parseEvent(text) + logger.withFields({ + peer: peer.id, + peerAuthenticated: authenticatedPeer?.authenticated, + peerModule: authenticatedPeer?.name, + peerModuleIndex: authenticatedPeer?.index, + }).debug('received event') + + if (authenticatedPeer) { + markPeerAlive(authenticatedPeer, { parentId: event.metadata?.event.id }) + + if (authenticatedPeer.authenticated && isExtensionModuleIdentity(event.metadata?.source)) { + authenticatedPeer.identity = event.metadata.source + } + } + + switch (event.type) { + case 'transport:connection:heartbeat': { + const p = peers.get(peer.id) + if (p) { + markPeerAlive(p, { + parentId: event.metadata?.event.id, + logMessage: 'heartbeat recovered, marking healthy', + }) + + // recover from unhealthy → healthy } - catch (err) { - if (isAiriWebSocketEventFormatError(err)) { - send(peer, RESPONSES.error(ServerErrorMessages.invalidEventFormat)) - return - } - const errorMessage = errorMessageFrom(err) ?? 'Unknown JSON parsing error' - send(peer, RESPONSES.error(createInvalidJsonServerErrorMessage(errorMessage))) + if (event.data.kind === MessageHeartbeatKind.Ping) { + send(peer, RESPONSES.heartbeat(MessageHeartbeatKind.Pong, heartbeatMessage, event.metadata?.event.id)) + } + + return + } + + case 'module:authenticate': { + const clientToken = typeof event.data.token === 'string' ? event.data.token : '' + if (authToken && !timingSafeCompare(clientToken, authToken)) { + logger.withFields({ peer: peer.id, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).log('authentication failed') + send(peer, RESPONSES.error(ServerErrorMessages.invalidToken, event.metadata?.event.id)) return } - logger.withFields({ - peer: peer.id, - peerAuthenticated: authenticatedPeer?.authenticated, - peerModule: authenticatedPeer?.name, - peerModuleIndex: authenticatedPeer?.index, - }).debug('received event') - - if (authenticatedPeer) { - markPeerAlive(authenticatedPeer, { parentId: event.metadata?.event.id }) - - if (authenticatedPeer.authenticated && isExtensionModuleIdentity(event.metadata?.source)) { - authenticatedPeer.identity = event.metadata.source - } + send(peer, RESPONSES.authenticated(event.metadata?.event.id)) + const p = peers.get(peer.id) + if (p) { + p.authenticated = true } - switch (event.type) { - case 'transport:connection:heartbeat': { - const p = peers.get(peer.id) - if (p) { - markPeerAlive(p, { - parentId: event.metadata?.event.id, - logMessage: 'heartbeat recovered, marking healthy', - }) + sendRegistrySync(peer, event.metadata?.event.id) - // recover from unhealthy → healthy - } + return + } - if (event.data.kind === MessageHeartbeatKind.Ping) { - send(peer, RESPONSES.heartbeat(MessageHeartbeatKind.Pong, heartbeatMessage, event.metadata?.event.id)) - } + case 'peer:authenticate': { + const clientToken = typeof event.data.token === 'string' ? event.data.token : '' + if (authToken && !timingSafeCompare(clientToken, authToken)) { + logger.withFields({ peer: peer.id, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).log('peer authentication failed') + send(peer, RESPONSES.error(ServerErrorMessages.invalidToken, event.metadata?.event.id)) - return - } + return + } - case 'module:authenticate': { - const clientToken = typeof event.data.token === 'string' ? event.data.token : '' - if (authToken && !timingSafeCompare(clientToken, authToken)) { - logger.withFields({ peer: peer.id, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).log('authentication failed') - send(peer, RESPONSES.error(ServerErrorMessages.invalidToken, event.metadata?.event.id)) + const authenticatedPeerId = event.data.peerId ?? peer.id + send(peer, RESPONSES.peerAuthenticated(authenticatedPeerId, event.metadata?.event.id)) + const p = peers.get(peer.id) + if (p) { + p.authenticated = true + p.peerIds ??= new Set() + p.peerIds.add(peer.id) + p.peerIds.add(authenticatedPeerId) + } - return - } + sendRegistrySync(peer, event.metadata?.event.id) - send(peer, RESPONSES.authenticated(event.metadata?.event.id)) - const p = peers.get(peer.id) - if (p) { - p.authenticated = true - } + return + } - sendRegistrySync(peer, event.metadata?.event.id) + case 'extension:authenticate': { + const clientToken = typeof event.data.token === 'string' ? event.data.token : '' + if (authToken && !timingSafeCompare(clientToken, authToken)) { + logger.withFields({ peer: peer.id, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).log('extension authentication failed') + send(peer, RESPONSES.error(ServerErrorMessages.invalidToken, event.metadata?.event.id)) - return - } + return + } - case 'peer:authenticate': { - const clientToken = typeof event.data.token === 'string' ? event.data.token : '' - if (authToken && !timingSafeCompare(clientToken, authToken)) { - logger.withFields({ peer: peer.id, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).log('peer authentication failed') - send(peer, RESPONSES.error(ServerErrorMessages.invalidToken, event.metadata?.event.id)) + const p = peers.get(peer.id) + if (p) { + p.authenticated = true + p.extensionIdentity = event.data.identity + } - return - } + send(peer, RESPONSES.extensionAuthenticated(event.data.identity, event.metadata?.event.id)) + sendRegistrySync(peer, event.metadata?.event.id) - const authenticatedPeerId = event.data.peerId ?? peer.id - send(peer, RESPONSES.peerAuthenticated(authenticatedPeerId, event.metadata?.event.id)) - const p = peers.get(peer.id) - if (p) { - p.authenticated = true - p.peerIds ??= new Set() - p.peerIds.add(peer.id) - p.peerIds.add(authenticatedPeerId) - } + return + } - sendRegistrySync(peer, event.metadata?.event.id) + case 'extension:announce': { + const p = peers.get(peer.id) + if (!p) { + return + } - return - } + if (authToken && !p.authenticated) { + send(peer, RESPONSES.error(ServerErrorMessages.mustAuthenticateBeforeAnnouncing)) - case 'extension:authenticate': { - const clientToken = typeof event.data.token === 'string' ? event.data.token : '' - if (authToken && !timingSafeCompare(clientToken, authToken)) { - logger.withFields({ peer: peer.id, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).log('extension authentication failed') - send(peer, RESPONSES.error(ServerErrorMessages.invalidToken, event.metadata?.event.id)) + return + } - return - } + if (!isExtensionIdentity(event.data.identity)) { + send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceIdentityInvalid)) - const p = peers.get(peer.id) - if (p) { - p.authenticated = true - p.extensionIdentity = event.data.identity - } + return + } - send(peer, RESPONSES.extensionAuthenticated(event.data.identity, event.metadata?.event.id)) - sendRegistrySync(peer, event.metadata?.event.id) + p.extensionIdentity = event.data.identity - return - } + send(peer, { + type: 'extension:announced', + data: event.data, + metadata: createEventMetadata(instanceId, event.metadata?.event.id), + }) - case 'extension:announce': { - const p = peers.get(peer.id) - if (!p) { - return - } - - if (authToken && !p.authenticated) { - send(peer, RESPONSES.error(ServerErrorMessages.mustAuthenticateBeforeAnnouncing)) - - return - } - - if (!isExtensionIdentity(event.data.identity)) { - send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceIdentityInvalid)) - - return - } - - p.extensionIdentity = event.data.identity - - send(peer, { + for (const other of peers.values()) { + if (other.authenticated && !(other.peer.id === peer.id)) { + send(other.peer, { type: 'extension:announced', data: event.data, metadata: createEventMetadata(instanceId, event.metadata?.event.id), }) - - for (const other of peers.values()) { - if (other.authenticated && !(other.peer.id === peer.id)) { - send(other.peer, { - type: 'extension:announced', - data: event.data, - metadata: createEventMetadata(instanceId, event.metadata?.event.id), - }) - } - } - - return } + } - case 'extension:module:announce': { - const p = peers.get(peer.id) - if (!p) { - return - } + return + } - if (authToken && !p.authenticated) { - send(peer, RESPONSES.error(ServerErrorMessages.mustAuthenticateBeforeAnnouncing)) + case 'extension:module:announce': { + const p = peers.get(peer.id) + if (!p) { + return + } - return - } + if (authToken && !p.authenticated) { + send(peer, RESPONSES.error(ServerErrorMessages.mustAuthenticateBeforeAnnouncing)) - const { name, identity } = event.data - if (!name || typeof name !== 'string') { - send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceNameInvalid)) + return + } - return - } + const { name, identity } = event.data + if (!name || typeof name !== 'string') { + send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceNameInvalid)) - if (!isExtensionModuleIdentity(identity)) { - send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceIdentityInvalid)) + return + } - return - } + if (!isExtensionModuleIdentity(identity)) { + send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceIdentityInvalid)) - if (p.extensionIdentity && identity.extension.id !== p.extensionIdentity.id) { - send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceIdentityInvalid)) + return + } - return - } + if (p.extensionIdentity && identity.extension.id !== p.extensionIdentity.id) { + send(peer, RESPONSES.error(ServerErrorMessages.moduleAnnounceIdentityInvalid)) - p.extensionIdentity = identity.extension - registerExtensionModulePeer(p, { name, identity }) + return + } - send(peer, { + p.extensionIdentity = identity.extension + registerExtensionModulePeer(p, { name, identity }) + + send(peer, { + type: 'extension:module:announced', + data: event.data, + metadata: createEventMetadata(instanceId, event.metadata?.event.id), + }) + + for (const other of peers.values()) { + if (other.authenticated && !(other.peer.id === peer.id)) { + send(other.peer, { type: 'extension:module:announced', data: event.data, metadata: createEventMetadata(instanceId, event.metadata?.event.id), }) - - for (const other of peers.values()) { - if (other.authenticated && !(other.peer.id === peer.id)) { - send(other.peer, { - type: 'extension:module:announced', - data: event.data, - metadata: createEventMetadata(instanceId, event.metadata?.event.id), - }) - } - } - - return } + } - case 'ui:configure': { - const data = event.data as { - moduleName?: string - moduleIndex?: number - identity?: MetadataEventSource - config?: Record - } - const moduleName = data.moduleName ?? (isExtensionModuleIdentity(data.identity) ? data.identity.id : '') ?? '' - const moduleIndex = data.moduleIndex - const config = data.config + return + } - if (moduleName === '') { - send(peer, RESPONSES.error(ServerErrorMessages.uiConfigureModuleNameInvalid)) + case 'ui:configure': { + const data = event.data as { + moduleName?: string + moduleIndex?: number + identity?: MetadataEventSource + config?: Record + } + const moduleName = data.moduleName ?? (isExtensionModuleIdentity(data.identity) ? data.identity.id : '') ?? '' + const moduleIndex = data.moduleIndex + const config = data.config - return - } - if (typeof moduleIndex !== 'undefined') { - if (!Number.isInteger(moduleIndex) || moduleIndex < 0) { - send(peer, RESPONSES.error(ServerErrorMessages.uiConfigureModuleIndexInvalid)) + if (moduleName === '') { + send(peer, RESPONSES.error(ServerErrorMessages.uiConfigureModuleNameInvalid)) - return - } - } + return + } + if (typeof moduleIndex !== 'undefined') { + if (!Number.isInteger(moduleIndex) || moduleIndex < 0) { + send(peer, RESPONSES.error(ServerErrorMessages.uiConfigureModuleIndexInvalid)) - const target = findModulePeer(moduleName, moduleIndex, data.identity) - if (target) { - send(target.peer, { - type: 'module:configure', - data: { config: config || {} }, - // NOTICE: this will forward the original event metadata as-is - metadata: event.metadata, - }) - } - else { - send(peer, RESPONSES.error(ServerErrorMessages.moduleNotFound)) - } - - return - } - - case 'module:consumer:register': { - const p = peers.get(peer.id) - if (!p?.authenticated) { - send(peer, RESPONSES.notAuthenticated(event.metadata?.event.id)) - return - } - - const data = event.data as { - event?: string - mode?: 'consumer' | 'consumer-group' - group?: string - priority?: number - } - - if (!data.event || typeof data.event !== 'string') { - send(peer, RESPONSES.error(ServerErrorMessages.moduleConsumerEventInvalid, event.metadata?.event.id)) - return - } - - registerConsumer( - peer.id, - data.event, - normalizeConsumerMode(data.mode, data.group), - data.group, - normalizeConsumerPriority(data.priority), - ) - return - } - - case 'module:consumer:unregister': { - const p = peers.get(peer.id) - if (!p?.authenticated) { - send(peer, RESPONSES.notAuthenticated(event.metadata?.event.id)) - return - } - - const data = event.data as { - event?: string - mode?: 'consumer' | 'consumer-group' - group?: string - } - - if (!data.event || typeof data.event !== 'string') { - send(peer, RESPONSES.error(ServerErrorMessages.moduleConsumerEventInvalid, event.metadata?.event.id)) - return - } - - unregisterConsumer(peer.id, data.event, normalizeConsumerMode(data.mode, data.group), data.group) return } } - // default case + const target = findModulePeer(moduleName, moduleIndex, data.identity) + if (target) { + send(target.peer, { + type: 'module:configure', + data: { config: config || {} }, + // NOTICE: this will forward the original event metadata as-is + metadata: event.metadata, + }) + } + else { + send(peer, RESPONSES.error(ServerErrorMessages.moduleNotFound)) + } + + return + } + + case 'module:consumer:register': { const p = peers.get(peer.id) if (!p?.authenticated) { - logger.withFields({ peer: peer.id, peerName: p?.name, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).debug('not authenticated') send(peer, RESPONSES.notAuthenticated(event.metadata?.event.id)) - return } - const payload = stringifyEvent(event) - const allowBypass = options?.routing?.allowBypass !== false - const shouldBypass = Boolean(event.route?.bypass && allowBypass && isDevtoolsPeer(p)) - const destinations = shouldBypass ? undefined : collectDestinations(event) - const delivery = shouldBypass ? undefined : resolveEventDelivery(event) - const effectiveRoutingMiddleware = shouldBypass ? [] : routingMiddleware - const decision = forEachEventMiddlewares({ - event, - fromPeer: p, - peers, - destinations, - middleware: effectiveRoutingMiddleware, - }) + const data = event.data as { + event?: string + mode?: 'consumer' | 'consumer-group' + group?: string + priority?: number + } - if (decision?.type === 'drop') { - logger.withFields({ peer: peer.id, peerName: p.name, event }).debug('routing dropped event') + if (!data.event || typeof data.event !== 'string') { + send(peer, RESPONSES.error(ServerErrorMessages.moduleConsumerEventInvalid, event.metadata?.event.id)) return } - const selectedConsumer = selectConsumer(event, peer.id, delivery) - if (delivery && (delivery.mode === 'consumer' || delivery.mode === 'consumer-group')) { - if (!selectedConsumer) { - logger.withFields({ peer: peer.id, peerName: p.name, event, delivery }).warn('no consumer registered for event delivery') - if (delivery.required) { - send(peer, RESPONSES.error(ServerErrorMessages.noConsumerRegistered, event.metadata?.event.id)) - } - return - } - - try { - logger.withFields({ - fromPeer: peer.id, - fromPeerName: p.name, - toPeer: selectedConsumer.peer.id, - toPeerName: selectedConsumer.name, - event, - delivery, - }).debug('sending event to selected consumer') - - selectedConsumer.peer.send(payload) - } - catch (err) { - logger.withFields({ - fromPeer: peer.id, - fromPeerName: p.name, - toPeer: selectedConsumer.peer.id, - toPeerName: selectedConsumer.name, - event, - delivery, - }).withError(err).error('failed to send event to selected consumer, removing peer') - - peers.delete(selectedConsumer.peer.id) - unregisterModulePeer(selectedConsumer, 'consumer send failed') - } - return - } - - const targetIds = decision?.type === 'targets' ? decision.targetIds : undefined - const shouldBroadcast = decision?.type === 'broadcast' || !targetIds - - logger.withFields({ peer: peer.id, peerName: p.name, event }).debug('broadcasting event to peers') - - for (const [id, other] of peers.entries()) { - if (id === peer.id) { - logger.withFields({ peer: peer.id, peerName: p.name, event }).debug('not sending event to self') - continue - } - - if (!other.authenticated) { - logger.withFields({ fromPeer: peer.id, toPeer: other.peer.id, toPeerName: other.name, event }).debug('not sending event to unauthenticated peer') - continue - } - - if (!shouldBroadcast && targetIds && !targetIds.has(id)) { - continue - } - - if (shouldBroadcast && destinations !== undefined && !matchesDestinations(destinations, other)) { - continue - } - - try { - logger.withFields({ fromPeer: peer.id, fromPeerName: p.name, toPeer: other.peer.id, toPeerName: other.name, event }).debug('sending event to peer') - other.peer.send(payload) - } - catch (err) { - logger.withFields({ fromPeer: peer.id, fromPeerName: p.name, toPeer: other.peer.id, toPeerName: other.name, event }).withError(err).error('failed to send event to peer, removing peer') - logger.withFields({ peer: peer.id, peerName: other.name }).debug('removing closed peer') - peers.delete(id) - - unregisterModulePeer(other, 'send failed') - } - } - }, - error: (peer, error) => { - logger.withFields({ peer: peer.id }).withError(error).error('an error occurred') - }, - close: (peer, details) => { - const p = peers.get(peer.id) - const now = Date.now() - const peerName = p?.name - const peerIndex = p?.index - const peerHealthy = p?.healthy - const peerMissedHeartbeats = p?.missedHeartbeats - const safeDetails = details ?? {} - const closeCode = typeof safeDetails.code === 'number' ? safeDetails.code : undefined - const closeReason = typeof safeDetails.reason === 'string' ? safeDetails.reason : undefined - const closeWasClean = typeof (safeDetails as { wasClean?: unknown }).wasClean === 'boolean' - ? (safeDetails as { wasClean?: unknown }).wasClean - : undefined - const heartbeatLastSeenAt = p?.lastHeartbeatAt - const heartbeatSilentForMs = heartbeatLastSeenAt ? now - heartbeatLastSeenAt : undefined - const likelyHeartbeatExpiry = Boolean( - p - && typeof heartbeatSilentForMs === 'number' - && heartbeatSilentForMs > heartbeatTtlMs, + registerConsumer( + peer.id, + data.event, + normalizeConsumerMode(data.mode, data.group), + data.group, + normalizeConsumerPriority(data.priority), ) - const likelySilentNetworkClose = closeCode === 1005 + return + } - if (p) { - peers.delete(peer.id) - unregisterModulePeer(p, 'connection closed') + case 'module:consumer:unregister': { + const p = peers.get(peer.id) + if (!p?.authenticated) { + send(peer, RESPONSES.notAuthenticated(event.metadata?.event.id)) + return } + const data = event.data as { + event?: string + mode?: 'consumer' | 'consumer-group' + group?: string + } + + if (!data.event || typeof data.event !== 'string') { + send(peer, RESPONSES.error(ServerErrorMessages.moduleConsumerEventInvalid, event.metadata?.event.id)) + return + } + + unregisterConsumer(peer.id, data.event, normalizeConsumerMode(data.mode, data.group), data.group) + return + } + } + + // default case + const p = peers.get(peer.id) + if (!p?.authenticated) { + logger.withFields({ peer: peer.id, peerName: p?.name, peerRemote: peer.remoteAddress, peerRequest: peer.request?.url }).debug('not authenticated') + send(peer, RESPONSES.notAuthenticated(event.metadata?.event.id)) + + return + } + + const payload = stringifyEvent(event) + const allowBypass = options?.routing?.allowBypass !== false + const shouldBypass = Boolean(event.route?.bypass && allowBypass && isDevtoolsPeer(p)) + const destinations = shouldBypass ? undefined : collectDestinations(event) + const delivery = shouldBypass ? undefined : resolveEventDelivery(event) + const effectiveRoutingMiddleware = shouldBypass ? [] : routingMiddleware + const decision = forEachEventMiddlewares({ + event, + fromPeer: p, + peers, + destinations, + middleware: effectiveRoutingMiddleware, + }) + + if (decision?.type === 'drop') { + logger.withFields({ peer: peer.id, peerName: p.name, event }).debug('routing dropped event') + return + } + + const selectedConsumer = selectConsumer(event, peer.id, delivery) + if (delivery && (delivery.mode === 'consumer' || delivery.mode === 'consumer-group')) { + if (!selectedConsumer) { + logger.withFields({ peer: peer.id, peerName: p.name, event, delivery }).warn('no consumer registered for event delivery') + if (delivery.required) { + send(peer, RESPONSES.error(ServerErrorMessages.noConsumerRegistered, event.metadata?.event.id)) + } + return + } + + try { logger.withFields({ - peer: peer.id, - peerRemote: peer.remoteAddress, - details, - closeCode, - closeReason, - closeWasClean, - activePeers: peers.size, - peerAuthenticated: p?.authenticated, - peerName, - peerIndex, - peerHealthy, - peerMissedHeartbeats, - heartbeatLastSeenAt, - heartbeatSilentForMs, - heartbeatTtlMs, - healthCheckIntervalMs, - likelyHeartbeatExpiry, - likelySilentNetworkClose, - }).log('closed') - }, - }, - dispose: () => { - clearInterval(healthCheckInterval) - closeAllPeers() - resetRoutingState(true) - }, + fromPeer: peer.id, + fromPeerName: p.name, + toPeer: selectedConsumer.peer.id, + toPeerName: selectedConsumer.name, + event, + delivery, + }).debug('sending event to selected consumer') + + selectedConsumer.peer.send(payload) + } + catch (err) { + logger.withFields({ + fromPeer: peer.id, + fromPeerName: p.name, + toPeer: selectedConsumer.peer.id, + toPeerName: selectedConsumer.name, + event, + delivery, + }).withError(err).error('failed to send event to selected consumer, removing peer') + + removeFailedPeer(selectedConsumer, 'consumer send failed') + } + return + } + + const targetIds = decision?.type === 'targets' ? decision.targetIds : undefined + const shouldBroadcast = decision?.type === 'broadcast' || !targetIds + + logger.withFields({ peer: peer.id, peerName: p.name, event }).debug('broadcasting event to peers') + + for (const [id, other] of peers.entries()) { + if (id === peer.id) { + logger.withFields({ peer: peer.id, peerName: p.name, event }).debug('not sending event to self') + continue + } + + if (!other.authenticated) { + logger.withFields({ fromPeer: peer.id, toPeer: other.peer.id, toPeerName: other.name, event }).debug('not sending event to unauthenticated peer') + continue + } + + if (!shouldBroadcast && targetIds && !targetIds.has(id)) { + continue + } + + if (shouldBroadcast && destinations !== undefined && !matchesDestinations(destinations, other)) { + continue + } + + try { + logger.withFields({ fromPeer: peer.id, fromPeerName: p.name, toPeer: other.peer.id, toPeerName: other.name, event }).debug('sending event to peer') + other.peer.send(payload) + } + catch (err) { + logger.withFields({ fromPeer: peer.id, fromPeerName: p.name, toPeer: other.peer.id, toPeerName: other.name, event }).withError(err).error('failed to send event to peer, removing peer') + logger.withFields({ peer: peer.id, peerName: other.name }).debug('removing closed peer') + removeFailedPeer(other, 'send failed') + } + } }) - app.get('/ws', defineWebSocketHandler(websocketGateway.handler)) + function handlePeerError(peer: Peer, error: unknown) { + logger.withFields({ peer: peer.id }).withError(error).error('an error occurred') + } + + function handlePeerClose(peer: Peer, details?: WsCloseDetails) { + const p = peers.get(peer.id) + const now = Date.now() + const peerName = p?.name + const peerIndex = p?.index + const peerHealthy = p?.healthy + const peerSilentFor = p?.missedHeartbeats + const safeDetails = details ?? {} + const closeCode = typeof safeDetails.code === 'number' ? safeDetails.code : undefined + const closeReason = typeof safeDetails.reason === 'string' ? safeDetails.reason : undefined + const closeWasClean = typeof (safeDetails as { wasClean?: unknown }).wasClean === 'boolean' + ? (safeDetails as { wasClean?: unknown }).wasClean + : undefined + const heartbeatLastSeenAt = p?.lastHeartbeatAt + const heartbeatSilentForMs = heartbeatLastSeenAt != null ? now - heartbeatLastSeenAt : undefined + const likelyHeartbeatExpiry = Boolean( + p + && typeof heartbeatSilentForMs === 'number' + && heartbeatSilentForMs > heartbeatTtlMs, + ) + const likelySilentNetworkClose = closeCode === 1005 + + const unregisterReason = likelyHeartbeatExpiry ? 'heartbeat expired' : closeReason === 'server shutdown' ? 'server shutdown' : 'connection closed' + + if (p) { + peers.delete(peer.id) + unregisterModulePeer(p, unregisterReason) + } + + logger.withFields({ + peer: peer.id, + peerRemote: peer.remoteAddress, + details, + closeCode, + closeReason, + closeWasClean, + activePeers: peers.size, + peerAuthenticated: p?.authenticated, + peerName, + peerIndex, + peerHealthy, + peerMissedHeartbeats: p?.missedHeartbeats, + peerSilentFor, + heartbeatLastSeenAt, + heartbeatSilentForMs, + heartbeatTtlMs, + healthCheckIntervalMs, + likelyHeartbeatExpiry, + likelySilentNetworkClose, + }).log('closed') + } + + function removeFailedPeer(peerInfo: AuthenticatedPeer, reason: string) { + const managedPeer = wsServer.peers.get(peerInfo.peer.id) + if (managedPeer) { + wsServer.remove(managedPeer.id, { reason }) + return + } + + handlePeerClose(peerInfo.peer, { reason }) + } + + wsServer.onPeerClose(({ peerId, details }) => { + const peerInfo = peers.get(peerId) + if (!peerInfo) { + return + } + + handlePeerClose(peerInfo.peer, details) + }) + + wsServer.onPeerHealthChange(({ peer, healthy, silentFor }) => { + const peerInfo = peers.get(peer.id) + if (!peerInfo) { + return + } + + // REVIEW: better-ws now reports silence duration in milliseconds, while the + // AIRI runtime peer state still exposes the legacy missedHeartbeats field. + // Rename this business-facing field with the server-runtime state cleanup. + peerInfo.missedHeartbeats = silentFor + + if (healthy) { + peerInfo.healthy = true + logger.withFields({ peer: peer.id, peerName: peerInfo.name }).debug('peer activity recovered, marking healthy') + broadcastPeerHealthy(peerInfo) + + return + } + + peerInfo.healthy = false + logger.withFields({ peer: peer.id, peerName: peerInfo.name, silentFor }).debug('heartbeat late, marking unhealthy') + broadcastPeerUnhealthy(peerInfo, 'heartbeat late') + }) + + function unregisterClosedLivenessPeers() { + for (const [id, peerInfo] of peers.entries()) { + if (wsServer.peers.has(id)) { + continue + } + + logger.withFields({ peer: id, peerName: peerInfo.name, silentFor: peerInfo.missedHeartbeats }).debug('heartbeat silent timeout expired, dropping peer') + peers.delete(id) + unregisterModulePeer(peerInfo, 'heartbeat expired') + } + } + + let healthCheckInterval: ReturnType | undefined = setInterval(() => { + try { + wsServer.checkLiveness(Date.now()) + } + catch (error) { + logger.withError(error as Error).debug('websocket liveness check failed while closing expired peers') + } + unregisterClosedLivenessPeers() + }, healthCheckIntervalMs) + if (typeof healthCheckInterval === 'object') { + healthCheckInterval.unref?.() + } + + function clearHealthCheckInterval() { + if (!healthCheckInterval) { + return + } + + clearInterval(healthCheckInterval) + healthCheckInterval = undefined + } + + app.get('/ws', toH3Handler(wsServer, { + readMessage(message: CrossWsMessage) { + return { text: () => message.text() } + }, + state(rawPeer: CrossWsPeer) { + return { rawPeer } + }, + error({ peer, rawPeer, error }) { + handlePeerError(peer ? rawPeerFrom(peer) ?? airiPeerFromRaw(rawPeer) : airiPeerFromRaw(rawPeer), error) + }, + })) function closeAllPeers() { logger.withFields({ totalPeers: peers.size }).log('closing all peers') @@ -1114,30 +1121,25 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => peerName: peerInfo.name, }).debug('closing peer') - try { - peerInfo.peer.close?.() - } - catch (error) { - logger - .withFields({ - peer: peerInfo.peer.id, - peerName: peerInfo.name, - }) - .withError(error as Error) - .debug('failed to close peer during shutdown') - - // Leave the peer registered until forced disposal cleanup. + const managedPeer = wsServer.peers.get(peerInfo.peer.id) + if (managedPeer) { + try { + managedPeer.close(undefined, 'server shutdown') + } + catch (error) { + logger + .withFields({ + peer: peerInfo.peer.id, + peerName: peerInfo.name, + }) + .withError(error as Error) + .debug('failed to close peer during shutdown') + } continue } - // Some websocket runtimes may never emit `close` - // during abrupt shutdown sequences. Remove peers - // synchronously after initiating a successful close - // so shutdown cleanup is deterministic. - peers.delete(peerInfo.peer.id) - try { - unregisterModulePeer(peerInfo, 'server shutdown') + handlePeerClose(peerInfo.peer, { reason: 'server shutdown' }) } catch (error) { logger @@ -1157,7 +1159,10 @@ export function setupApp(options?: AppOptions): { app: H3, closeAllPeers: () => } disposed = true - websocketGateway.dispose() + clearHealthCheckInterval() + closeAllPeers() + wsServer.close() + resetRoutingState(true) } return { diff --git a/packages/server-runtime/src/server-ws/airi/index.test.ts b/packages/server-runtime/src/server-ws/airi/codec.test.ts similarity index 50% rename from packages/server-runtime/src/server-ws/airi/index.test.ts rename to packages/server-runtime/src/server-ws/airi/codec.test.ts index de27d48bb..a7e73f529 100644 --- a/packages/server-runtime/src/server-ws/airi/index.test.ts +++ b/packages/server-runtime/src/server-ws/airi/codec.test.ts @@ -1,16 +1,16 @@ import type { WebSocketEvent } from '@proj-airi/server-shared/types' -import { stringify } from 'superjson' +import { stringify as stringifySuperJson } from 'superjson' import { describe, expect, it } from 'vitest' import { - AiriWebSocketEventFormatError, - createResponses, heartbeatFrameFrom, + InvalidEventError, parseEvent, -} from '.' + stringifyEvent, +} from './codec' -describe('airi websocket protocol codec', () => { +describe('airi websocket codec', () => { it('parses superjson encoded events', () => { const event: WebSocketEvent = { type: 'module:authenticate', @@ -25,7 +25,7 @@ describe('airi websocket protocol codec', () => { }, } - expect(parseEvent(stringify(event))).toEqual(event) + expect(parseEvent(stringifySuperJson(event))).toEqual(event) }) it('falls back to plain JSON events', () => { @@ -45,61 +45,48 @@ describe('airi websocket protocol codec', () => { expect(parseEvent(JSON.stringify(event))).toEqual(event) }) - it('rejects payloads without event type', () => { + it('rejects invalid event envelopes', () => { expect(() => parseEvent('null')) - .toThrow(AiriWebSocketEventFormatError) + .toThrow(InvalidEventError) expect(() => parseEvent(JSON.stringify({ data: {} }))) - .toThrow(AiriWebSocketEventFormatError) - }) - - it('rejects payloads with non-string event type', () => { + .toThrow(InvalidEventError) expect(() => parseEvent(JSON.stringify({ type: 0, data: {} }))) - .toThrow(AiriWebSocketEventFormatError) - }) - - it('rejects payloads without object event data', () => { + .toThrow(InvalidEventError) expect(() => parseEvent(JSON.stringify({ type: 'module:authenticate' }))) - .toThrow(AiriWebSocketEventFormatError) + .toThrow(InvalidEventError) expect(() => parseEvent(JSON.stringify({ type: 'module:authenticate', data: null }))) - .toThrow(AiriWebSocketEventFormatError) + .toThrow(InvalidEventError) expect(() => parseEvent(JSON.stringify({ type: 'module:authenticate', data: 'secret' }))) - .toThrow(AiriWebSocketEventFormatError) - }) - - it('rejects payloads with array event data', () => { + .toThrow(InvalidEventError) expect(() => parseEvent(JSON.stringify({ type: 'module:authenticate', data: [] }))) - .toThrow(AiriWebSocketEventFormatError) + .toThrow(InvalidEventError) }) - it('classifies raw ping and pong control frames', () => { + it('keeps validation cause and source on invalid event errors', () => { + const source = { type: 'module:authenticate', data: 'secret' } + + try { + parseEvent(JSON.stringify(source)) + expect.unreachable('Expected invalid event parsing to throw.') + } + catch (error) { + expect(error).toBeInstanceOf(InvalidEventError) + expect(error).toMatchObject({ source }) + expect((error as InvalidEventError).cause).toEqual(expect.arrayContaining([ + expect.objectContaining({ + message: 'Expected event data to be a non-array object.', + }), + ])) + } + }) + + it('detects raw ping and pong control frames', () => { expect(heartbeatFrameFrom('ping')).toBe('ping') expect(heartbeatFrameFrom('pong')).toBe('pong') expect(heartbeatFrameFrom('{"type":"ping"}')).toBeUndefined() }) - /** - * @example - * expect(responses.peerAuthenticated('peer-1').type).toBe('peer:authenticated') - * expect(responses.extensionAuthenticated({ id: 'airi-extension-chess' }).type).toBe('extension:authenticated') - */ - it('creates peer and extension authentication responses separately', () => { - const responses = createResponses('server-1') - - expect(responses.peerAuthenticated('peer-1')).toMatchObject({ - type: 'peer:authenticated', - data: { - authenticated: true, - peerId: 'peer-1', - }, - }) - expect(responses.extensionAuthenticated({ id: 'airi-extension-chess' })).toMatchObject({ - type: 'extension:authenticated', - data: { - authenticated: true, - identity: { - id: 'airi-extension-chess', - }, - }, - }) + it('preserves raw string events when stringifying', () => { + expect(stringifyEvent('raw-payload')).toBe('raw-payload') }) }) diff --git a/packages/server-runtime/src/server-ws/airi/codec.ts b/packages/server-runtime/src/server-ws/airi/codec.ts new file mode 100644 index 000000000..eeb426d03 --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/codec.ts @@ -0,0 +1,81 @@ +import type { WebSocketBaseEvent, WebSocketEvent } from '@proj-airi/server-shared/types' + +import { MessageHeartbeatKind } from '@proj-airi/server-shared/types' +import { parse, stringify } from 'superjson' +import { check, objectWithRest, pipe, safeParse, string, unknown } from 'valibot' + +const invalidAiriWebSocketEventFormatMessage = 'Invalid WebSocket event format.' + +const eventDataSchema = pipe( + unknown(), + check( + value => Boolean(value) && typeof value === 'object' && !Array.isArray(value), + 'Expected event data to be a non-array object.', + ), +) + +const eventEnvelopeSchema = objectWithRest({ + type: string(), + data: eventDataSchema, +}, unknown()) + +interface InvalidEventErrorOptions { + cause?: unknown + source?: unknown +} + +/** Error thrown when parsed websocket text is not an AIRI event envelope. */ +export class InvalidEventError extends Error { + readonly source?: unknown + + constructor(options: InvalidEventErrorOptions = {}) { + super(invalidAiriWebSocketEventFormatMessage, { cause: options.cause }) + this.name = 'InvalidEventError' + this.source = options.source + } +} + +/** Checks whether an error came from AIRI websocket event envelope validation. */ +export function isInvalidEventError(error: unknown): error is InvalidEventError { + return error instanceof InvalidEventError +} + +/** Detects raw ping/pong text frames that should not enter the event protocol. */ +export function heartbeatFrameFrom(text: string): MessageHeartbeatKind | undefined { + if (text === MessageHeartbeatKind.Ping || text === MessageHeartbeatKind.Pong) { + return text + } +} + +/** Parses one AIRI websocket protocol event from SuperJSON or plain JSON text. */ +export function parseEvent(text: string): WebSocketEvent { + // NOTICE: + // SDK clients send events using superjson.stringify, so websocket runtime code must + // use superjson.parse instead of message.json() or plain JSON.parse first. + // JSON.parse on a superjson-encoded string returns the wrapper object + // `{ json: {...}, meta: {...} }` with no protocol `type`, which breaks routing. + // Keep this until all AIRI websocket clients share one non-wrapper wire format. + let parsed: WebSocketEvent | undefined + try { + parsed = parse(text) + } + catch { + parsed = undefined + } + + const potentialEvent = (parsed && typeof parsed === 'object' && 'type' in parsed) + ? parsed + : JSON.parse(text) + + const result = safeParse(eventEnvelopeSchema, potentialEvent) + if (!result.success) { + throw new InvalidEventError({ cause: result.issues, source: potentialEvent }) + } + + return potentialEvent as WebSocketEvent +} + +/** Serializes one AIRI websocket protocol event with the existing SuperJSON wire format. */ +export function stringifyEvent(event: WebSocketBaseEvent | string) { + return typeof event === 'string' ? event : stringify(event) +} diff --git a/packages/server-runtime/src/server-ws/core/index.test.ts b/packages/server-runtime/src/server-ws/airi/consumers.test.ts similarity index 96% rename from packages/server-runtime/src/server-ws/core/index.test.ts rename to packages/server-runtime/src/server-ws/airi/consumers.test.ts index f956795d4..337b4e2cb 100644 --- a/packages/server-runtime/src/server-ws/core/index.test.ts +++ b/packages/server-runtime/src/server-ws/airi/consumers.test.ts @@ -1,13 +1,13 @@ -import type { ServerWsStickyAssignment } from '.' +import type { ConsumerStickyAssignment } from './consumers' import { describe, expect, it } from 'vitest' import { createConsumerOrchestrator, selectConsumerPeerId, -} from '.' +} from './consumers' -describe('server-ws consumer selection', () => { +describe('airi websocket consumer selection', () => { it('selects highest priority then earliest registration', () => { expect(selectConsumerPeerId({ eventType: 'event:test', @@ -36,7 +36,7 @@ describe('server-ws consumer selection', () => { }) it('preserves sticky assignment for the same sticky key', () => { - const stickyAssignments = new Map() + const stickyAssignments = new Map() const delivery = { mode: 'consumer-group' as const, group: 'workers', selection: 'sticky' as const, stickyKey: 'job-1' } const candidates = [ { peerId: 'a', priority: 0, registeredAt: 1, authenticated: true }, @@ -73,7 +73,7 @@ describe('server-ws consumer selection', () => { }) }) -describe('server-ws consumer registry', () => { +describe('airi websocket consumer registry', () => { it('registers and unregisters consumers', () => { const registry = createConsumerOrchestrator() @@ -106,7 +106,7 @@ describe('server-ws consumer registry', () => { }) it('keeps sticky assignments isolated for delimiter-like event and group names', () => { - const stickyAssignments = new Map() + const stickyAssignments = new Map() const candidates = [ { peerId: 'event::group-target', priority: 0, registeredAt: 1, authenticated: true }, { peerId: 'other-target', priority: 0, registeredAt: 2, authenticated: true }, diff --git a/packages/server-runtime/src/server-ws/airi/consumers.ts b/packages/server-runtime/src/server-ws/airi/consumers.ts new file mode 100644 index 000000000..f9003e276 --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/consumers.ts @@ -0,0 +1,301 @@ +import type { DeliveryConfig } from '@proj-airi/server-shared/types' + +const DEFAULT_CONSUMER_GROUP = 'default' + +interface ConsumerRegistryRef { + event: string + group: string +} + +/** + * Candidate peer metadata used for AIRI consumer selection. + */ +export interface ConsumerSelectionCandidate { + /** Peer id available to receive the event. */ + peerId: string + /** Higher values are selected before lower values. */ + priority: number + /** Timestamp captured when the peer registered as a consumer. */ + registeredAt: number + /** Whether the peer has completed protocol-level authentication. */ + authenticated: boolean + /** Explicit `false` excludes the peer from selection. */ + healthy?: boolean +} + +/** + * Stored AIRI consumer registration. + */ +export interface ConsumerRegistration { + /** Protocol event type consumed by the peer. */ + event: string + /** Normalized consumer group name. */ + group: string + /** Peer id that registered for the event/group pair. */ + peerId: string + /** Higher values are selected before lower values. */ + priority: number + /** Timestamp captured when the peer registered as a consumer. */ + registeredAt: number +} + +/** + * Sticky AIRI consumer assignment stored by the consumer selector. + */ +export interface ConsumerStickyAssignment { + /** Protocol event type the sticky assignment belongs to. */ + event: string + /** Normalized consumer group the sticky assignment belongs to. */ + group: string + /** Peer selected for the sticky key. */ + peerId: string +} + +/** + * Checks whether a delivery mode targets the AIRI consumer registry. + */ +export function isConsumerDeliveryMode(mode: unknown): mode is 'consumer' | 'consumer-group' { + return mode === 'consumer' || mode === 'consumer-group' +} + +/** + * Normalizes delivery mode for AIRI consumer registration. + * + * Before: + * - undefined with group "workers" + * + * After: + * - "consumer-group" + */ +export function normalizeConsumerMode(mode: unknown, group?: string): 'consumer' | 'consumer-group' { + if (isConsumerDeliveryMode(mode)) { + return mode + } + + return group ? 'consumer-group' : 'consumer' +} + +/** + * Normalizes AIRI consumer priority. + * + * Before: + * - NaN + * + * After: + * - 0 + */ +export function normalizeConsumerPriority(priority: unknown) { + return typeof priority === 'number' && Number.isFinite(priority) + ? priority + : 0 +} + +function normalizeConsumerGroup(mode: 'consumer' | 'consumer-group', group?: string) { + if (mode === 'consumer') { + return DEFAULT_CONSUMER_GROUP + } + + return group || DEFAULT_CONSUMER_GROUP +} + +function sortConsumers(entries: Array>) { + return [...entries].sort((left, right) => { + if (right.priority !== left.priority) { + return right.priority - left.priority + } + + return left.registeredAt - right.registeredAt + }) +} + +/** + * Selects a concrete peer for AIRI consumer-style delivery modes. + * + * Sticky and round-robin state are keyed with structured JSON tuples so event, + * group, and sticky key values may contain delimiter-like text safely. + */ +export function selectConsumerPeerId(options: { + eventType: string + fromPeerId: string + delivery?: DeliveryConfig + candidates: ConsumerSelectionCandidate[] + roundRobinCursor?: Map + stickyAssignments?: Map +}) { + const { candidates, delivery, eventType, fromPeerId } = options + if (!delivery || !isConsumerDeliveryMode(delivery.mode)) { + return + } + + const normalizedGroup = normalizeConsumerGroup(delivery.mode, delivery.group) + const registryKey = JSON.stringify([eventType, normalizedGroup]) + const availableEntries = sortConsumers( + candidates + .filter(entry => entry.peerId !== fromPeerId) + .filter(entry => entry.authenticated && entry.healthy !== false), + ) + + if (availableEntries.length === 0) { + return + } + + const selection = delivery.selection ?? 'first' + if (selection === 'sticky' && delivery.stickyKey) { + const stickyRegistryKey = JSON.stringify([eventType, normalizedGroup, delivery.stickyKey]) + const stickyAssignment = options.stickyAssignments?.get(stickyRegistryKey) + if (stickyAssignment && stickyAssignment.peerId !== fromPeerId) { + const stickyCandidate = availableEntries.find(entry => entry.peerId === stickyAssignment.peerId) + if (stickyCandidate) { + return stickyAssignment.peerId + } + } + + const selected = availableEntries[0] + options.stickyAssignments?.set(stickyRegistryKey, { event: eventType, group: normalizedGroup, peerId: selected.peerId }) + return selected.peerId + } + + if (selection === 'round-robin') { + const cursor = options.roundRobinCursor?.get(registryKey) ?? 0 + const selected = availableEntries[cursor % availableEntries.length] + options.roundRobinCursor?.set(registryKey, (cursor + 1) % availableEntries.length) + return selected.peerId + } + + return availableEntries[0].peerId +} + +/** + * Creates the AIRI consumer delivery orchestrator for websocket peers. + * + * The orchestrator owns registration, unregister, listing, selection, and + * sticky/round-robin cleanup state for AIRI consumer routing. + */ +export function createConsumerOrchestrator() { + const consumerRegistry = new Map>>() + const consumerKeysByPeer = new Map>() + const deliveryRoundRobinCursor = new Map() + const stickyAssignments = new Map() + + function removeStickyAssignmentsFor(event: string, group: string, peerId?: string) { + for (const [stickyKey, assignment] of stickyAssignments.entries()) { + if (peerId && assignment.peerId !== peerId) { + continue + } + + if (assignment.event === event && assignment.group === group) { + stickyAssignments.delete(stickyKey) + } + } + } + + return { + register(input: { peerId: string, event: string, mode: 'consumer' | 'consumer-group', group?: string, priority?: number }) { + const normalizedGroup = normalizeConsumerGroup(input.mode, input.group) + const registryKey = JSON.stringify([input.event, normalizedGroup]) + let groups = consumerRegistry.get(input.event) + if (!groups) { + groups = new Map() + consumerRegistry.set(input.event, groups) + } + + let peersForGroup = groups.get(normalizedGroup) + if (!peersForGroup) { + peersForGroup = new Map() + groups.set(normalizedGroup, peersForGroup) + } + + const didGrowMembership = !peersForGroup.has(input.peerId) + peersForGroup.set(input.peerId, { + event: input.event, + group: normalizedGroup, + peerId: input.peerId, + priority: normalizeConsumerPriority(input.priority), + registeredAt: Date.now(), + }) + if (didGrowMembership) { + deliveryRoundRobinCursor.delete(registryKey) + } + + let registrations = consumerKeysByPeer.get(input.peerId) + if (!registrations) { + registrations = new Map() + consumerKeysByPeer.set(input.peerId, registrations) + } + registrations.set(registryKey, { event: input.event, group: normalizedGroup }) + }, + unregister(input: { peerId: string, event: string, mode: 'consumer' | 'consumer-group', group?: string }) { + const normalizedGroup = normalizeConsumerGroup(input.mode, input.group) + const registryKey = JSON.stringify([input.event, normalizedGroup]) + const groups = consumerRegistry.get(input.event) + const peersForGroup = groups?.get(normalizedGroup) + const didDelete = peersForGroup?.delete(input.peerId) ?? false + + if (!didDelete) { + return + } + + deliveryRoundRobinCursor.delete(registryKey) + if (peersForGroup?.size === 0) { + groups?.delete(normalizedGroup) + } + if (groups?.size === 0) { + consumerRegistry.delete(input.event) + } + + const registrations = consumerKeysByPeer.get(input.peerId) + registrations?.delete(registryKey) + if (registrations?.size === 0) { + consumerKeysByPeer.delete(input.peerId) + } + + removeStickyAssignmentsFor(input.event, normalizedGroup, input.peerId) + }, + unregisterPeer(peerId: string) { + const registrations = consumerKeysByPeer.get(peerId) + if (!registrations?.size) { + return + } + + for (const registration of registrations.values()) { + const { event, group } = registration + const groups = consumerRegistry.get(event) + const peersForGroup = groups?.get(group) + peersForGroup?.delete(peerId) + deliveryRoundRobinCursor.delete(JSON.stringify([event, group])) + if (peersForGroup?.size === 0) { + groups?.delete(group) + } + if (groups?.size === 0) { + consumerRegistry.delete(event) + } + + removeStickyAssignmentsFor(event, group, peerId) + } + + consumerKeysByPeer.delete(peerId) + }, + listFor(input: { event: string, mode: 'consumer' | 'consumer-group', group?: string }) { + const normalizedGroup = normalizeConsumerGroup(input.mode, input.group) + return [...consumerRegistry.get(input.event)?.get(normalizedGroup)?.values() ?? []] + }, + select(input: { + eventType: string + fromPeerId: string + delivery?: DeliveryConfig + candidates: ConsumerSelectionCandidate[] + }) { + return selectConsumerPeerId({ + ...input, + roundRobinCursor: deliveryRoundRobinCursor, + stickyAssignments, + }) + }, + clear() { + consumerRegistry.clear() + consumerKeysByPeer.clear() + deliveryRoundRobinCursor.clear() + stickyAssignments.clear() + }, + } +} diff --git a/packages/server-runtime/src/server-ws/airi/index.ts b/packages/server-runtime/src/server-ws/airi/index.ts index 00b4b0c80..bc0219a26 100644 --- a/packages/server-runtime/src/server-ws/airi/index.ts +++ b/packages/server-runtime/src/server-ws/airi/index.ts @@ -1,354 +1,27 @@ -import type { DeliveryConfig, ExtensionIdentity, MessageHeartbeat, MetadataEventSource, WebSocketBaseEvent, WebSocketEvent } from '@proj-airi/server-shared/types' - -import type { - RouteContext, - RouteDecision, - RouteMiddleware, -} from '../../middlewares' -import type { Peer } from '../../types' - -import { ServerErrorMessages } from '@proj-airi/server-shared' -import { - getProtocolEventMetadata, - MessageHeartbeatKind, - WebSocketEventSource, -} from '@proj-airi/server-shared/types' -import { nanoid } from 'nanoid' -import { parse, stringify } from 'superjson' - -import packageJSON from '../../../package.json' - -import { createEventCodec, createGatewayLifecycle } from '../core' - -const invalidAiriWebSocketEventFormatMessage = 'Invalid WebSocket event format.' - -/** - * Close details surfaced by the websocket runtime for AIRI peer shutdown logging. - */ -export interface AiriServerWsCloseDetails { - /** WebSocket close code when the runtime reports one. */ - code?: number - /** WebSocket close reason when the runtime reports one. */ - reason?: string - /** Whether the runtime considers the close clean. */ - wasClean?: unknown -} - -/** - * Error thrown when a websocket message parses as JSON but is not an AIRI event envelope. - * - * Use when: - * - The runtime must distinguish malformed event envelopes from invalid JSON text - * - * Expects: - * - Callers convert this to the protocol `invalidEventFormat` response - * - * Returns: - * - A typed error for invalid AIRI websocket event envelopes - */ -export class AiriWebSocketEventFormatError extends Error { - constructor() { - super(invalidAiriWebSocketEventFormatMessage) - this.name = 'AiriWebSocketEventFormatError' - } -} - -/** - * Creates the AIRI websocket gateway wrapper. - * - * Use when: - * - `setupApp(...)` needs a gateway object to mount on `/ws` - * - * Expects: - * - `handler` preserves the existing AIRI websocket lifecycle behavior - * - * Returns: - * - A gateway object compatible with H3 `defineWebSocketHandler(...)` - */ -export function createGateway(input: { - handler: { - open: (peer: Peer) => void - message: (peer: Peer, message: { text: () => string }) => void - error: (peer: Peer, error: unknown) => void - close: (peer: Peer, details?: AiriServerWsCloseDetails) => void - } - dispose?: () => void -}) { - return createGatewayLifecycle({ - handler: input.handler, - dispose: input.dispose, - }) -} - -/** - * Creates metadata for events emitted by the AIRI websocket runtime. - * - * Use when: - * - The server sends protocol events to connected peers - * - Response events should preserve parent event correlation - * - * Expects: - * - `serverInstanceId` identifies the active server runtime instance - * - * Returns: - * - AIRI protocol metadata with server source and event id - */ -export function createEventMetadata( - serverInstanceId: string, - parentId?: string, -): { source: MetadataEventSource, event: { id: string, parentId?: string } } { - return { - event: { - id: nanoid(), - parentId, - }, - source: { - kind: 'plugin', - plugin: { - id: WebSocketEventSource.Server, - version: packageJSON.version, - }, - id: serverInstanceId, - }, - } -} - -/** - * Creates AIRI server response event factories. - * - * Use when: - * - WebSocket handlers need stable response event shapes - * - * Expects: - * - `serverInstanceId` identifies the current server runtime - * - * Returns: - * - Factory methods for protocol responses emitted by the server - */ -export function createResponses(serverInstanceId: string) { - return { - authenticated(parentId?: string) { - return { - type: 'module:authenticated', - data: { authenticated: true }, - metadata: createEventMetadata(serverInstanceId, parentId), - } satisfies WebSocketEvent> - }, - peerAuthenticated(peerId: string, parentId?: string) { - return { - type: 'peer:authenticated', - data: { authenticated: true, peerId }, - metadata: createEventMetadata(serverInstanceId, parentId), - } satisfies WebSocketEvent> - }, - extensionAuthenticated(identity: ExtensionIdentity, parentId?: string) { - return { - type: 'extension:authenticated', - data: { identity, authenticated: true }, - metadata: createEventMetadata(serverInstanceId, parentId), - } satisfies WebSocketEvent> - }, - notAuthenticated(parentId?: string) { - return { - type: 'error', - data: { message: ServerErrorMessages.notAuthenticated }, - metadata: createEventMetadata(serverInstanceId, parentId), - } satisfies WebSocketEvent> - }, - error(message: string, parentId?: string) { - return { - type: 'error', - data: { message }, - metadata: createEventMetadata(serverInstanceId, parentId), - } satisfies WebSocketEvent> - }, - heartbeat(kind: MessageHeartbeatKind, message: MessageHeartbeat | string, parentId?: string) { - return { - type: 'transport:connection:heartbeat', - data: { kind, message, at: Date.now() }, - metadata: createEventMetadata(serverInstanceId, parentId), - } satisfies WebSocketEvent> - }, - } -} - -/** - * Checks whether an error came from AIRI websocket event envelope validation. - * - * Use when: - * - Message handlers need to map invalid envelopes to protocol errors - * - * Expects: - * - Parser code throws {@link AiriWebSocketEventFormatError} for envelope failures - * - * Returns: - * - `true` when the error should become `ServerErrorMessages.invalidEventFormat` - */ -export function isAiriWebSocketEventFormatError(error: unknown): error is AiriWebSocketEventFormatError { - return error instanceof AiriWebSocketEventFormatError -} - -/** - * Detects raw websocket heartbeat control frames surfaced as text payloads. - * - * Use when: - * - A websocket runtime forwards ping/pong frames through the normal message callback - * - The runtime should ignore transport heartbeats instead of treating them as protocol JSON - * - * Expects: - * - Raw text payloads such as `ping` and `pong` - * - * Returns: - * - The heartbeat kind when the text is a control frame, otherwise `undefined` - */ -export function heartbeatFrameFrom(text: string): MessageHeartbeatKind | undefined { - if (text === MessageHeartbeatKind.Ping || text === MessageHeartbeatKind.Pong) { - return text - } -} - -/** - * Parses one AIRI websocket protocol event. - * - * Use when: - * - Reading text messages from WebSocket peers - * - * Expects: - * - SDK clients may send `superjson.stringify(...)` - * - External clients may send plain JSON - * - * Returns: - * - A WebSocket event with a string `type` - */ -export function parseEvent(text: string): WebSocketEvent { - // NOTICE: - // SDK clients send events using superjson.stringify, so websocket runtime code must - // use superjson.parse instead of message.json() or plain JSON.parse first. - // JSON.parse on a superjson-encoded string returns the wrapper object - // `{ json: {...}, meta: {...} }` with no protocol `type`, which breaks routing. - // Keep this until all AIRI websocket clients share one non-wrapper wire format. - let parsed: WebSocketEvent | undefined - try { - parsed = parse(text) - } - catch { - parsed = undefined - } - - const potentialEvent = (parsed && typeof parsed === 'object' && 'type' in parsed) - ? parsed - : JSON.parse(text) - - if ( - !potentialEvent - || typeof potentialEvent !== 'object' - || !('type' in potentialEvent) - || typeof potentialEvent.type !== 'string' - || !('data' in potentialEvent) - || !potentialEvent.data - || typeof potentialEvent.data !== 'object' - || Array.isArray(potentialEvent.data) - ) { - throw new AiriWebSocketEventFormatError() - } - - return potentialEvent as WebSocketEvent -} - -/** - * Serializes one AIRI websocket protocol event. - * - * Use when: - * - Sending AIRI events through WebSocket peers - * - * Expects: - * - `event` is already protocol-shaped - * - * Returns: - * - SuperJSON text payload matching existing runtime behavior - */ -export function stringifyEvent(event: WebSocketBaseEvent | string) { - return typeof event === 'string' ? event : stringify(event) -} - -/** - * Resolves the effective event delivery policy. - * - * Use when: - * - Protocol defaults should be merged with route-level delivery overrides - * - Routing needs to know whether the event should broadcast or target one consumer - * - * Expects: - * - Route delivery to override protocol metadata field-by-field - * - * Returns: - * - The merged broadcast/consumer delivery policy, or `undefined` when unrestricted - */ -export function resolveEventDelivery(event: WebSocketEvent): DeliveryConfig | undefined { - const eventMetadata = getProtocolEventMetadata(event.type) - const defaultDelivery = eventMetadata?.delivery - const routeDelivery = event.route?.delivery - - if (!defaultDelivery && !routeDelivery) { - return undefined - } - - return { - ...defaultDelivery, - ...routeDelivery, - } -} - -/** - * Creates event serializer hooks used by server websocket adapters. - * - * Use when: - * - A gateway wants protocol-specific parsing, stringifying, and control-frame detection - * - * Expects: - * - Callers route raw control frames before protocol events - * - * Returns: - * - A reusable `server-ws/core` codec configured for AIRI events - */ -export function createEventSerializer() { - return createEventCodec({ - parse: parseEvent, - stringify: stringifyEvent, - detectControlFrame: heartbeatFrameFrom, - }) -} - -/** - * Iterates event middlewares in declaration order until one returns a decision. - * - * Use when: - * - The websocket runtime needs the first route decision from configured middleware - * - * Expects: - * - Middleware functions are ordered by caller policy - * - * Returns: - * - The first route decision, or `undefined` when no middleware decided - */ -export function forEachEventMiddlewares(input: { - event: WebSocketEvent - fromPeer: RouteContext['fromPeer'] - peers: Map - destinations?: RouteContext['destinations'] - middleware: RouteMiddleware[] -}): RouteDecision | undefined { - const context: RouteContext = { - event: input.event, - fromPeer: input.fromPeer, - peers: input.peers, - destinations: input.destinations, - } - - for (const middleware of input.middleware) { - const result = middleware(context) - if (result) { - return result - } - } -} +export { + heartbeatFrameFrom, + InvalidEventError, + isInvalidEventError, + parseEvent, + stringifyEvent, +} from './codec' +export { + createConsumerOrchestrator, + isConsumerDeliveryMode, + normalizeConsumerMode, + normalizeConsumerPriority, + selectConsumerPeerId, +} from './consumers' +export type { + ConsumerRegistration, + ConsumerSelectionCandidate, + ConsumerStickyAssignment, +} from './consumers' +export { + resolveHealthCheckIntervalMs, + serverWsDefaultHeartbeatTtlMs, + serverWsHealthCheckIntervalDivisor, + serverWsMinimumHealthCheckIntervalMs, +} from './liveness' +export { createEventMetadata, createResponses } from './responses' +export { forEachEventMiddlewares, resolveEventDelivery } from './routing' diff --git a/packages/server-runtime/src/server-ws/airi/liveness.test.ts b/packages/server-runtime/src/server-ws/airi/liveness.test.ts new file mode 100644 index 000000000..ab6337d98 --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/liveness.test.ts @@ -0,0 +1,21 @@ +import { describe, expect, it } from 'vitest' + +import { + resolveHealthCheckIntervalMs, + serverWsDefaultHeartbeatTtlMs, + serverWsMinimumHealthCheckIntervalMs, +} from './liveness' + +describe('airi websocket liveness policy', () => { + it('uses the AIRI default heartbeat TTL', () => { + expect(serverWsDefaultHeartbeatTtlMs).toBe(60_000) + expect(resolveHealthCheckIntervalMs(serverWsDefaultHeartbeatTtlMs)).toBe(12_000) + }) + + it('keeps health checks at least five seconds apart', () => { + expect(serverWsMinimumHealthCheckIntervalMs).toBe(5_000) + expect(resolveHealthCheckIntervalMs(1_000)).toBe(serverWsMinimumHealthCheckIntervalMs) + expect(resolveHealthCheckIntervalMs(24_999)).toBe(serverWsMinimumHealthCheckIntervalMs) + expect(resolveHealthCheckIntervalMs(25_000)).toBe(serverWsMinimumHealthCheckIntervalMs) + }) +}) diff --git a/packages/server-runtime/src/server-ws/airi/liveness.ts b/packages/server-runtime/src/server-ws/airi/liveness.ts new file mode 100644 index 000000000..514472d27 --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/liveness.ts @@ -0,0 +1,13 @@ +/** Default heartbeat read timeout. */ +export const serverWsDefaultHeartbeatTtlMs = 60_000 + +/** Number of liveness checks scheduled within one heartbeat TTL. */ +export const serverWsHealthCheckIntervalDivisor = 5 + +/** Minimum interval to avoid busy liveness loops. */ +export const serverWsMinimumHealthCheckIntervalMs = 5_000 + +/** Resolves the AIRI heartbeat health-check interval in milliseconds. */ +export function resolveHealthCheckIntervalMs(heartbeatTtlMs: number) { + return Math.max(serverWsMinimumHealthCheckIntervalMs, Math.floor(heartbeatTtlMs / serverWsHealthCheckIntervalDivisor)) +} diff --git a/packages/server-runtime/src/server-ws/airi/responses.test.ts b/packages/server-runtime/src/server-ws/airi/responses.test.ts new file mode 100644 index 000000000..36820acbd --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/responses.test.ts @@ -0,0 +1,70 @@ +import { WebSocketEventSource } from '@proj-airi/server-shared/types' +import { describe, expect, it } from 'vitest' + +import packageJSON from '../../../package.json' + +import { createEventMetadata, createResponses } from './responses' + +describe('airi websocket responses', () => { + it('creates server metadata with source and parent event ids', () => { + const metadata = createEventMetadata('server-1', 'parent-event-1') + + expect(metadata.source).toEqual({ + kind: 'plugin', + plugin: { + id: WebSocketEventSource.Server, + version: packageJSON.version, + }, + id: 'server-1', + }) + expect(metadata.event).toEqual({ + id: expect.any(String), + parentId: 'parent-event-1', + }) + }) + + it('creates peer and extension authentication response shapes', () => { + const responses = createResponses('server-1') + + expect(responses.peerAuthenticated('peer-1', 'event-1')).toMatchObject({ + type: 'peer:authenticated', + data: { + authenticated: true, + peerId: 'peer-1', + }, + metadata: { + source: { + id: 'server-1', + plugin: { + id: WebSocketEventSource.Server, + version: packageJSON.version, + }, + }, + event: { + parentId: 'event-1', + }, + }, + }) + expect(responses.extensionAuthenticated({ id: 'airi-extension-chess' }, 'event-2')).toMatchObject({ + type: 'extension:authenticated', + data: { + authenticated: true, + identity: { + id: 'airi-extension-chess', + }, + }, + metadata: { + source: { + id: 'server-1', + plugin: { + id: WebSocketEventSource.Server, + version: packageJSON.version, + }, + }, + event: { + parentId: 'event-2', + }, + }, + }) + }) +}) diff --git a/packages/server-runtime/src/server-ws/airi/responses.ts b/packages/server-runtime/src/server-ws/airi/responses.ts new file mode 100644 index 000000000..2f092d6a2 --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/responses.ts @@ -0,0 +1,76 @@ +import type { ExtensionIdentity, MessageHeartbeat, MessageHeartbeatKind, MetadataEventSource, WebSocketEvent } from '@proj-airi/server-shared/types' + +import { ServerErrorMessages } from '@proj-airi/server-shared' +import { WebSocketEventSource } from '@proj-airi/server-shared/types' +import { nanoid } from 'nanoid' + +import packageJSON from '../../../package.json' + +/** Creates AIRI server event metadata and preserves optional parent correlation. */ +export function createEventMetadata( + serverInstanceId: string, + parentId?: string, +): { source: MetadataEventSource, event: { id: string, parentId?: string } } { + return { + event: { + id: nanoid(), + parentId, + }, + source: { + kind: 'plugin', + plugin: { + id: WebSocketEventSource.Server, + version: packageJSON.version, + }, + id: serverInstanceId, + }, + } +} + +/** Creates AIRI server response event factories. */ +export function createResponses(serverInstanceId: string) { + return { + authenticated(parentId?: string) { + return { + type: 'module:authenticated', + data: { authenticated: true }, + metadata: createEventMetadata(serverInstanceId, parentId), + } satisfies WebSocketEvent> + }, + peerAuthenticated(peerId: string, parentId?: string) { + return { + type: 'peer:authenticated', + data: { authenticated: true, peerId }, + metadata: createEventMetadata(serverInstanceId, parentId), + } satisfies WebSocketEvent> + }, + extensionAuthenticated(identity: ExtensionIdentity, parentId?: string) { + return { + type: 'extension:authenticated', + data: { identity, authenticated: true }, + metadata: createEventMetadata(serverInstanceId, parentId), + } satisfies WebSocketEvent> + }, + notAuthenticated(parentId?: string) { + return { + type: 'error', + data: { message: ServerErrorMessages.notAuthenticated }, + metadata: createEventMetadata(serverInstanceId, parentId), + } satisfies WebSocketEvent> + }, + error(message: string, parentId?: string) { + return { + type: 'error', + data: { message }, + metadata: createEventMetadata(serverInstanceId, parentId), + } satisfies WebSocketEvent> + }, + heartbeat(kind: MessageHeartbeatKind, message: MessageHeartbeat | string, parentId?: string) { + return { + type: 'transport:connection:heartbeat', + data: { kind, message, at: Date.now() }, + metadata: createEventMetadata(serverInstanceId, parentId), + } satisfies WebSocketEvent> + }, + } +} diff --git a/packages/server-runtime/src/server-ws/airi/routing.ts b/packages/server-runtime/src/server-ws/airi/routing.ts new file mode 100644 index 000000000..0ffcb82e5 --- /dev/null +++ b/packages/server-runtime/src/server-ws/airi/routing.ts @@ -0,0 +1,67 @@ +import type { DeliveryConfig, WebSocketEvent } from '@proj-airi/server-shared/types' + +import type { RouteContext, RouteDecision, RouteMiddleware } from '../../middlewares' + +import { getProtocolEventMetadata } from '@proj-airi/server-shared/types' + +/** + * Resolves the effective event delivery policy. + * + * Use when: + * - Protocol defaults should be merged with route-level delivery overrides + * - Routing needs to know whether the event should broadcast or target one consumer + * + * Expects: + * - Route delivery to override protocol metadata field-by-field + * + * Returns: + * - The merged broadcast/consumer delivery policy, or `undefined` when unrestricted + */ +export function resolveEventDelivery(event: WebSocketEvent): DeliveryConfig | undefined { + const eventMetadata = getProtocolEventMetadata(event.type) + const defaultDelivery = eventMetadata?.delivery + const routeDelivery = event.route?.delivery + + if (!defaultDelivery && !routeDelivery) { + return undefined + } + + return { + ...defaultDelivery, + ...routeDelivery, + } +} + +/** + * Iterates event middlewares in declaration order until one returns a decision. + * + * Use when: + * - The websocket runtime needs the first route decision from configured middleware + * + * Expects: + * - Middleware functions are ordered by caller policy + * + * Returns: + * - The first route decision, or `undefined` when no middleware decided + */ +export function forEachEventMiddlewares(input: { + event: WebSocketEvent + fromPeer: RouteContext['fromPeer'] + peers: Map + destinations?: RouteContext['destinations'] + middleware: RouteMiddleware[] +}): RouteDecision | undefined { + const context: RouteContext = { + event: input.event, + fromPeer: input.fromPeer, + peers: input.peers, + destinations: input.destinations, + } + + for (const middleware of input.middleware) { + const result = middleware(context) + if (result) { + return result + } + } +} diff --git a/packages/server-runtime/src/server-ws/core/index.ts b/packages/server-runtime/src/server-ws/core/index.ts deleted file mode 100644 index 1219cfc48..000000000 --- a/packages/server-runtime/src/server-ws/core/index.ts +++ /dev/null @@ -1,543 +0,0 @@ -/** - * Delivery settings used by the reusable websocket gateway. - * - * @param TMode - Delivery mode literals accepted by the adapter. - */ -export interface ServerWsDeliveryConfig { - /** - * Delivery mode selected by the protocol adapter. - * - * @default undefined - */ - mode?: TMode - /** - * Optional consumer group. - * - * @default "default" for consumer delivery modes. - */ - group?: string - /** - * Selection strategy within the target consumer set. - * - * @default "first" - */ - selection?: 'first' | 'priority' | 'sticky' | 'round-robin' - /** - * Sticky routing key used when `selection` is `sticky`. - * - * @default undefined - */ - stickyKey?: string - /** - * Whether missing consumers should be surfaced as an error by the adapter. - * - * @default false - */ - required?: boolean -} - -/** - * Delivery settings accepted by the reusable consumer registry. - * - * @param TMode - Consumer delivery mode literals accepted by the adapter. - */ -export type ServerWsConsumerDeliveryConfig = ServerWsDeliveryConfig - -/** - * Candidate peer metadata used for consumer selection. - */ -export interface ServerWsConsumerSelectionCandidate { - /** Peer id available to receive the event. */ - peerId: string - /** Higher values are selected before lower values. */ - priority: number - /** Timestamp captured when the peer registered as a consumer. */ - registeredAt: number - /** Whether the peer has completed protocol-level authentication. */ - authenticated: boolean - /** Explicit `false` excludes the peer from selection. */ - healthy?: boolean -} - -/** - * Stored consumer registration. - */ -export interface ServerWsConsumerRegistration { - /** Protocol event type consumed by the peer. */ - event: string - /** Normalized consumer group name. */ - group: string - /** Peer id that registered for the event/group pair. */ - peerId: string - /** Higher values are selected before lower values. */ - priority: number - /** Timestamp captured when the peer registered as a consumer. */ - registeredAt: number -} - -/** - * Describes protocol-agnostic text encoding and decoding for websocket events. - * - * @param TEvent - Event envelope shape owned by the protocol adapter. - */ -export interface ServerWsEventCodec { - /** Parses one text payload into a protocol event. */ - parse: (text: string) => TEvent - /** Serializes one protocol event or pre-serialized payload for peer sending. */ - stringify: (event: TEvent | string) => string - /** Detects raw transport control payloads that should not enter protocol routing. */ - detectControlFrame?: (text: string) => string | undefined -} - -/** - * Describes a websocket handler object accepted by H3 `defineWebSocketHandler`. - * - * @param TPeer - Runtime peer object accepted by lifecycle callbacks. - * @param TMessage - Runtime message object accepted by the message callback. - * @param TCloseDetails - Runtime close details object accepted by the close callback. - */ -export interface ServerWsGatewayHandler { - /** Called when a peer opens a websocket connection. */ - open?: (peer: TPeer) => void - /** Called when a peer sends one websocket message. */ - message?: (peer: TPeer, message: TMessage) => void - /** Called when the websocket runtime reports an error. */ - error?: (peer: TPeer, error: unknown) => void - /** Called when a peer closes a websocket connection. */ - close?: (peer: TPeer, details?: TCloseDetails) => void -} - -/** - * Minimal websocket peer shape used by the reusable gateway. - */ -export interface ServerWsPeer { - /** Stable peer id assigned by the websocket runtime. */ - get id(): string - /** Sends one payload to the peer. */ - send: (data: unknown, options?: { compress?: boolean }) => number | void | undefined - /** Closes the peer connection when the runtime exposes an explicit close hook. */ - close?: () => void - /** WebSocket ready state when exposed by the runtime. */ - readyState?: number - /** Request metadata associated with the websocket upgrade. */ - request?: { - /** Request URL associated with the websocket upgrade. */ - url?: string - /** Request headers associated with the websocket upgrade. */ - headers?: Headers - } - /** Remote peer address when exposed by the runtime. */ - remoteAddress?: string -} - -/** Default heartbeat read timeout used by the websocket gateway. */ -export const serverWsDefaultHeartbeatTtlMs = 60_000 - -/** Miss count where a peer becomes unhealthy but remains connected. */ -export const serverWsHealthCheckMissesUnhealthy = 5 - -/** Miss count where a peer is considered dead and should be closed. */ -export const serverWsHealthCheckMissesDead = serverWsHealthCheckMissesUnhealthy * 2 - -const DEFAULT_CONSUMER_GROUP = 'default' - -interface ServerWsConsumerRegistryRef { - event: string - group: string -} - -/** - * Sticky consumer assignment stored by the reusable consumer selector. - */ -export interface ServerWsStickyAssignment { - /** Protocol event type the sticky assignment belongs to. */ - event: string - /** Normalized consumer group the sticky assignment belongs to. */ - group: string - /** Peer selected for the sticky key. */ - peerId: string -} - -/** - * Creates a websocket event codec from explicit parser and serializer callbacks. - * - * Use when: - * - A protocol adapter wants to plug its own event envelope into `server-ws/core` - * - * Expects: - * - Parser and serializer preserve the adapter's current wire format - * - * Returns: - * - A protocol-agnostic codec object consumed by gateway code - */ -export function createEventCodec(codec: ServerWsEventCodec) { - return codec -} - -/** - * Wraps websocket lifecycle callbacks and disposal as a reusable mount object. - * - * Use when: - * - Adapters need one stable lifecycle shape for server mounting - * - * Expects: - * - `handler` contains already-bound protocol behavior - * - * Returns: - * - A handler plus idempotent disposal hook - */ -export function createGatewayLifecycle(input: { - handler: ServerWsGatewayHandler - dispose?: () => void -}) { - let disposed = false - - return { - handler: input.handler, - dispose: () => { - if (disposed) { - return - } - - disposed = true - input.dispose?.() - }, - } -} - -/** - * Resolves the interval used for heartbeat health checks. - * - * Use when: - * - Gateway code needs to convert heartbeat TTL into periodic miss checks - * - * Expects: - * - Very small TTL values should still avoid busy intervals - * - * Returns: - * - Interval in milliseconds - */ -export function resolveServerWsHealthCheckIntervalMs(heartbeatTtlMs: number) { - return Math.max(5_000, Math.floor(heartbeatTtlMs / serverWsHealthCheckMissesUnhealthy)) -} - -/** - * Creates a typed peer store around websocket peer state. - * - * Use when: - * - A gateway needs stable peer lookup, iteration, and cleanup - * - * Expects: - * - `TState` contains protocol-specific peer state - * - * Returns: - * - A small registry over peers keyed by peer id - */ -export function createServerWsPeerStore() { - const peers = new Map() - - return { - peers, - get(peerId: string) { - return peers.get(peerId) - }, - set(peerId: string, state: TState) { - peers.set(peerId, state) - return state - }, - delete(peerId: string) { - return peers.delete(peerId) - }, - clear() { - peers.clear() - }, - values() { - return peers.values() - }, - entries() { - return peers.entries() - }, - size() { - return peers.size - }, - } -} - -/** - * Checks whether a delivery mode targets the consumer registry. - * - * Use when: - * - A protocol adapter receives broad delivery modes but must call consumer-only APIs - * - * Expects: - * - Non-consumer modes such as `broadcast` should remain outside the consumer registry - * - * Returns: - * - `true` for `consumer` and `consumer-group` - */ -export function isConsumerDeliveryMode(mode: unknown): mode is ServerWsConsumerDeliveryConfig['mode'] { - return mode === 'consumer' || mode === 'consumer-group' -} - -/** - * Normalizes delivery mode for consumer registration. - * - * Before: - * - undefined with group "workers" - * - * After: - * - "consumer-group" - */ -export function normalizeConsumerMode(mode: unknown, group?: string): 'consumer' | 'consumer-group' { - if (isConsumerDeliveryMode(mode)) { - return mode! - } - - return group ? 'consumer-group' : 'consumer' -} - -/** - * Normalizes consumer priority. - * - * Before: - * - NaN - * - * After: - * - 0 - */ -export function normalizeConsumerPriority(priority: unknown) { - return typeof priority === 'number' && Number.isFinite(priority) - ? priority - : 0 -} - -function normalizeConsumerGroup(mode: ServerWsConsumerDeliveryConfig['mode'], group?: string) { - if (mode === 'consumer') { - return DEFAULT_CONSUMER_GROUP - } - - return group || DEFAULT_CONSUMER_GROUP -} - -function getConsumerRegistryKey(event: string, group: string) { - return JSON.stringify([event, group]) -} - -function getStickyRegistryKey(event: string, group: string, stickyKey: string) { - return JSON.stringify([event, group, stickyKey]) -} - -function sortConsumers(entries: Array>) { - return [...entries].sort((left, right) => { - if (right.priority !== left.priority) { - return right.priority - left.priority - } - - return left.registeredAt - right.registeredAt - }) -} - -/** - * Selects a concrete consumer peer for consumer-style delivery modes. - * - * Use when: - * - An event should be sent to exactly one registered consumer - * - Sticky or round-robin routing needs to be resolved against live peer metadata - * - * Expects: - * - Candidates already describe authenticated and health state - * - * Returns: - * - The selected peer id, or `undefined` when no eligible consumer is available - */ -export function selectConsumerPeerId(options: { - eventType: string - fromPeerId: string - delivery?: ServerWsDeliveryConfig - candidates: ServerWsConsumerSelectionCandidate[] - roundRobinCursor?: Map - stickyAssignments?: Map -}) { - const { candidates, delivery, eventType, fromPeerId } = options - if (!delivery || (delivery.mode !== 'consumer' && delivery.mode !== 'consumer-group')) { - return - } - - const normalizedGroup = normalizeConsumerGroup(delivery.mode, delivery.group) - const registryKey = getConsumerRegistryKey(eventType, normalizedGroup) - const availableEntries = sortConsumers( - candidates - .filter(entry => entry.peerId !== fromPeerId) - .filter(entry => entry.authenticated && entry.healthy !== false), - ) - - if (availableEntries.length === 0) { - return - } - - const selection = delivery.selection ?? 'first' - if (selection === 'sticky' && delivery.stickyKey) { - const stickyRegistryKey = getStickyRegistryKey(eventType, normalizedGroup, delivery.stickyKey) - const stickyAssignment = options.stickyAssignments?.get(stickyRegistryKey) - if (stickyAssignment && stickyAssignment.peerId !== fromPeerId) { - const stickyCandidate = availableEntries.find(entry => entry.peerId === stickyAssignment.peerId) - if (stickyCandidate) { - return stickyAssignment.peerId - } - } - - const selected = availableEntries[0] - options.stickyAssignments?.set(stickyRegistryKey, { event: eventType, group: normalizedGroup, peerId: selected.peerId }) - return selected.peerId - } - - if (selection === 'round-robin') { - const cursor = options.roundRobinCursor?.get(registryKey) ?? 0 - const selected = availableEntries[cursor % availableEntries.length] - options.roundRobinCursor?.set(registryKey, (cursor + 1) % availableEntries.length) - return selected.peerId - } - - return availableEntries[0].peerId -} - -/** - * Creates a reusable consumer delivery orchestrator for websocket peers. - * - * Use when: - * - A protocol adapter supports one-consumer delivery or consumer groups - * - * Expects: - * - Peer liveness is checked by the caller before delivery - * - * Returns: - * - Registration, unregister, listing, selection, and cleanup helpers - */ -export function createConsumerOrchestrator() { - const consumerRegistry = new Map>>() - const consumerKeysByPeer = new Map>() - const deliveryRoundRobinCursor = new Map() - const stickyAssignments = new Map() - - function removeStickyAssignmentsFor(event: string, group: string, peerId?: string) { - for (const [stickyKey, assignment] of stickyAssignments.entries()) { - if (peerId && assignment.peerId !== peerId) { - continue - } - - if (assignment.event === event && assignment.group === group) { - stickyAssignments.delete(stickyKey) - } - } - } - - return { - register(input: { peerId: string, event: string, mode: ServerWsConsumerDeliveryConfig['mode'], group?: string, priority?: number }) { - const normalizedGroup = normalizeConsumerGroup(input.mode, input.group) - const registryKey = getConsumerRegistryKey(input.event, normalizedGroup) - let groups = consumerRegistry.get(input.event) - if (!groups) { - groups = new Map() - consumerRegistry.set(input.event, groups) - } - - let peersForGroup = groups.get(normalizedGroup) - if (!peersForGroup) { - peersForGroup = new Map() - groups.set(normalizedGroup, peersForGroup) - } - - const didGrowMembership = !peersForGroup.has(input.peerId) - peersForGroup.set(input.peerId, { - event: input.event, - group: normalizedGroup, - peerId: input.peerId, - priority: normalizeConsumerPriority(input.priority), - registeredAt: Date.now(), - }) - if (didGrowMembership) { - deliveryRoundRobinCursor.delete(registryKey) - } - - let registrations = consumerKeysByPeer.get(input.peerId) - if (!registrations) { - registrations = new Map() - consumerKeysByPeer.set(input.peerId, registrations) - } - registrations.set(registryKey, { event: input.event, group: normalizedGroup }) - }, - unregister(input: { peerId: string, event: string, mode: ServerWsConsumerDeliveryConfig['mode'], group?: string }) { - const normalizedGroup = normalizeConsumerGroup(input.mode, input.group) - const registryKey = getConsumerRegistryKey(input.event, normalizedGroup) - const groups = consumerRegistry.get(input.event) - const peersForGroup = groups?.get(normalizedGroup) - const didDelete = peersForGroup?.delete(input.peerId) ?? false - - if (!didDelete) { - return - } - - deliveryRoundRobinCursor.delete(registryKey) - if (peersForGroup?.size === 0) { - groups?.delete(normalizedGroup) - } - if (groups?.size === 0) { - consumerRegistry.delete(input.event) - } - - const registrations = consumerKeysByPeer.get(input.peerId) - registrations?.delete(registryKey) - if (registrations?.size === 0) { - consumerKeysByPeer.delete(input.peerId) - } - - removeStickyAssignmentsFor(input.event, normalizedGroup, input.peerId) - }, - unregisterPeer(peerId: string) { - const registrations = consumerKeysByPeer.get(peerId) - if (!registrations?.size) { - return - } - - for (const registration of registrations.values()) { - const { event, group } = registration - const groups = consumerRegistry.get(event) - const peersForGroup = groups?.get(group) - peersForGroup?.delete(peerId) - deliveryRoundRobinCursor.delete(getConsumerRegistryKey(event, group)) - if (peersForGroup?.size === 0) { - groups?.delete(group) - } - if (groups?.size === 0) { - consumerRegistry.delete(event) - } - - removeStickyAssignmentsFor(event, group, peerId) - } - - consumerKeysByPeer.delete(peerId) - }, - listFor(input: { event: string, mode: ServerWsConsumerDeliveryConfig['mode'], group?: string }) { - const normalizedGroup = normalizeConsumerGroup(input.mode, input.group) - return [...consumerRegistry.get(input.event)?.get(normalizedGroup)?.values() ?? []] - }, - select(input: { - eventType: string - fromPeerId: string - delivery?: ServerWsDeliveryConfig - candidates: ServerWsConsumerSelectionCandidate[] - }) { - return selectConsumerPeerId({ - ...input, - roundRobinCursor: deliveryRoundRobinCursor, - stickyAssignments, - }) - }, - clear() { - consumerRegistry.clear() - consumerKeysByPeer.clear() - deliveryRoundRobinCursor.clear() - stickyAssignments.clear() - }, - } -} diff --git a/packages/server-runtime/src/server.test.ts b/packages/server-runtime/src/server.test.ts index 8d912ff6a..367866fe4 100644 --- a/packages/server-runtime/src/server.test.ts +++ b/packages/server-runtime/src/server.test.ts @@ -12,6 +12,7 @@ const serveMocks = vi.hoisted(() => { const closeCall = vi.fn(async () => {}) const disposeCall = vi.fn(() => {}) + const createH3CrossWsPluginCall = vi.fn(() => ({ name: 'better-ws-h3-plugin' })) const setupAppCall = vi.fn(() => ({ app: { fetch: vi.fn(async () => ({ crossws: {} })), @@ -22,6 +23,7 @@ const serveMocks = vi.hoisted(() => { return { closeCall, + createH3CrossWsPluginCall, disposeCall, rejectServe: (error: Error) => rejectServe?.(error), resolveServe: () => resolveServe?.(), @@ -34,15 +36,14 @@ vi.mock('h3', () => ({ H3: class { get = vi.fn() }, - defineWebSocketHandler: vi.fn(handler => handler), serve: vi.fn(() => ({ serve: serveMocks.serveCall, close: serveMocks.closeCall, })), })) -vi.mock('crossws/server', () => ({ - plugin: vi.fn(() => ({})), +vi.mock('@proj-airi/better-ws/server/h3', () => ({ + createH3CrossWsPlugin: serveMocks.createH3CrossWsPluginCall, })) vi.mock('./index', () => ({ @@ -72,6 +73,9 @@ describe('createServer', async () => { await Promise.all([firstStart, secondStart]) expect(serveMocks.serveCall).toHaveBeenCalledTimes(1) + expect(serveMocks.createH3CrossWsPluginCall).toHaveBeenCalledWith(expect.objectContaining({ + fetch: expect.any(Function), + })) }) it('clears the single-flight state when start fails', async () => { diff --git a/packages/server-runtime/src/server/index.ts b/packages/server-runtime/src/server/index.ts index 2c1d8994e..a68b21133 100644 --- a/packages/server-runtime/src/server/index.ts +++ b/packages/server-runtime/src/server/index.ts @@ -1,3 +1,5 @@ +import type { H3CrossWsApp, H3CrossWsResponse } from '@proj-airi/better-ws/server/h3' + import type { AppOptions } from '..' import { isIP } from 'node:net' @@ -5,7 +7,7 @@ import { networkInterfaces } from 'node:os' import { useLogg } from '@guiiai/logg' import { merge } from '@moeru/std' -import { plugin as ws } from 'crossws/server' +import { createH3CrossWsPlugin } from '@proj-airi/better-ws/server/h3' import { serve } from 'h3' import { normalizeLoggerConfig, setupApp } from '..' @@ -151,13 +153,15 @@ export function createServer(opts?: ServerOptions): Server { startTask = (async () => { const secureEnabled = options?.tlsConfig != null const h3App = setupApp(options) + const crossWsApp = { + fetch: async request => await h3App.app.fetch(request) as H3CrossWsResponse, + } satisfies H3CrossWsApp const port = options.port const hostname = options.hostname const instance = serve(h3App.app, { - // @ts-expect-error - the .crossws property wasn't extended in types - plugins: [ws({ resolve: async req => (await h3App.app.fetch(req)).crossws })], + plugins: [createH3CrossWsPlugin(crossWsApp)], port, hostname, tls: options?.tlsConfig || undefined, diff --git a/packages/server-runtime/src/setupApp.liveness.test.ts b/packages/server-runtime/src/setupApp.liveness.test.ts new file mode 100644 index 000000000..2f7ed19aa --- /dev/null +++ b/packages/server-runtime/src/setupApp.liveness.test.ts @@ -0,0 +1,235 @@ +import type { WebSocketEvent } from '@proj-airi/server-shared/types' + +import type { Peer } from './types' + +import { parse, stringify } from 'superjson' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import { setupApp } from './index' + +interface TestWebSocketHandler { + open?: (peer: Peer) => void + message?: (peer: Peer, message: { text: () => string }) => void + close?: (peer: Peer, details?: { code?: number, reason?: string, wasClean?: unknown }) => void +} + +interface TestWsServer { + accept: ( + adapter: { id: string, send: (message: { text: () => string }) => void | number, close?: () => void }, + options: { state: { rawPeer: Peer } }, + ) => void + peers: { + get: (peerId: string) => { receive: (message: { text: () => string }) => void } | undefined + } + remove: (peerId: string, details?: { code?: number, reason?: string, wasClean?: unknown }) => void +} + +const h3Mocks = vi.hoisted(() => ({ + handlers: new Map(), +})) + +vi.mock('h3', () => ({ + H3: class { + get(path: string, handler: unknown) { + h3Mocks.handlers.set(path, handler) + } + }, +})) + +vi.mock('@proj-airi/better-ws/server/h3', () => ({ + toH3Handler: vi.fn((server: TestWsServer, options: { state: (peer: Peer) => { rawPeer: Peer } }) => ({ + open(peer: Peer) { + server.accept({ + id: peer.id, + send: message => peer.send(message.text()), + close: () => peer.close?.(), + }, { + state: options.state(peer), + }) + }, + message(peer: Peer, message: { text: () => string }) { + server.peers.get(peer.id)?.receive(message) + }, + close(peer: Peer, details?: { code?: number, reason?: string, wasClean?: unknown }) { + server.remove(peer.id, details) + }, + })), +})) + +function createPeer(id: string) { + const sent: string[] = [] + const send: Peer['send'] = (data) => { + sent.push(String(data)) + } + + return { + peer: { + id, + send: vi.fn(send), + close: vi.fn(), + request: { url: `/ws?id=${id}` }, + remoteAddress: '127.0.0.1', + } satisfies Peer, + sent, + } +} + +function wsHandler() { + const handler = h3Mocks.handlers.get('/ws') as TestWebSocketHandler | undefined + if (!handler) { + throw new Error('Expected setupApp to register a /ws websocket handler.') + } + + return handler +} + +function sendEvent( + handler: TestWebSocketHandler, + peer: Peer, + event: WebSocketEvent, +) { + handler.message?.(peer, { text: () => stringify(event) }) +} + +function decodeEvents(sent: string[]) { + return sent.map(message => parse(message)) +} + +function createExtensionModuleAnnounceEvent(): WebSocketEvent { + return { + type: 'extension:module:announce', + data: { + name: 'memory', + possibleEvents: [], + identity: { + id: 'memory-module-1', + extension: { + id: 'extension-1', + }, + }, + }, + metadata: { + source: { + kind: 'plugin', + id: 'extension-1', + plugin: { + id: 'extension-1', + }, + }, + event: { + id: 'announce-1', + }, + }, + } +} + +describe('setupApp websocket liveness', () => { + beforeEach(() => { + h3Mocks.handlers.clear() + vi.useFakeTimers({ now: 0 }) + }) + + afterEach(() => { + vi.useRealTimers() + }) + + it('broadcasts extension module unhealthy events from better-ws liveness checks', () => { + const runtime = setupApp({ heartbeat: { readTimeout: 20_000 } }) + const handler = wsHandler() + const observer = createPeer('observer') + const modulePeer = createPeer('module-peer') + + handler.open?.(observer.peer) + handler.open?.(modulePeer.peer) + sendEvent(handler, modulePeer.peer, createExtensionModuleAnnounceEvent()) + observer.sent.length = 0 + + vi.advanceTimersByTime(25_000) + + expect(decodeEvents(observer.sent)).toEqual(expect.arrayContaining([ + expect.objectContaining({ + type: 'registry:modules:health:unhealthy', + data: { + name: 'memory', + identity: { + id: 'memory-module-1', + extension: { + id: 'extension-1', + }, + }, + reason: 'heartbeat late', + }, + }), + ])) + + runtime.dispose() + }) + + it('de-announces expired extension modules when better-ws removes stale peers', () => { + const runtime = setupApp({ heartbeat: { readTimeout: 20_000 } }) + const handler = wsHandler() + const observer = createPeer('observer') + const modulePeer = createPeer('module-peer') + + handler.open?.(observer.peer) + handler.open?.(modulePeer.peer) + sendEvent(handler, modulePeer.peer, createExtensionModuleAnnounceEvent()) + observer.sent.length = 0 + + vi.advanceTimersByTime(25_000) + handler.message?.(observer.peer, { text: () => 'pong' }) + observer.sent.length = 0 + vi.advanceTimersByTime(25_000) + + expect(modulePeer.peer.close).toHaveBeenCalledOnce() + expect(decodeEvents(observer.sent)).toEqual(expect.arrayContaining([ + expect.objectContaining({ + type: 'extension:module:de-announced', + data: expect.objectContaining({ + name: 'memory', + reason: 'heartbeat expired', + }), + }), + ])) + + runtime.dispose() + }) + + it('de-announces extension modules before accepting a same-id reconnect', () => { + const runtime = setupApp({ heartbeat: { readTimeout: 20_000 } }) + const handler = wsHandler() + const observer = createPeer('observer') + const firstModulePeer = createPeer('module-peer') + const secondModulePeer = createPeer('module-peer') + + handler.open?.(observer.peer) + handler.open?.(firstModulePeer.peer) + sendEvent(handler, firstModulePeer.peer, createExtensionModuleAnnounceEvent()) + observer.sent.length = 0 + + handler.open?.(secondModulePeer.peer) + + expect(decodeEvents(observer.sent)).toEqual(expect.arrayContaining([ + expect.objectContaining({ + type: 'extension:module:de-announced', + data: expect.objectContaining({ + name: 'memory', + reason: 'connection closed', + }), + }), + ])) + + runtime.dispose() + }) + + it('closes each raw peer once during runtime disposal', () => { + const runtime = setupApp({ heartbeat: { readTimeout: 20_000 } }) + const handler = wsHandler() + const peer = createPeer('peer-1') + + handler.open?.(peer.peer) + runtime.dispose() + + expect(peer.peer.close).toHaveBeenCalledOnce() + }) +}) diff --git a/packages/server-runtime/src/types/conn.ts b/packages/server-runtime/src/types/conn.ts index 0720f7ce4..fde6ac1e8 100644 --- a/packages/server-runtime/src/types/conn.ts +++ b/packages/server-runtime/src/types/conn.ts @@ -52,5 +52,10 @@ export interface AuthenticatedPeer extends NamedPeer { extensionModules?: Map lastHeartbeatAt?: number healthy?: boolean + /** + * REVIEW: Legacy field name kept during the better-ws migration. + * The value now stores peer silence duration in milliseconds, not a miss count. + * Rename this with the server-runtime peer state cleanup. + */ missedHeartbeats?: number } diff --git a/packages/server-sdk/package.json b/packages/server-sdk/package.json index 3bd111bd0..5ce8bb6f6 100644 --- a/packages/server-sdk/package.json +++ b/packages/server-sdk/package.json @@ -38,8 +38,9 @@ }, "dependencies": { "@moeru/std": "catalog:", + "@proj-airi/better-ws": "workspace:^", "@proj-airi/server-shared": "workspace:^", - "crossws": "catalog:", - "superjson": "catalog:" + "superjson": "catalog:", + "valibot": "catalog:" } } diff --git a/packages/server-sdk/src/client.ts b/packages/server-sdk/src/client.ts index 0b5f2a1c9..f50aa9209 100644 --- a/packages/server-sdk/src/client.ts +++ b/packages/server-sdk/src/client.ts @@ -1,3 +1,9 @@ +import type { + Client as BetterWsClient, + ClientConnector, + PrepareContext, + ReconnectOptions, +} from '@proj-airi/better-ws' import type { ExtensionIdentity, ExtensionModuleIdentity, @@ -9,15 +15,16 @@ import type { WebSocketEvents, } from '@proj-airi/server-shared/types' -import type { WebSocketLike, WebSocketLikeConstructor, WebSocketMessageEventLike } from './websocket-like' - -import NativeWebSocket from 'crossws/websocket' -import superjson from 'superjson' - -import { errorMessageFrom, sleep } from '@moeru/std' +import { errorMessageFrom } from '@moeru/std' +import { createClient as createBetterWsClient } from '@proj-airi/better-ws' +import { createCrossWsConnector } from '@proj-airi/better-ws/client/crossws' import { isTerminalAuthenticationServerErrorMessage, parseServerErrorMessage } from '@proj-airi/server-shared' import { MessageHeartbeat, MessageHeartbeatKind } from '@proj-airi/server-shared/types' +import { parseEvent, stringifyEvent } from './codec' + +export type { ClientConnector, ClientEvents } from '@proj-airi/better-ws' + export type ClientStatus = | 'idle' | 'connecting' @@ -49,7 +56,7 @@ export interface ClientOptions { url?: string name: string token?: string - websocketConstructor?: WebSocketLikeConstructor + connector?: ClientConnector> /** * Selects the connection handshake owned by this client. * @@ -63,7 +70,7 @@ export interface ClientOptions { identity?: ExtensionModuleIdentity dependencies?: ModuleDependency[] configSchema?: ModuleConfigSchema - heartbeat?: ClientHeartbeatOptions + heartbeat?: false | ClientHeartbeatOptions autoConnect?: boolean autoReconnect?: boolean @@ -78,13 +85,33 @@ export interface ClientOptions { onAnySend?: (data: WebSocketEvent) => void } -interface ConnectionAttempt { - announced: boolean - authenticated: boolean - promise: Promise - reject: (error: Error) => void - resolve: () => void - socket: WebSocketLike +interface NormalizedClientOptions { + url: string + name: string + token?: string + connector: ClientConnector> + handshake: 'module' | 'manual' + connectTimeoutMs: number + possibleEvents: Array> + extension: ExtensionIdentity + identity: ExtensionModuleIdentity + dependencies: ModuleDependency[] + configSchema?: ModuleConfigSchema + heartbeat: false | Required + autoConnect: boolean + autoReconnect: boolean + maxReconnectAttempts: number + onError: (error: unknown) => void + onClose: () => void + onReady: () => void + onStateChange: (context: ClientStateChangeContext) => void + onAnyMessage: (data: WebSocketEvent) => void + onAnySend: (data: WebSocketEvent) => void +} + +interface ProtocolWaitResult { + ready: boolean + error?: Error } function createInstanceId() { @@ -95,19 +122,11 @@ function createEventId() { return `${Date.now().toString(36)}-${Math.random().toString(36).slice(2, 10)}` } -function createDeferredPromise() { - let resolve!: () => void - let reject!: (error: Error) => void +function normalizeHeartbeatOptions(heartbeat?: false | ClientHeartbeatOptions): false | Required { + if (heartbeat === false) { + return false + } - const promise = new Promise((innerResolve, innerReject) => { - resolve = innerResolve - reject = innerReject - }) - - return { promise, reject, resolve } -} - -function normalizeHeartbeatOptions(heartbeat?: ClientHeartbeatOptions): Required { const readTimeout = heartbeat?.readTimeout ?? 30_000 const pingInterval = heartbeat?.pingInterval ?? Math.max(1_000, Math.floor(readTimeout / 2)) @@ -118,73 +137,115 @@ function normalizeHeartbeatOptions(heartbeat?: ClientHeartbeatOptions): Required } } +/** Wraps a text websocket connector with AIRI protocol serialization. */ +export function createTextProtocolConnector( + textConnector: ClientConnector, +): ClientConnector> { + return { + async connect(events) { + const connection = await textConnector.connect({ + message(text) { + try { + events.message(parseEvent(text)) + } + catch (error) { + events.error(error) + } + }, + close: details => events.close(details), + error: error => events.error(error), + }) + + return { + send: message => connection.send(stringifyEvent(message)), + close: (code, reason) => connection.close?.(code, reason), + ping: connection.ping, + pong: connection.pong, + } + }, + } +} + +function createDefaultProtocolConnector(url: string): ClientConnector> { + return createTextProtocolConnector(createCrossWsConnector({ url })) +} + +function normalizeOptions(options: ClientOptions): NormalizedClientOptions { + const url = options.url ?? 'ws://localhost:6121/ws' + const extension = options.extension ?? { id: options.name } + const identity = options.identity ?? { + id: createInstanceId(), + extension, + } + + return { + url, + name: options.name, + token: options.token, + connector: options.connector ?? createDefaultProtocolConnector(url), + handshake: options.handshake ?? 'module', + connectTimeoutMs: options.connectTimeoutMs ?? 15_000, + possibleEvents: options.possibleEvents ?? [], + extension, + identity, + dependencies: options.dependencies ?? [], + configSchema: options.configSchema, + heartbeat: normalizeHeartbeatOptions(options.heartbeat), + autoConnect: options.autoConnect ?? true, + autoReconnect: options.autoReconnect ?? true, + maxReconnectAttempts: options.maxReconnectAttempts ?? -1, + onError: options.onError ?? (() => {}), + onClose: options.onClose ?? (() => {}), + onReady: options.onReady ?? (() => {}), + onStateChange: options.onStateChange ?? (() => {}), + onAnyMessage: options.onAnyMessage ?? (() => {}), + onAnySend: options.onAnySend ?? (() => {}), + } +} + +function createConnectionTimeoutError(timeout: number) { + return new Error(`Connection timed out after ${timeout}ms`) +} + +function createAbortError() { + return new Error('Connection aborted') +} + export class Client { - private websocket?: WebSocketLike - private shouldClose = false - private connectTask?: Promise - private heartbeatTimer?: ReturnType - private lastPingAt = 0 - private lastReadAt = 0 - private reconnectAttempts = 0 - private pendingReconnect = false - private connectionAttempt?: ConnectionAttempt - private failureReason?: Error - private status: ClientStatus = 'idle' - private readonly identity: ExtensionModuleIdentity - private readonly heartbeat: Required - private readonly websocketConstructor: WebSocketLikeConstructor - - private readonly opts: - & Required, 'token' | 'heartbeat' | 'websocketConstructor' | 'configSchema'>> - & Pick, 'token' | 'heartbeat' | 'configSchema'> - + private readonly opts: NormalizedClientOptions + private readonly transport: BetterWsClient> private readonly eventListeners = new Map< keyof WebSocketEvents, - Set<(data: WebSocketBaseEvent) => void | Promise> + Set<(data: WebSocketBaseEvent) => void | Promise> >() private readonly stateListeners = new Set<(context: ClientStateChangeContext) => void>() + private status: ClientStatus = 'idle' + private connectTask?: Promise + private failureReason?: Error constructor(options: ClientOptions) { - const { websocketConstructor, ...clientOptions } = options - const extension = options.extension ?? { - id: options.name, - } - const identity = options.identity ?? { - id: createInstanceId(), - extension, - } + this.opts = normalizeOptions(options) + this.transport = createBetterWsClient>({ + connector: this.opts.connector, + reconnect: this.createReconnectOptions(), + heartbeat: this.createHeartbeatOptions(), + prepare: context => this.prepareProtocolConnection(context), + }) - const heartbeat = normalizeHeartbeatOptions(options.heartbeat) - - this.opts = { - url: 'ws://localhost:6121/ws', - connectTimeoutMs: 15_000, - onAnyMessage: () => {}, - onAnySend: () => {}, - possibleEvents: [], - dependencies: [], - configSchema: undefined, - onError: () => {}, - onClose: () => {}, - onReady: () => {}, - onStateChange: () => {}, - autoConnect: true, - autoReconnect: true, - maxReconnectAttempts: -1, - handshake: 'module', - ...clientOptions, - extension, - heartbeat, - identity, - } - - this.identity = identity - this.heartbeat = heartbeat - this.websocketConstructor = websocketConstructor ?? (NativeWebSocket as unknown as WebSocketLikeConstructor) + this.transport.onMessage(({ message }) => { + void this.handleMessage(message) + }) + this.transport.onStateChange(({ previousState, state }) => { + this.handleTransportStateChange(previousState, state) + }) if (this.opts.autoConnect) { - void this.connect() + void this.connect().catch((error) => { + const normalized = this.normalizeError(error, 'Failed to connect websocket client') + this.failureReason = normalized + this.opts.onError(normalized) + }) } } @@ -197,7 +258,7 @@ export class Client { } get isSocketOpen() { - return this.websocket?.readyState === this.websocketConstructor.OPEN + return this.transport.state === 'open' || this.transport.state === 'preparing' || this.transport.state === 'ready' } get lastError() { @@ -205,21 +266,23 @@ export class Client { } async connect(options?: ConnectOptions) { - if (this.shouldClose) { - throw new Error('Client is closed') - } - if (this.status === 'ready') { return } - if (this.connectTask) { - return this.waitForConnection(this.connectTask, options) + if (!this.connectTask && this.transport.state === 'reconnecting') { + return this.waitForConnection(this.waitForReady(), options) } - this.connectTask = this.runConnectLoop().finally(() => { - this.connectTask = undefined - }) + if (!this.connectTask && (this.transport.state === 'open' || this.transport.state === 'preparing')) { + return this.waitForConnection(this.waitForReady(), options) + } + + if (!this.connectTask) { + this.connectTask = this.transport.connect().finally(() => { + this.connectTask = undefined + }) + } return this.waitForConnection(this.connectTask, options) } @@ -250,7 +313,7 @@ export class Client { this.eventListeners.set(event, listeners) } - listeners.add(callback as any) + listeners.add(callback as (data: WebSocketBaseEvent) => void | Promise) return () => { this.offEvent(event, callback) @@ -267,25 +330,24 @@ export class Client { } if (callback) { - listeners.delete(callback as any) + listeners.delete(callback as (data: WebSocketBaseEvent) => void | Promise) if (!listeners.size) { this.eventListeners.delete(event) } + return } - else { - this.eventListeners.delete(event) - } + + this.eventListeners.delete(event) } send(data: WebSocketEventOptionalSource): boolean { - if (!this.isSocketOpen || !this.websocket) { + const payload = this.createPayload(data) + const result = this.transport.send(payload) + if (!result.ok) { return false } - const payload = this.createPayload(data) - this.opts.onAnySend?.(payload) - this.websocket.send(superjson.stringify(payload)) - + this.opts.onAnySend(payload) return true } @@ -295,260 +357,220 @@ export class Client { } } - sendRaw(data: string | ArrayBufferLike | ArrayBufferView): boolean { - if (!this.isSocketOpen || !this.websocket) { + close(code?: number, reason?: string): void { + this.transport.close(code, reason) + } + + private createReconnectOptions(): false | ReconnectOptions { + if (!this.opts.autoReconnect) { return false } - this.websocket.send(data) - return true - } - - close(): void { - this.shouldClose = true - this.pendingReconnect = false - this.transitionTo('closing') - this.stopHeartbeat() - this.rejectAttempt(new Error('Client closed')) - - const websocket = this.websocket - this.websocket = undefined - - if (websocket && websocket.readyState !== this.websocketConstructor.CLOSED && websocket.readyState !== this.websocketConstructor.CLOSING) { - websocket.close() - } - - this.transitionTo('closed') - } - - private async runConnectLoop() { - const reconnectingFromReady = this.pendingReconnect - this.pendingReconnect = false - - while (!this.shouldClose) { - const reconnecting = this.reconnectAttempts > 0 - this.transitionTo(reconnecting ? 'reconnecting' : 'connecting') - - try { - await this.connectOnce({ reconnectingFromReady }) - this.reconnectAttempts = 0 - return - } - catch (error) { - const normalizedError = error instanceof Error ? error : new Error(errorMessageFrom(error) ?? 'Failed to connect websocket client') - this.failureReason = normalizedError - this.opts.onError?.(normalizedError) - - if (this.shouldClose) { - throw normalizedError + return { + retries: (attempt, error) => { + const normalized = this.normalizeError(error, 'Failed to connect websocket client') + if (isTerminalAuthenticationServerErrorMessage(normalized.message)) { + return false } - if (isTerminalAuthenticationServerErrorMessage(normalizedError.message)) { - this.transitionTo('failed') - throw normalizedError - } - - if (!this.opts.autoReconnect && reconnecting) { - this.transitionTo('failed') - throw normalizedError - } - - if (!this.canRetry()) { - this.transitionTo('failed') - throw normalizedError - } - - const delay = this.getReconnectDelay(this.reconnectAttempts) - this.reconnectAttempts += 1 - await sleep(delay) - } - } - - throw new Error('Client is closed') - } - - private connectOnce(options: { reconnectingFromReady?: boolean } = {}): Promise { - const WebSocketConstructor = this.websocketConstructor - const ws = new WebSocketConstructor(this.opts.url) - this.websocket = ws - this.lastReadAt = Date.now() - this.lastPingAt = 0 - - const deferred = createDeferredPromise() - const attempt: ConnectionAttempt = { - announced: false, - authenticated: !this.opts.token, - promise: deferred.promise, - reject: deferred.reject, - resolve: deferred.resolve, - socket: ws, - } - - this.connectionAttempt = attempt - - const isCurrentSocket = () => this.websocket === ws - const connectTimeoutMs = this.opts.connectTimeoutMs - const connectTimer = setTimeout(() => { - if (ws.readyState === WebSocket.OPEN) { - return - } - - ws.close() - deferred.reject(new Error(`Connection timeout after ${connectTimeoutMs}ms`)) - }, connectTimeoutMs) - - const clearConnectTimer = () => { - clearTimeout(connectTimer) - } - - ws.onmessage = (event: WebSocketMessageEventLike) => { - if (!isCurrentSocket()) { - return - } - - void this.handleMessage(event) - } - - ws.onerror = (event: unknown) => { - clearConnectTimer() - - if (!isCurrentSocket()) { - return - } - - // Extract error from WebSocket error event which may vary in shape - const error = (event as any)?.error instanceof Error - ? (event as any).error - : new Error('WebSocket error') - if (this.connectionAttempt) { - this.handleSocketFailure(error, ws) - } - else { - this.opts.onError?.(error) - void this.reconnectAfterProtocolError(error) - } - } - - ws.onclose = () => { - clearConnectTimer() - - if (!isCurrentSocket()) { - return - } - - const wasReady = this.status === 'ready' - this.cleanupSocket(ws) - this.opts.onClose?.() - - if (this.shouldClose) { - return - } - - if (wasReady && this.opts.autoReconnect) { - this.pendingReconnect = true - this.transitionTo('idle') - void this.connect() - return - } - - this.rejectAttempt(new Error('WebSocket closed')) - } - - ws.onopen = () => { - clearConnectTimer() - - if (!isCurrentSocket()) { - return - } - - this.startHeartbeat() - - if (this.opts.handshake === 'manual') { - if (options.reconnectingFromReady) { - attempt.authenticated = false - attempt.announced = false - this.reconnectAttempts = 0 - this.transitionTo('authenticating') + return this.opts.maxReconnectAttempts === -1 || attempt <= this.opts.maxReconnectAttempts + }, + onFailed: (error) => { + const normalized = this.normalizeError(error, 'Failed to connect websocket client') + if (this.failureReason === normalized) { return } - attempt.authenticated = true - attempt.announced = true - this.reconnectAttempts = 0 - this.transitionTo('ready') - this.resolveAttempt() - this.opts.onReady?.() + this.failureReason = normalized + this.opts.onError(normalized) + }, + } + } + + private createHeartbeatOptions() { + if (!this.opts.heartbeat) { + return false + } + + return { + mode: 'message' as const, + interval: this.opts.heartbeat.pingInterval, + timeout: this.opts.heartbeat.readTimeout, + message: () => this.createPayload({ + type: 'transport:connection:heartbeat', + data: { + kind: MessageHeartbeatKind.Ping, + message: this.opts.heartbeat ? this.opts.heartbeat.message : MessageHeartbeat.Ping, + at: Date.now(), + }, + } as WebSocketEventOptionalSource), + } + } + + private async prepareProtocolConnection(context: PrepareContext>): Promise { + if (this.opts.handshake === 'manual') { + if (!context.reconnecting) { return } - if (this.opts.token) { - attempt.authenticated = false - this.transitionTo('authenticating') - this.tryAuthenticate() + this.transitionTo('authenticating') + await this.waitForManualReconnectHandshake(context) + return + } + + if (this.opts.token) { + this.transitionTo('authenticating') + context.send(this.createPayload({ + type: 'module:authenticate', + data: { token: this.opts.token }, + } as WebSocketEventOptionalSource)) + + await context.waitFor((message) => { + const result = this.consumePrepareMessage(message) + if (result.error) { + throw result.error + } + + return message.type === 'module:authenticated' && message.data.authenticated === true + }, { timeout: this.opts.connectTimeoutMs }) + } + + this.transitionTo('announcing') + context.send(this.createPayload({ + type: 'extension:module:announce', + data: { + name: this.opts.name, + identity: this.opts.identity, + possibleEvents: this.opts.possibleEvents, + configSchema: this.opts.configSchema, + dependencies: this.opts.dependencies, + }, + } as WebSocketEventOptionalSource)) + + await context.waitFor((message) => { + const result = this.consumePrepareMessage(message) + if (result.error) { + throw result.error } - else { - attempt.authenticated = true - this.transitionTo('announcing') - this.tryAnnounce() + + return result.ready + }, { timeout: this.opts.connectTimeoutMs }) + } + + private async waitForManualReconnectHandshake(context: PrepareContext>): Promise { + await context.waitFor((message) => { + const result = this.consumePrepareMessage(message) + if (result.error) { + throw result.error + } + + return message.type === 'peer:authenticated' && message.data.authenticated === true + }, { timeout: this.opts.connectTimeoutMs }) + + this.transitionTo('announcing') + + await context.waitFor((message) => { + const result = this.consumePrepareMessage(message) + if (result.error) { + throw result.error + } + + return message.type === 'extension:announced' + && message.data.identity.id === this.opts.extension.id + }, { timeout: this.opts.connectTimeoutMs }) + } + + private consumePrepareMessage(message: WebSocketEvent): ProtocolWaitResult { + const error = this.errorFromServerEvent(message) + if (error) { + this.failureReason = error + return { ready: false, error } + } + + if (message.type === 'extension:module:announced') { + return { ready: this.isSelfModuleAnnouncement(message) } + } + + if (message.type === 'registry:modules:sync') { + return { ready: this.hasSelfModuleInRegistrySync(message) } + } + + return { ready: false } + } + + private async handleMessage(message: WebSocketEvent): Promise { + this.opts.onAnyMessage(message) + + const error = this.errorFromServerEvent(message) + if (error) { + this.failureReason = error + this.opts.onError(error) + } + + if (message.type === 'transport:connection:heartbeat' && message.data.kind === MessageHeartbeatKind.Ping) { + this.send({ + type: 'transport:connection:heartbeat', + data: { + kind: MessageHeartbeatKind.Pong, + message: MessageHeartbeat.Pong, + at: Date.now(), + }, + } as WebSocketEventOptionalSource) + } + + const listeners = this.eventListeners.get(message.type) + if (!listeners?.size) { + return + } + + const results = await Promise.allSettled( + Array.from(listeners).map(listener => Promise.resolve(listener(message as WebSocketBaseEvent))), + ) + + for (const result of results) { + if (result.status === 'rejected') { + this.failureReason = this.normalizeError(result.reason, 'Client event listener failed') + this.opts.onError(result.reason) } } - - return attempt.promise } - private handleSocketFailure(error: Error, socket?: WebSocketLike) { - if (socket && this.websocket !== socket) { + private handleTransportStateChange(previousState: BetterWsClient>['state'], state: BetterWsClient>['state']) { + if (state === 'ready') { + this.transitionTo('ready') + this.opts.onReady() return } - const currentSocket = socket ?? this.websocket - this.cleanupSocket(socket) - - if (currentSocket && currentSocket.readyState !== this.websocketConstructor.CLOSED && currentSocket.readyState !== this.websocketConstructor.CLOSING) { - currentSocket.close() - } - - this.rejectAttempt(error) - } - - private cleanupSocket(socket?: WebSocketLike) { - if (socket && this.websocket !== socket) { + const nextStatus = this.mapTransportStatus(state) + if (!nextStatus) { return } - this.stopHeartbeat() + this.transitionTo(nextStatus) - if (!socket || this.websocket === socket) { - this.websocket = undefined + if (previousState === 'ready' && state === 'reconnecting') { + this.opts.onClose() + } + else if (state === 'closed') { + this.opts.onClose() } } - private rejectAttempt(error: Error) { - if (!this.connectionAttempt) { - return + private mapTransportStatus(state: BetterWsClient>['state']): ClientStatus | undefined { + switch (state) { + case 'idle': + case 'connecting': + case 'reconnecting': + case 'closing': + case 'closed': + case 'failed': + return state + case 'open': + case 'preparing': + case 'ready': + return undefined } - - const attempt = this.connectionAttempt - this.connectionAttempt = undefined - attempt.reject(error) - } - - private resolveAttempt() { - if (!this.connectionAttempt) { - return - } - - const attempt = this.connectionAttempt - this.connectionAttempt = undefined - attempt.resolve() - } - - private canRetry() { - return this.opts.maxReconnectAttempts === -1 || this.reconnectAttempts < this.opts.maxReconnectAttempts - } - - private getReconnectDelay(attempts: number) { - return Math.min(2 ** attempts * 1_000, 30_000) } private transitionTo(status: ClientStatus) { @@ -560,7 +582,7 @@ export class Client { this.status = status const context = { previousStatus, status } - this.opts.onStateChange?.(context) + this.opts.onStateChange(context) for (const listener of this.stateListeners) { listener(context) @@ -574,12 +596,12 @@ export class Client { const timeout = options?.timeout if (typeof timeout !== 'undefined' && timeout <= 0) { - throw new Error(`Connection timed out after ${timeout}ms`) + throw createConnectionTimeoutError(timeout) } const abortSignal = options?.abortSignal if (abortSignal?.aborted) { - throw new Error('Connection aborted') + throw createAbortError() } let timeoutHandle: ReturnType | undefined @@ -591,15 +613,12 @@ export class Client { new Promise((_, reject) => { if (typeof timeout !== 'undefined') { timeoutHandle = setTimeout(() => { - reject(new Error(`Connection timed out after ${timeout}ms`)) + reject(createConnectionTimeoutError(timeout)) }, timeout) } if (abortSignal) { - const onAbort = () => { - reject(new Error('Connection aborted')) - } - + const onAbort = () => reject(createAbortError()) abortSignal.addEventListener('abort', onAbort, { once: true }) removeAbortListener = () => abortSignal.removeEventListener('abort', onAbort) } @@ -615,345 +634,80 @@ export class Client { } } - private tryAnnounce() { - this.sendOrThrow({ - type: 'extension:module:announce', - data: { - name: this.opts.name, - identity: this.identity, - possibleEvents: this.opts.possibleEvents, - dependencies: this.opts.dependencies, - configSchema: this.opts.configSchema, - }, + private waitForReady(): Promise { + if (this.status === 'ready') { + return Promise.resolve() + } + + return new Promise((resolve, reject) => { + const dispose = this.onConnectionStateChange(({ status }) => { + if (status === 'ready') { + dispose() + resolve() + return + } + + if (status === 'failed' || status === 'closed') { + dispose() + reject(this.failureReason ?? new Error(`Client connection ended with status: ${status}`)) + } + }) }) } - private tryAuthenticate() { - if (!this.opts.token) { - return + private errorFromServerEvent(message: WebSocketEvent): Error | undefined { + if (message.type !== 'error') { + return undefined } - this.sendOrThrow({ - type: 'module:authenticate', - data: { token: this.opts.token }, - }) + const errorMessage = typeof message.data.message === 'string' + ? message.data.message + : 'Unknown server error' + const parsed = parseServerErrorMessage(errorMessage) + + if (parsed.code === 'unknown') { + return new Error(errorMessage) + } + + return new Error(parsed.message) } - private async handleMessage(event: WebSocketMessageEventLike) { - this.lastReadAt = Date.now() - - try { - const data = this.parseMessage(event.data as string) - this.opts.onAnyMessage?.(data) - - await this.handleControlMessage(data) - await this.dispatchMessage(data) - } - catch (error) { - const normalizedError = error instanceof Error ? error : new Error(errorMessageFrom(error) ?? 'Failed to handle websocket message') - this.opts.onError?.(normalizedError) - - if (this.connectionAttempt && this.status !== 'ready') { - this.handleSocketFailure(normalizedError) - } - } + private normalizeError(error: unknown, fallback: string): Error { + return error instanceof Error + ? error + : new Error(errorMessageFrom(error) ?? fallback) } - private parseMessage(raw: string): WebSocketEvent { - try { - const parsed = superjson.parse | undefined>(raw) - if (parsed && typeof parsed === 'object' && 'type' in parsed) { - return parsed - } - } - catch { - // Try standard JSON next. - } - - const parsed = JSON.parse(raw) as WebSocketEvent - if (!parsed || typeof parsed !== 'object' || !('type' in parsed)) { - throw new Error('Received invalid websocket message') - } - - return parsed + private isSelfModuleAnnouncement(event: WebSocketBaseEvent<'extension:module:announced', WebSocketEvents['extension:module:announced']>) { + return event.data.name === this.opts.name && event.data.identity?.id === this.opts.identity.id } - private async handleControlMessage(data: WebSocketEvent) { - switch (data.type) { - case 'error': { - const message = data.data?.message - if (!message || typeof message !== 'string') { - return - } - - const parsedServerError = parseServerErrorMessage(message) - if (parsedServerError.authentication) { - const error = new Error(message) - if (parsedServerError.terminal) { - this.shouldClose = true - this.handleSocketFailure(error) - this.transitionTo('failed') - return - } - - await this.reconnectAfterProtocolError(error) - return - } - - if (parsedServerError.code !== 'unknown') { - throw new Error(parsedServerError.message) - } - - throw new Error(message) - } - - case 'module:authenticated': { - if (data.data.authenticated) { - if (!this.connectionAttempt || this.connectionAttempt.authenticated) { - return - } - - this.connectionAttempt.authenticated = true - this.transitionTo('announcing') - this.tryAnnounce() - return - } - - throw new Error('Authentication failed') - } - - case 'peer:authenticated': { - if (this.opts.handshake !== 'manual' || this.status !== 'authenticating' || !this.connectionAttempt) { - return - } - - if (data.data.authenticated) { - this.connectionAttempt.authenticated = true - this.transitionTo('announcing') - return - } - - throw new Error('Peer authentication failed') - } - - case 'extension:announced': { - if (this.opts.handshake !== 'manual' || this.status !== 'announcing' || !this.connectionAttempt) { - return - } - - this.connectionAttempt.announced = true - this.reconnectAttempts = 0 - this.transitionTo('ready') - this.resolveAttempt() - this.opts.onReady?.() - return - } - - case 'extension:module:announced': { - if (!this.isSelfAnnouncement(data)) { - return - } - - if (this.status === 'ready') { - return - } - - if (this.connectionAttempt) { - this.connectionAttempt.announced = true - } - - this.reconnectAttempts = 0 - this.transitionTo('ready') - this.resolveAttempt() - this.opts.onReady?.() - return - } - - case 'registry:modules:sync': { - // Fallback: If the status is stuck at 'announcing' but the sync already contains this module, - // it means the announce succeeded; the server simply didn't send back 'extension:module:announced' - if (this.status !== 'announcing' || !this.connectionAttempt) { - return - } - - const syncData = data.data as { - modules?: Array<{ - name: string - identity?: { id?: string } - }> - } | unknown - const modules = Array.isArray((syncData as any)?.modules) ? (syncData as any).modules : [] - - const selfRegistered = modules.some( - m => m.name === this.opts.name - && m.identity?.id === this.identity.id, - ) - - if (!selfRegistered) { - return - } - - if (this.connectionAttempt) { - this.connectionAttempt.announced = true - } - - this.reconnectAttempts = 0 - this.transitionTo('ready') - this.resolveAttempt() - this.opts.onReady?.() - return - } - - case 'transport:connection:heartbeat': { - if (data.data.kind === MessageHeartbeatKind.Ping) { - this.sendHeartbeatPong() - } - } - } - } - - private isSelfAnnouncement(event: WebSocketBaseEvent<'extension:module:announced', WebSocketEvents['extension:module:announced']>) { - return event.data.name === this.opts.name && event.data.identity?.id === this.identity.id - } - - private async dispatchMessage(data: WebSocketEvent) { - const listeners = this.eventListeners.get(data.type) - if (!listeners?.size) { - return - } - - // Cast is necessary here because the Set stores callbacks from potentially different event types, - // but we're only calling listeners registered for this specific event type - const results = await Promise.allSettled( - Array.from(listeners).map(listener => Promise.resolve((listener as (data: WebSocketEvent) => void | Promise)(data))), + private hasSelfModuleInRegistrySync(event: WebSocketBaseEvent<'registry:modules:sync', WebSocketEvents['registry:modules:sync']>) { + return event.data.modules.some(module => + module.name === this.opts.name + && module.identity?.id === this.opts.identity.id, ) - - for (const result of results) { - if (result.status === 'rejected') { - this.opts.onError?.(result.reason) - } - } } private createPayload(data: WebSocketEventOptionalSource) { return { ...data, metadata: { - ...data?.metadata, - source: data?.metadata?.source ?? this.identity, + ...data.metadata, + source: data.metadata?.source ?? { + kind: 'plugin', + ...this.opts.identity, + plugin: { id: this.opts.extension.id }, + }, event: { - id: data?.metadata?.event?.id ?? createEventId(), - ...data?.metadata?.event, + ...data.metadata?.event, + id: data.metadata?.event?.id ?? createEventId(), }, }, } as WebSocketEvent } +} - private startHeartbeat() { - if (!this.heartbeat.readTimeout || !this.heartbeat.pingInterval) { - return - } - - this.stopHeartbeat() - this.lastReadAt = Date.now() - this.lastPingAt = 0 - - const interval = Math.max(1_000, Math.min(this.heartbeat.pingInterval, Math.floor(this.heartbeat.readTimeout / 2))) - this.heartbeatTimer = setInterval(() => { - if (!this.isSocketOpen) { - return - } - - const now = Date.now() - if (now - this.lastReadAt > this.heartbeat.readTimeout) { - void this.reconnectAfterProtocolError(new Error(`Read timeout after ${this.heartbeat.readTimeout}ms`)) - return - } - - if (now - this.lastPingAt >= this.heartbeat.pingInterval) { - this.sendHeartbeatPing() - } - }, interval) - } - - private stopHeartbeat() { - if (!this.heartbeatTimer) { - return - } - - clearInterval(this.heartbeatTimer) - this.heartbeatTimer = undefined - } - - private sendNativeHeartbeat(kind: 'ping' | 'pong') { - const websocket = this.websocket as WebSocketLike & { - ping?: () => void - pong?: () => void - } - - if (kind === 'ping') { - websocket.ping?.() - } - else { - websocket.pong?.() - } - } - - private sendHeartbeatPing() { - this.lastPingAt = Date.now() - this.send({ - type: 'transport:connection:heartbeat', - data: { - kind: MessageHeartbeatKind.Ping, - message: this.heartbeat.message, - at: Date.now(), - }, - }) - this.sendNativeHeartbeat('ping') - } - - private sendHeartbeatPong() { - this.send({ - type: 'transport:connection:heartbeat', - data: { - kind: MessageHeartbeatKind.Pong, - message: MessageHeartbeat.Pong, - at: Date.now(), - }, - }) - this.sendNativeHeartbeat('pong') - } - - private async reconnectAfterProtocolError(error: Error) { - if (this.shouldClose || this.pendingReconnect) { - return - } - - this.pendingReconnect = true - const hadSocket = !!this.websocket - - if (!this.connectionAttempt || this.status === 'ready') { - this.opts.onError?.(error) - } - - const websocket = this.websocket - this.cleanupSocket(websocket) - this.rejectAttempt(error) - - if (websocket && websocket.readyState !== this.websocketConstructor.CLOSED && websocket.readyState !== this.websocketConstructor.CLOSING) { - websocket.close() - } - - if (hadSocket) { - this.opts.onClose?.() - } - - if (!this.opts.autoReconnect) { - this.transitionTo('failed') - return - } - - this.transitionTo('idle') - void this.connect() - } +export function createClient(options: ClientOptions): Client { + return new Client(options) } diff --git a/packages/server-sdk/src/codec.ts b/packages/server-sdk/src/codec.ts new file mode 100644 index 000000000..642d47e49 --- /dev/null +++ b/packages/server-sdk/src/codec.ts @@ -0,0 +1,81 @@ +import type { WebSocketBaseEvent, WebSocketEvent } from '@proj-airi/server-shared/types' + +import { parse, stringify } from 'superjson' +import { check, objectWithRest, pipe, safeParse, string, unknown } from 'valibot' + +const invalidMessage = 'Invalid AIRI websocket message.' + +const eventDataSchema = pipe( + unknown(), + check( + value => Boolean(value) && typeof value === 'object' && !Array.isArray(value), + 'Expected event data to be a non-array object.', + ), +) + +const eventEnvelopeSchema = objectWithRest({ + type: string(), + data: eventDataSchema, +}, unknown()) + +/** Options for websocket message validation failures. */ +export interface InvalidMessageErrorOptions { + /** Original parser or validator failure that made the websocket message unusable. */ + cause?: unknown + /** Parsed candidate event when available; otherwise the original websocket text. */ + source?: unknown +} + +/** Error thrown when websocket text cannot be parsed as an AIRI event envelope. */ +export class InvalidMessageError extends Error { + readonly source?: unknown + + constructor(options: InvalidMessageErrorOptions = {}) { + super(invalidMessage, { cause: options.cause }) + this.name = 'InvalidMessageError' + this.source = options.source + } +} + +/** Parses one AIRI websocket protocol event from SuperJSON or plain JSON text. */ +export function parseEvent(text: string): WebSocketEvent { + let superJsonParsed: WebSocketEvent | undefined + let superJsonError: unknown + + try { + superJsonParsed = parse>(text) + } + catch (error) { + superJsonError = error + } + + const potentialEvent = superJsonParsed && typeof superJsonParsed === 'object' && 'type' in superJsonParsed + ? superJsonParsed + : parsePlainJson(text, superJsonError) + + const result = safeParse(eventEnvelopeSchema, potentialEvent) + if (!result.success) { + throw new InvalidMessageError({ cause: result.issues, source: potentialEvent }) + } + + return potentialEvent as WebSocketEvent +} + +/** Serializes one AIRI websocket protocol event with SuperJSON. */ +export function stringifyEvent( + event: WebSocketBaseEvent | WebSocketEvent, +) { + return stringify(event) +} + +function parsePlainJson(text: string, superJsonError: unknown): unknown { + try { + return JSON.parse(text) + } + catch (jsonError) { + throw new InvalidMessageError({ + cause: superJsonError ?? jsonError, + source: text, + }) + } +} diff --git a/packages/server-sdk/src/extension-peer.ts b/packages/server-sdk/src/extension-peer.ts index d4e5c2561..4b8ee17e3 100644 --- a/packages/server-sdk/src/extension-peer.ts +++ b/packages/server-sdk/src/extension-peer.ts @@ -9,9 +9,9 @@ import type { WebSocketEvents, } from '@proj-airi/server-shared/types' -import type { Client, ClientOptions, ConnectOptions } from './client' +import type { ClientOptions, ConnectOptions } from './client' -import { Client as WebSocketClient } from './client' +import { createClient } from './client' /** * Describes the client operations required by {@link WebSocketExtensionPeer}. @@ -19,15 +19,10 @@ import { Client as WebSocketClient } from './client' * @param C - Optional custom protocol event map carried by the websocket client. */ export interface ExtensionPeerClient { - /** Opens the underlying websocket client connection. */ connect: (options?: ConnectOptions) => Promise - /** Sends one typed websocket event and reports whether it was accepted by the transport. */ send: (data: WebSocketEventOptionalSource) => boolean - /** Sends one typed websocket event or throws when the transport is unavailable. */ sendOrThrow: (data: WebSocketEventOptionalSource) => void - /** Closes the underlying websocket client connection. */ close: () => void - /** Registers a typed event listener when backed by the standard server-sdk Client. */ onEvent?: >( event: E, callback: (data: WebSocketBaseEvent[E]>) => void | Promise, @@ -57,32 +52,19 @@ export interface AnnounceExtensionModuleInput { } /** - * Options for creating a websocket-backed extension peer. + * Options for creating a websocket-backed extension peer. Supplying `client` lets + * tests and embedding runtimes provide their own protocol client implementation. * * @param C - Optional custom protocol event map carried by the websocket client. */ export interface WebSocketExtensionPeerOptions { - /** Extension session identity announced after peer authentication. */ extension: ExtensionIdentity - /** Optional prebuilt client used by tests or embedding runtimes. */ client?: ExtensionPeerClient - /** Standard server-sdk Client options used when `client` is not supplied. */ clientOptions?: Omit, 'name' | 'identity'> } /** - * Provides extension-level protocol helpers over the existing websocket Client. - * - * Use when: - * - A remote extension talks to an AIRI host over websocket transport - * - Authoring/runtime code should say peer/extension/module explicitly instead of sending raw websocket events - * - * Expects: - * - The underlying client owns websocket lifecycle and serialization - * - The host interprets `peer:*`, `extension:*`, and `extension:module:*` protocol events - * - * Returns: - * - A thin transport peer that delegates connection and event sending to server-sdk Client + * Provides extension-level protocol helpers over a server-sdk protocol client. */ export class WebSocketExtensionPeer { private readonly client: ExtensionPeerClient @@ -90,43 +72,19 @@ export class WebSocketExtensionPeer { constructor(options: WebSocketExtensionPeerOptions) { this.extension = options.extension - this.client = options.client ?? new WebSocketClient({ + this.client = options.client ?? createClient({ ...options.clientOptions, name: options.extension.id, handshake: 'manual', autoConnect: options.clientOptions?.autoConnect ?? false, autoReconnect: options.clientOptions?.autoReconnect ?? false, - }) as Client + }) } - /** - * Opens the underlying websocket connection. - * - * Use when: - * - The extension transport should begin peer authentication or announcement - * - * Expects: - * - The wrapped client can reach the configured websocket URL - * - * Returns: - * - Resolves when the wrapped client reports readiness - */ connect(options?: ConnectOptions): Promise { return this.client.connect(options) } - /** - * Sends transport-level peer authentication. - * - * Use when: - * - A websocket connection needs to authenticate before extension session grant - * - * Expects: - * - The websocket connection is already open or the client can queue/send immediately - * - * Returns: - * - Nothing; send failures are surfaced by the wrapped client - */ authenticatePeer(input: { token?: string, peerId?: string } = {}): void { this.client.sendOrThrow({ type: 'peer:authenticate', @@ -134,18 +92,6 @@ export class WebSocketExtensionPeer { }) } - /** - * Announces the extension session after peer authentication. - * - * Use when: - * - The remote peer has permission to enter extension setup - * - * Expects: - * - Permissions represent the extension-level ceiling grant or declaration snapshot - * - * Returns: - * - Nothing; send failures are surfaced by the wrapped client - */ announceExtension(input: { permissions?: ModulePermissionDeclaration } = {}): void { this.client.sendOrThrow({ type: 'extension:announce', @@ -156,18 +102,6 @@ export class WebSocketExtensionPeer { }) } - /** - * Announces one module registered by the current extension. - * - * Use when: - * - A websocket extension dynamically registers module capabilities - * - * Expects: - * - `id` is stable within this extension session - * - * Returns: - * - Nothing; send failures are surfaced by the wrapped client - */ announceModule(input: AnnounceExtensionModuleInput): void { this.client.sendOrThrow({ type: 'extension:module:announce', @@ -186,34 +120,10 @@ export class WebSocketExtensionPeer { }) } - /** - * Sends a typed websocket event through the wrapped client. - * - * Use when: - * - Runtime code has a protocol event not covered by helper methods - * - * Expects: - * - Callers pass a server-shared websocket event - * - * Returns: - * - Whether the event was accepted by the underlying transport - */ send(data: WebSocketEventOptionalSource): boolean { return this.client.send(data) } - /** - * Registers one event listener when the wrapped client supports typed listeners. - * - * Use when: - * - The remote extension needs to observe host protocol events - * - * Expects: - * - Test doubles may omit listener support - * - * Returns: - * - A disposer that removes the listener - */ onEvent>( event: E, callback: (data: WebSocketBaseEvent[E]>) => void | Promise, @@ -225,35 +135,12 @@ export class WebSocketExtensionPeer { return this.client.onEvent(event, callback) } - /** - * Closes the underlying websocket client. - * - * Use when: - * - The extension transport is disposed - * - * Expects: - * - Close is idempotent in the wrapped client - * - * Returns: - * - Nothing - */ close(): void { this.client.close() } } -/** - * Creates a websocket extension peer over server-sdk Client. - * - * Use when: - * - Code prefers a function factory over direct class construction - * - * Expects: - * - `extension.id` is the stable extension id - * - * Returns: - * - A {@link WebSocketExtensionPeer} ready to connect and announce - */ +/** Creates a websocket extension peer over a server-sdk protocol client. */ export function createWebSocketExtensionPeer( options: WebSocketExtensionPeerOptions, ): WebSocketExtensionPeer { diff --git a/packages/server-sdk/src/index.ts b/packages/server-sdk/src/index.ts index 954be7f61..d0255bb05 100644 --- a/packages/server-sdk/src/index.ts +++ b/packages/server-sdk/src/index.ts @@ -1,5 +1,5 @@ export * from './client' +export * from './codec' export * from './extension-peer' -export type * from './websocket-like' export type * from '@proj-airi/server-shared/types' export { ContextUpdateStrategy, WebSocketEventSource } from '@proj-airi/server-shared/types' diff --git a/packages/server-sdk/src/websocket-like.ts b/packages/server-sdk/src/websocket-like.ts deleted file mode 100644 index d5e0763aa..000000000 --- a/packages/server-sdk/src/websocket-like.ts +++ /dev/null @@ -1,30 +0,0 @@ -export interface WebSocketMessageEventLike { - data: T -} - -export interface WebSocketErrorEventLike { - error?: Error | unknown -} - -export interface WebSocketLike { - readonly readyState: number - - onopen?: (event?: unknown) => void - onmessage?: (event: WebSocketMessageEventLike) => void - onerror?: (event: WebSocketErrorEventLike | unknown) => void - onclose?: (event?: unknown) => void - - send: (data: string | ArrayBufferLike | ArrayBufferView) => void - close: (code?: number, reason?: string) => void - - ping?: () => void - pong?: () => void -} - -export interface WebSocketLikeConstructor { - readonly OPEN: number - readonly CLOSING: number - readonly CLOSED: number - - new (url: string): WebSocketLike -} diff --git a/packages/server-sdk/test/client.test.ts b/packages/server-sdk/test/client.test.ts index adc29039d..8154b1f19 100644 --- a/packages/server-sdk/test/client.test.ts +++ b/packages/server-sdk/test/client.test.ts @@ -1,505 +1,480 @@ +import type { ClientConnection, ClientConnector, ClientEvents } from '@proj-airi/better-ws' import type { WebSocketEvent, WebSocketEventOf } from '@proj-airi/server-shared/types' -import superjson from 'superjson' - import { afterEach, describe, expect, it, vi } from 'vitest' import { Client } from '../src/client' -import { createWebSocketExtensionPeer } from '../src/extension-peer' -const { InjectedMockWebSocket, MockWebSocket } = vi.hoisted(() => { - class MockWebSocket { - static readonly CONNECTING = 0 - static readonly OPEN = 1 - static readonly CLOSING = 2 - static readonly CLOSED = 3 +class Deferred { + promise: Promise + resolve!: (value: T) => void + reject!: (error: unknown) => void - static instances: MockWebSocket[] = [] + constructor() { + this.promise = new Promise((resolve, reject) => { + this.resolve = resolve + this.reject = reject + }) + } +} - readonly sent: Array> = [] - readyState = MockWebSocket.CONNECTING - onclose?: () => void - onerror?: (event: { error?: Error } | unknown) => void - onmessage?: (event: { data: string | ArrayBufferLike | ArrayBufferView }) => void - onopen?: () => void +class FakeConnection implements ClientConnection> { + readonly sent: Array> = [] + readonly pongs: number[] = [] + closed = false - constructor(public readonly url: string) { - MockWebSocket.instances.push(this) - } + constructor(private readonly events: ClientEvents>) {} - send(data: string | ArrayBufferLike | ArrayBufferView) { - this.sent.push(data) - } - - close() { - this.readyState = MockWebSocket.CLOSED - this.onclose?.() - } - - ping() {} - pong() {} + send(message: WebSocketEvent) { + this.sent.push(message) + return true } - class InjectedMockWebSocket extends MockWebSocket { - static instances: InjectedMockWebSocket[] = [] - - constructor(url: string) { - super(url) - InjectedMockWebSocket.instances.push(this) + close() { + if (this.closed) { + return } + + this.closed = true + this.events.close({ code: 1000, reason: 'closed', wasClean: true }) } + pong() { + this.pongs.push(Date.now()) + return true + } +} + +class FakeConnector implements ClientConnector> { + readonly attempts: Array<{ + deferred: Deferred>> + events: ClientEvents> + connection?: FakeConnection + }> = [] + + connect(events: ClientEvents>) { + const deferred = new Deferred>>() + this.attempts.push({ deferred, events }) + return deferred.promise + } + + open(index = this.attempts.length - 1) { + const attempt = this.attempts[index] + if (!attempt) { + throw new Error(`Missing fake connector attempt at index ${index}.`) + } + + const connection = new FakeConnection(attempt.events) + attempt.connection = connection + attempt.deferred.resolve(connection) + return connection + } + + reject(error: unknown, index = this.attempts.length - 1) { + const attempt = this.attempts[index] + if (!attempt) { + throw new Error(`Missing fake connector attempt at index ${index}.`) + } + + attempt.deferred.reject(error) + } + + emit(message: WebSocketEvent, index = this.attempts.length - 1) { + const attempt = this.attempts[index] + if (!attempt) { + throw new Error(`Missing fake connector attempt at index ${index}.`) + } + + attempt.events.message(message) + } +} + +function serverEvent( + type: E, + data: WebSocketEventOf['data'], +): WebSocketEventOf { return { - InjectedMockWebSocket, - MockWebSocket, - } -}) - -vi.mock('crossws/websocket', () => ({ - default: MockWebSocket, -})) - -function lastSocket() { - const socket = MockWebSocket.instances.at(-1) - if (!socket) { - throw new Error('No mock websocket instance created') - } - - return socket + type, + data, + metadata: { + source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, + event: { id: `${type}-1` }, + }, + } as WebSocketEventOf } -function parseSent(socket: InstanceType, index = -1) { - const payload = socket.sent.at(index) - if (!payload) { - throw new Error(`No sent payload at index ${index}`) - } - if (typeof payload === 'string') { - return superjson.parse(payload) - } - - const textDecoder = new TextDecoder() - const decoded = textDecoder.decode(payload) - - return superjson.parse(decoded) +async function flushMicrotasks() { + await Promise.resolve() + await Promise.resolve() } -function emitOpen(socket: InstanceType) { - socket.readyState = MockWebSocket.OPEN - socket.onopen?.() -} - -function emitMessage(socket: InstanceType, event: WebSocketEvent) { - socket.onmessage?.({ - data: superjson.stringify(event), - }) +async function flushAsyncTasks() { + await flushMicrotasks() + await new Promise(resolve => setTimeout(resolve, 0)) } afterEach(() => { - MockWebSocket.instances.length = 0 - InjectedMockWebSocket.instances.length = 0 vi.useRealTimers() }) describe('client', () => { - it('resolves connect only after authentication and self announcement', async () => { + it('routes default autoConnect failures through onError without unhandled rejections', async () => { + const connector = new FakeConnector() + const onError = vi.fn() + const unhandledRejections: unknown[] = [] + const onUnhandledRejection = (reason: unknown) => { + unhandledRejections.push(reason) + } + process.on('unhandledRejection', onUnhandledRejection) + + let client: Client | undefined + try { + client = new Client({ + autoReconnect: false, + connector, + name: 'test-plugin', + onError, + }) + + const failure = new Error('server unavailable') + connector.reject(failure) + await flushAsyncTasks() + + expect(onError).toHaveBeenCalledWith(failure) + expect(unhandledRejections).toEqual([]) + } + finally { + client?.close() + process.off('unhandledRejection', onUnhandledRejection) + } + }) + + it('runs module authentication and announcement in the better-ws prepare step', async () => { + const connector = new FakeConnector() const client = new Client({ autoConnect: false, autoReconnect: false, + connector, name: 'test-plugin', token: 'secret', }) const connected = client.connect() - const socket = lastSocket() + const connection = connector.open() + await flushMicrotasks() - emitOpen(socket) - - expect(parseSent(socket)).toMatchObject({ + expect(client.connectionStatus).toBe('authenticating') + expect(connection.sent.at(-1)).toMatchObject({ type: 'module:authenticate', data: { token: 'secret' }, }) - emitMessage(socket, { - type: 'module:authenticated', - data: { authenticated: true }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'auth-1' }, - }, - }) + connector.emit(serverEvent('module:authenticated', { authenticated: true })) + await flushMicrotasks() - const announceEvent = parseSent(socket) as WebSocketEventOf<'extension:module:announce'> + const announceEvent = connection.sent.at(-1) as WebSocketEventOf<'extension:module:announce'> + expect(client.connectionStatus).toBe('announcing') expect(announceEvent).toMatchObject({ type: 'extension:module:announce', data: { name: 'test-plugin' }, }) - emitMessage(socket, { - type: 'extension:module:announced', - data: { - name: 'test-plugin', - identity: announceEvent.data.identity, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'announce-1' }, - }, - }) + connector.emit(serverEvent('extension:module:announced', { + name: 'test-plugin', + identity: announceEvent.data.identity, + })) await expect(connected).resolves.toBeUndefined() expect(client.connectionStatus).toBe('ready') expect(client.isReady).toBe(true) }) - it('fails terminally on invalid token', async () => { - const client = new Client({ - autoConnect: false, - autoReconnect: true, - name: 'test-plugin', - token: 'wrong-token', - }) - - const connected = client.connect() - const socket = lastSocket() - - emitOpen(socket) - emitMessage(socket, { - type: 'error', - data: { message: 'invalid token' }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'error-1' }, - }, - }) - - await expect(connected).rejects.toThrow('invalid token') - expect(client.connectionStatus).toBe('failed') - }) - - it('returns an unsubscribe function from onEvent', () => { + it('accepts registry sync as the module announcement completion signal', async () => { + const connector = new FakeConnector() + const onReady = vi.fn() const client = new Client({ autoConnect: false, autoReconnect: false, + connector, name: 'test-plugin', - }) - - const listener = vi.fn() - const dispose = client.onEvent('input:text', listener) - - dispose() - expect(() => client.offEvent('input:text', listener)).not.toThrow() - }) - - it('uses an injected websocket constructor when provided', async () => { - const client = new Client({ - autoConnect: false, - autoReconnect: false, - name: 'test-plugin', - websocketConstructor: InjectedMockWebSocket, + onReady, }) const connected = client.connect() - const socket = InjectedMockWebSocket.instances.at(-1) + const connection = connector.open() + await flushMicrotasks() - expect(socket).toBeDefined() - expect(MockWebSocket.instances).toHaveLength(1) + const announceEvent = connection.sent.at(-1) as WebSocketEventOf<'extension:module:announce'> - if (!socket) { - throw new Error('No custom mock websocket instance created') - } - - emitOpen(socket) - const announceEvent = parseSent(socket) as WebSocketEventOf<'extension:module:announce'> - - emitMessage(socket, { - type: 'extension:module:announced', - data: { - name: 'test-plugin', - identity: announceEvent.data.identity, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'announce-1' }, - }, - }) - - await expect(connected).resolves.toBeUndefined() - }) - - it('supports manual handshake for extension peers without legacy module announce', async () => { - const client = new Client({ - autoConnect: false, - autoReconnect: false, - handshake: 'manual', - name: 'test-extension', - }) - - const connected = client.connect() - const socket = lastSocket() - - emitOpen(socket) + connector.emit(serverEvent('registry:modules:sync', { + modules: [{ name: 'test-plugin', identity: announceEvent.data.identity }], + })) + connector.emit(serverEvent('extension:module:announced', { + name: 'test-plugin', + identity: announceEvent.data.identity, + })) await expect(connected).resolves.toBeUndefined() expect(client.connectionStatus).toBe('ready') - expect(socket.sent).toHaveLength(0) + expect(onReady).toHaveBeenCalledTimes(1) }) - it('keeps manual reconnects non-ready until the peer reauthenticates and reannounces', async () => { + it('keeps manual automatic reconnects non-ready until peer authentication and extension announcement arrive', async () => { + vi.useFakeTimers() + + const connector = new FakeConnector() const onReady = vi.fn() const client = new Client({ autoConnect: false, autoReconnect: true, + connector, handshake: 'manual', name: 'test-extension', onReady, }) const connected = client.connect() - const firstSocket = lastSocket() - - emitOpen(firstSocket) - + const firstConnection = connector.open(0) + await flushMicrotasks() await expect(connected).resolves.toBeUndefined() + expect(client.connectionStatus).toBe('ready') expect(onReady).toHaveBeenCalledTimes(1) - firstSocket.close() - const secondSocket = lastSocket() + firstConnection.close() + await vi.advanceTimersByTimeAsync(1_000) + expect(connector.attempts).toHaveLength(2) - emitOpen(secondSocket) + const secondConnection = connector.open(1) + await flushMicrotasks() expect(client.connectionStatus).toBe('authenticating') - expect(onReady).toHaveBeenCalledTimes(1) + expect(client.isReady).toBe(false) + expect(secondConnection.sent).toEqual([]) - emitMessage(secondSocket, { - type: 'peer:authenticated', - data: { authenticated: true }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'peer-auth-1' }, - }, - }) + connector.emit(serverEvent('peer:authenticated', { authenticated: true }), 1) + await flushMicrotasks() expect(client.connectionStatus).toBe('announcing') expect(onReady).toHaveBeenCalledTimes(1) - emitMessage(secondSocket, { - type: 'extension:announced', - data: { - identity: { id: 'test-extension' }, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'extension-announce-1' }, - }, - }) + connector.emit(serverEvent('extension:announced', { + identity: { id: 'other-extension' }, + }), 1) + await flushMicrotasks() + + expect(client.connectionStatus).toBe('announcing') + expect(onReady).toHaveBeenCalledTimes(1) + + connector.emit(serverEvent('extension:announced', { + identity: { id: 'test-extension' }, + }), 1) await expect(client.ensureConnected()).resolves.toBeUndefined() expect(client.connectionStatus).toBe('ready') expect(onReady).toHaveBeenCalledTimes(2) }) - it('uses manual handshake when creating websocket extension peers', async () => { - const peer = createWebSocketExtensionPeer({ - extension: { - id: 'test-extension', - sessionId: 'session-1', - }, - clientOptions: { - autoReconnect: false, - }, - }) - - const connected = peer.connect() - const socket = lastSocket() - - emitOpen(socket) - await expect(connected).resolves.toBeUndefined() - - expect(socket.sent).toHaveLength(0) - - peer.authenticatePeer({ token: 'secret', peerId: 'peer-1' }) - expect(parseSent(socket)).toMatchObject({ - type: 'peer:authenticate', - data: { - token: 'secret', - peerId: 'peer-1', - }, - }) - }) - - it('supports timeout-aware ensureConnected without cancelling the shared connect task', async () => { - vi.useFakeTimers() - + it('injects source and event metadata before sending through better-ws', async () => { + const connector = new FakeConnector() + const onAnySend = vi.fn() const client = new Client({ autoConnect: false, autoReconnect: false, + connector, + handshake: 'manual', + name: 'test-extension', + onAnySend, + }) + + const connected = client.connect() + const connection = connector.open() + await connected + + const sent = client.send({ + type: 'input:text', + data: { text: 'hello' }, + }) + + expect(sent).toBe(true) + expect(connection.sent.at(-1)).toMatchObject({ + type: 'input:text', + data: { text: 'hello' }, + metadata: { + source: { + kind: 'plugin', + id: expect.any(String), + plugin: { id: 'test-extension' }, + }, + event: { id: expect.any(String) }, + }, + }) + expect(onAnySend).toHaveBeenCalledWith(connection.sent.at(-1)) + }) + + it('fails without retrying terminal authentication errors', async () => { + vi.useFakeTimers() + + const connector = new FakeConnector() + const onError = vi.fn() + const client = new Client({ + autoConnect: false, + autoReconnect: true, + connector, + name: 'test-plugin', + onError, + token: 'wrong-token', + }) + + const connected = client.connect() + connector.open() + await flushMicrotasks() + + connector.emit(serverEvent('error', { message: 'invalid token' })) + + await expect(connected).rejects.toThrow('invalid token') + expect(client.connectionStatus).toBe('failed') + expect(connector.attempts).toHaveLength(1) + expect(onError).toHaveBeenCalledTimes(1) + expect(onError).toHaveBeenCalledWith(expect.any(Error)) + }) + + it('dispatches typed events and answers transport heartbeat pings with pong', async () => { + const connector = new FakeConnector() + const listener = vi.fn() + const onAnyMessage = vi.fn() + const client = new Client({ + autoConnect: false, + autoReconnect: false, + connector, + handshake: 'manual', + name: 'test-extension', + onAnyMessage, + }) + + const connected = client.connect() + const connection = connector.open() + await connected + + client.onEvent('input:text', listener) + + const input = serverEvent('input:text', { text: 'hello' }) + connector.emit(input) + connector.emit(serverEvent('transport:connection:heartbeat', { + kind: 'ping', + message: 'ping', + })) + await flushMicrotasks() + + expect(listener).toHaveBeenCalledWith(input) + expect(onAnyMessage).toHaveBeenCalledWith(input) + expect(connection.sent.at(-1)).toMatchObject({ + type: 'transport:connection:heartbeat', + data: { kind: 'pong' }, + }) + }) + + it('keeps generated event ids when caller metadata has an undefined id', async () => { + const connector = new FakeConnector() + const client = new Client({ + autoConnect: false, + autoReconnect: false, + connector, + handshake: 'manual', + name: 'test-extension', + }) + + const connected = client.connect() + const connection = connector.open() + await connected + + client.send({ + type: 'input:text', + data: { text: 'hello' }, + metadata: { + event: { id: undefined }, + }, + }) + + expect(connection.sent.at(-1)?.metadata.event.id).toEqual(expect.any(String)) + }) + + it('can disable protocol heartbeat', async () => { + const connector = new FakeConnector() + const client = new Client({ + autoConnect: false, + autoReconnect: false, + connector, + heartbeat: false, + handshake: 'manual', + name: 'test-extension', + }) + + const connected = client.connect() + connector.open() + await connected + + expect(client.isReady).toBe(true) + }) + + it('races local timeouts without cancelling the shared connect task', async () => { + vi.useFakeTimers() + + const connector = new FakeConnector() + const client = new Client({ + autoConnect: false, + autoReconnect: false, + connector, name: 'test-plugin', }) const timedOut = client.ensureConnected({ timeout: 50 }) const timedOutAssertion = expect(timedOut).rejects.toThrow('Connection timed out after 50ms') - const socket = lastSocket() await vi.advanceTimersByTimeAsync(50) await timedOutAssertion - emitOpen(socket) - const announceEvent = parseSent(socket) as WebSocketEventOf<'extension:module:announce'> + const connection = connector.open() + await flushMicrotasks() - emitMessage(socket, { - type: 'extension:module:announced', - data: { - name: 'test-plugin', - identity: announceEvent.data.identity, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'announce-1' }, - }, - }) + const announceEvent = connection.sent.at(-1) as WebSocketEventOf<'extension:module:announce'> + connector.emit(serverEvent('extension:module:announced', { + name: 'test-plugin', + identity: announceEvent.data.identity, + })) await expect(client.ensureConnected()).resolves.toBeUndefined() expect(client.isReady).toBe(true) }) - it('supports abort-aware connect', async () => { + it('races local aborts without cancelling the shared connect task', async () => { + const connector = new FakeConnector() const client = new Client({ autoConnect: false, autoReconnect: false, + connector, name: 'test-plugin', }) const controller = new AbortController() const connecting = client.connect({ abortSignal: controller.signal }) - lastSocket() controller.abort() await expect(connecting).rejects.toThrow('Connection aborted') expect(client.connectionStatus).toBe('connecting') - }) - it('notifies external state listeners', async () => { - const client = new Client({ - autoConnect: false, - autoReconnect: false, + const connection = connector.open() + await flushMicrotasks() + + const announceEvent = connection.sent.at(-1) as WebSocketEventOf<'extension:module:announce'> + connector.emit(serverEvent('extension:module:announced', { name: 'test-plugin', - }) + identity: announceEvent.data.identity, + })) - const listener = vi.fn() - const dispose = client.onConnectionStateChange(listener) - const connected = client.connect() - const socket = lastSocket() - - emitOpen(socket) - - const announceEvent = parseSent(socket) as WebSocketEventOf<'extension:module:announce'> - - emitMessage(socket, { - type: 'extension:module:announced', - data: { - name: 'test-plugin', - identity: announceEvent.data.identity, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'announce-1' }, - }, - }) - - await connected - - expect(listener).toHaveBeenCalledWith({ previousStatus: 'idle', status: 'connecting' }) - expect(listener).toHaveBeenCalledWith({ previousStatus: 'connecting', status: 'announcing' }) - expect(listener).toHaveBeenCalledWith({ previousStatus: 'announcing', status: 'ready' }) - - dispose() - }) - - it('retries after connect timeout and eventually connects on a later socket', async () => { - vi.useFakeTimers() - - const client = new Client({ - autoConnect: false, - autoReconnect: true, - connectTimeoutMs: 50, - name: 'test-plugin', - }) - - const connecting = client.connect() - const firstSocket = lastSocket() - const firstCloseSpy = vi.spyOn(firstSocket, 'close') - - await vi.advanceTimersByTimeAsync(50) - expect(firstCloseSpy).toHaveBeenCalledTimes(1) - - await vi.advanceTimersByTimeAsync(1_000) - expect(MockWebSocket.instances).toHaveLength(2) - - const secondSocket = lastSocket() - emitOpen(secondSocket) - - const announceEvent = parseSent(secondSocket) as WebSocketEventOf<'extension:module:announce'> - - emitMessage(secondSocket, { - type: 'extension:module:announced', - data: { - name: 'test-plugin', - identity: announceEvent.data.identity, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'announce-retry-1' }, - }, - }) - - await expect(connecting).resolves.toBeUndefined() - expect(client.connectionStatus).toBe('ready') - }) - - it('does not emit onReady twice when sync fallback already moved status to ready', async () => { - const onReady = vi.fn() - const client = new Client({ - autoConnect: false, - autoReconnect: false, - name: 'test-plugin', - onReady, - }) - - const connecting = client.connect() - const socket = lastSocket() - emitOpen(socket) - - const announceEvent = parseSent(socket) as WebSocketEventOf<'extension:module:announce'> - - const selfIdentity = announceEvent.data.identity - - emitMessage(socket, { - type: 'registry:modules:sync', - data: { - modules: [{ name: 'test-plugin', identity: selfIdentity }], - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'sync-1' }, - }, - }) - - emitMessage(socket, { - type: 'extension:module:announced', - data: { - name: 'test-plugin', - identity: selfIdentity, - }, - metadata: { - source: { kind: 'plugin', plugin: { id: 'server' }, id: 'server-1' }, - event: { id: 'announce-1' }, - }, - }) - - await expect(connecting).resolves.toBeUndefined() - expect(onReady).toHaveBeenCalledTimes(1) + await expect(client.ready()).resolves.toBeUndefined() + expect(client.isReady).toBe(true) }) }) diff --git a/packages/server-sdk/test/codec.test.ts b/packages/server-sdk/test/codec.test.ts new file mode 100644 index 000000000..8d2eeeb1a --- /dev/null +++ b/packages/server-sdk/test/codec.test.ts @@ -0,0 +1,71 @@ +import type { WebSocketEvent } from '@proj-airi/server-shared/types' + +import { stringify as stringifySuperJson } from 'superjson' +import { describe, expect, it } from 'vitest' + +import { InvalidMessageError, parseEvent, stringifyEvent } from '../src/codec' + +describe('server-sdk codec', () => { + it('parses SuperJSON and plain JSON protocol events', () => { + const event: WebSocketEvent = { + type: 'module:authenticate', + data: { token: 'secret' }, + metadata: { + source: { kind: 'plugin', plugin: { id: 'plugin-1' }, id: 'plugin-1' }, + event: { id: 'event-1' }, + }, + } + + expect(parseEvent(stringifySuperJson(event))).toEqual(event) + expect(parseEvent(JSON.stringify(event))).toEqual(event) + }) + + it('throws debuggable errors for invalid messages', () => { + const source = { type: 'module:authenticate', data: 'secret' } + + try { + parseEvent(JSON.stringify(source)) + expect.unreachable('Expected invalid message parsing to throw.') + } + catch (error) { + expect(error).toBeInstanceOf(InvalidMessageError) + expect(error).toMatchObject({ source }) + expect((error as InvalidMessageError).cause).toEqual(expect.anything()) + } + }) + + it('rejects array event data with validation context', () => { + const source = { type: 'module:authenticate', data: ['secret'] } + + expect(() => parseEvent(JSON.stringify(source))).toThrow(InvalidMessageError) + + try { + parseEvent(JSON.stringify(source)) + } + catch (error) { + expect(error).toMatchObject({ source }) + expect((error as InvalidMessageError).cause).toEqual(expect.anything()) + } + }) + + it('wraps malformed event text with the original source', () => { + const source = '{not-json' + + try { + parseEvent(source) + expect.unreachable('Expected malformed message parsing to throw.') + } + catch (error) { + expect(error).toBeInstanceOf(InvalidMessageError) + expect(error).toMatchObject({ source }) + expect((error as InvalidMessageError).cause).toBeInstanceOf(Error) + } + }) + + it('stringifies protocol events with SuperJSON', () => { + expect(stringifyEvent({ + type: 'module:authenticate', + data: { token: 'secret' }, + })).toContain('module:authenticate') + }) +}) diff --git a/packages/server-sdk/test/extension-peer.test.ts b/packages/server-sdk/test/extension-peer.test.ts index 3d76af1aa..c10959580 100644 --- a/packages/server-sdk/test/extension-peer.test.ts +++ b/packages/server-sdk/test/extension-peer.test.ts @@ -1,16 +1,19 @@ -import type { WebSocketEventOptionalSource } from '@proj-airi/server-shared/types' +import type { ClientConnection, ClientConnector, ClientEvents } from '@proj-airi/better-ws' +import type { WebSocketBaseEvent, WebSocketEvent, WebSocketEventOptionalSource, WebSocketEvents } from '@proj-airi/server-shared/types' import type { ExtensionPeerClient } from '../src/extension-peer' -import type { WebSocketLike } from '../src/websocket-like' import { describe, expect, it, vi } from 'vitest' import { createWebSocketExtensionPeer } from '../src/extension-peer' +type Listener = (data: WebSocketBaseEvent) => void | Promise + class FakeClient implements ExtensionPeerClient { readonly sent: WebSocketEventOptionalSource[] = [] readonly connect = vi.fn(async () => {}) readonly close = vi.fn(() => {}) + readonly listeners = new Map>() send(data: WebSocketEventOptionalSource): boolean { this.sent.push(data) @@ -20,43 +23,49 @@ class FakeClient implements ExtensionPeerClient { sendOrThrow(data: WebSocketEventOptionalSource): void { this.sent.push(data) } + + onEvent( + event: E, + callback: (data: WebSocketBaseEvent) => void | Promise, + ) { + let listeners = this.listeners.get(event) + if (!listeners) { + listeners = new Set() + this.listeners.set(event, listeners) + } + + const listener = callback as Listener + listeners.add(listener) + + return () => { + listeners?.delete(listener) + } + } } -class FakeSocket implements WebSocketLike { - static readonly CONNECTING = 0 - static readonly OPEN = 1 - static readonly CLOSING = 2 - static readonly CLOSED = 3 +class FakeConnector implements ClientConnector { + readonly attempts: Array<{ + events: ClientEvents + connection: ClientConnection + }> = [] - onopen?: () => void - onclose?: () => void - onmessage?: (event: { data: string }) => void - onerror?: (event: unknown) => void - readyState = FakeSocket.CONNECTING - readonly sent: string[] = [] + connect(events: ClientEvents) { + const connection: ClientConnection = { + send: () => true, + close: () => events.close({ code: 1000, reason: 'closed', wasClean: true }), + } - constructor(readonly url: string) {} - - open() { - this.readyState = FakeSocket.OPEN - this.onopen?.() + this.attempts.push({ events, connection }) + return connection } +} - close(_code?: number, _reason?: string) { - this.readyState = FakeSocket.CLOSED - this.onclose?.() - } - - send(data: string | ArrayBufferLike | ArrayBufferView) { - this.sent.push(typeof data === 'string' ? data : new TextDecoder().decode(data)) - } +async function flushMicrotasks() { + await Promise.resolve() + await Promise.resolve() } describe('websocket extension peer', () => { - /** - * @example - * expect(fakeClient.sent.map(event => event.type)).toEqual(['peer:authenticate', 'extension:announce']) - */ it('authenticates the websocket peer separately from the extension session', async () => { const fakeClient = new FakeClient() const peer = createWebSocketExtensionPeer({ @@ -96,10 +105,6 @@ describe('websocket extension peer', () => { }) }) - /** - * @example - * expect(fakeClient.sent[0].type).toBe('extension:module:announce') - */ it('announces extension modules under the owning extension identity', () => { const fakeClient = new FakeClient() const peer = createWebSocketExtensionPeer({ @@ -132,36 +137,26 @@ describe('websocket extension peer', () => { }) }) - /** - * @example - * expect(sockets).toHaveLength(1) - */ - it('does not reconnect by default because manual extension handshakes are one-shot', async () => { - const sockets: FakeSocket[] = [] + it('creates a manual peer client without auto-connect or auto-reconnect by default', async () => { + const connector = new FakeConnector() const peer = createWebSocketExtensionPeer({ extension: { id: 'airi-extension-chess', sessionId: 'session-1', }, clientOptions: { - websocketConstructor: class extends FakeSocket { - constructor(url: string) { - super(url) - sockets.push(this) - } - }, - connectTimeoutMs: 10, + connector, }, }) - const connectPromise = peer.connect() - expect(sockets).toHaveLength(1) - sockets[0]!.open() - await connectPromise + expect(connector.attempts).toHaveLength(0) - sockets[0]!.close() - await new Promise(resolve => setTimeout(resolve, 0)) + await peer.connect() + expect(connector.attempts).toHaveLength(1) - expect(sockets).toHaveLength(1) + connector.attempts[0]!.connection.close() + await flushMicrotasks() + + expect(connector.attempts).toHaveLength(1) }) }) diff --git a/packages/stage-ui/src/stores/mods/api/channel-server.ts b/packages/stage-ui/src/stores/mods/api/channel-server.ts index 24e6a6ddb..2bafa8c5f 100644 --- a/packages/stage-ui/src/stores/mods/api/channel-server.ts +++ b/packages/stage-ui/src/stores/mods/api/channel-server.ts @@ -1,16 +1,16 @@ import type { + ClientConnector, ContextUpdate, InputContextUpdate, WebSocketBaseEvent, WebSocketEvent, WebSocketEventOptionalSource, WebSocketEvents, - WebSocketLikeConstructor, } from '@proj-airi/server-sdk' import type { CommonContentPart } from '@xsai/shared-chat' import { errorMessageFrom } from '@moeru/std' -import { Client, WebSocketEventSource } from '@proj-airi/server-sdk' +import { Client, createTextProtocolConnector, WebSocketEventSource } from '@proj-airi/server-sdk' import { isStageTamagotchi, isStageWeb } from '@proj-airi/stage-shared' import { useLocalStorage } from '@vueuse/core' import { nanoid } from 'nanoid' @@ -25,6 +25,8 @@ interface ChannelListenerEntry { boundClient?: Client } +type TextConnectorFactory = (url: string) => ClientConnector | undefined + function hasReconnectableWebSocketScheme(url: string | undefined) { if (!url) { return false @@ -51,7 +53,7 @@ export const useModsServerChannelStore = defineStore('mods:channels:proj-airi:se const connected = ref(false) const client = ref() const initializing = ref | null>(null) - const websocketConstructor = ref() + const textConnectorFactory = ref() const hasEverConnected = ref(false) const pendingSend = ref>([]) const pendingSendCount = computed(() => pendingSend.value.length) @@ -90,15 +92,15 @@ export const useModsServerChannelStore = defineStore('mods:channels:proj-airi:se async function initialize(options?: { token?: string possibleEvents?: Array - websocketConstructor?: WebSocketLikeConstructor + connector?: TextConnectorFactory }) { if (connected.value && client.value) return Promise.resolve() if (initializing.value) return initializing.value - if (options?.websocketConstructor) { - websocketConstructor.value = options.websocketConstructor + if (options?.connector) { + textConnectorFactory.value = options.connector } const possibleEvents = Array.from(new Set([ @@ -107,11 +109,16 @@ export const useModsServerChannelStore = defineStore('mods:channels:proj-airi:se ])) initializing.value = new Promise((resolve) => { + const currentWebSocketUrl = websocketUrl.value || defaultWebSocketUrl + const textConnector = textConnectorFactory.value?.(currentWebSocketUrl) + client.value = new Client({ name: isStageWeb() ? WebSocketEventSource.StageWeb : isStageTamagotchi() ? WebSocketEventSource.StageTamagotchi : WebSocketEventSource.StageWeb, - url: websocketUrl.value || defaultWebSocketUrl, + url: currentWebSocketUrl, token: options?.token ?? (websocketAuthToken.value || undefined), - websocketConstructor: websocketConstructor.value, + connector: textConnector + ? createTextProtocolConnector(textConnector) + : undefined, heartbeat: { // Keep client and server heartbeat windows aligned to reduce false-positive disconnects. readTimeout: 60_000, diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 869295102..ba5321493 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -3307,6 +3307,21 @@ importers: specifier: 'catalog:' version: 0.1.3 + packages/better-ws: + dependencies: + '@moeru/eventa': + specifier: 'catalog:' + version: 1.0.0-beta.8(electron@41.2.1)(h3@2.0.1-rc.20(crossws@0.4.5(srvx@0.11.15)))(hono@4.12.2) + crossws: + specifier: 'catalog:' + version: 0.4.5(srvx@0.11.15) + h3: + specifier: 'catalog:' + version: 2.0.1-rc.20(crossws@0.4.5(srvx@0.11.15)) + srvx: + specifier: 'catalog:' + version: 0.11.15 + packages/cap-vite: dependencies: '@capacitor/cli': @@ -3696,6 +3711,9 @@ importers: '@moeru/std': specifier: 'catalog:' version: 0.1.0-beta.17 + '@proj-airi/better-ws': + specifier: workspace:^ + version: link:../better-ws '@proj-airi/server-shared': specifier: workspace:^ version: link:../server-shared @@ -3714,6 +3732,9 @@ importers: superjson: specifier: 'catalog:' version: 2.2.6 + valibot: + specifier: 'catalog:' + version: 1.3.1(typescript@5.9.3) packages/server-schema: devDependencies: @@ -3726,15 +3747,18 @@ importers: '@moeru/std': specifier: 'catalog:' version: 0.1.0-beta.17 + '@proj-airi/better-ws': + specifier: workspace:^ + version: link:../better-ws '@proj-airi/server-shared': specifier: workspace:^ version: link:../server-shared - crossws: - specifier: 'catalog:' - version: 0.4.5(srvx@0.11.15) superjson: specifier: 'catalog:' version: 2.2.6 + valibot: + specifier: 'catalog:' + version: 1.3.1(typescript@5.9.3) packages/server-sdk-shared: dependencies: diff --git a/vitest.config.ts b/vitest.config.ts index c0bef9272..d7d860044 100644 --- a/vitest.config.ts +++ b/vitest.config.ts @@ -10,6 +10,7 @@ export default defineConfig({ 'packages/cap-vite', 'packages/core-agent', 'packages/vishot-runner-browser', + 'packages/better-ws', 'packages/plugin-sdk', 'packages/plugin-sdk-tamagotchi', 'packages/server-runtime',