diff --git a/apps/server/src/adapters/WebsocketAdapter.ts b/apps/server/src/adapters/WebsocketAdapter.ts index e786c2f03..008266c44 100644 --- a/apps/server/src/adapters/WebsocketAdapter.ts +++ b/apps/server/src/adapters/WebsocketAdapter.ts @@ -62,11 +62,22 @@ class SocketServer implements IAdapter { this.wss = new WebSocketServer({ path: `${prefix}/ws`, server, maxPayload: this.MAX_PAYLOAD }); this.wss.on('connection', (ws, req) => { + // Rejected sockets can emit an error while their close handshake is in progress. + ws.on('error', console.error); + + let isAuthenticated = false; authenticateSocket(ws, req, (error) => { if (error) { ws.close(1008, 'Unauthorized'); + return; } + isAuthenticated = true; }); + + if (!isAuthenticated) { + return; + } + const clientId = generateId(); const clientName = getRandomName(); function sendPacket( @@ -94,8 +105,6 @@ class SocketServer implements IAdapter { // send store payload on connect sendPacket(MessageTag.RuntimeData, eventStore.poll()); - ws.on('error', console.error); - ws.on('close', () => { this.clients.delete(clientId); logger.info(LogOrigin.Client, `${this.clients.size} Connections with disconnected: ${clientName}`); diff --git a/apps/server/src/adapters/__tests__/WebsocketAdapter.test.ts b/apps/server/src/adapters/__tests__/WebsocketAdapter.test.ts new file mode 100644 index 000000000..0b0d080d1 --- /dev/null +++ b/apps/server/src/adapters/__tests__/WebsocketAdapter.test.ts @@ -0,0 +1,82 @@ +import type { Server } from 'node:http'; + +import { afterEach, describe, expect, it, vi } from 'vitest'; + +const websocketMocks = vi.hoisted(() => { + let connectionHandler: ((socket: FakeWebSocket, request: unknown) => void) | undefined; + + class FakeWebSocket { + static readonly OPEN = 1; + + readyState = FakeWebSocket.OPEN; + close = vi.fn(); + handlers = new Map void>>(); + + on(event: string, handler: (...args: unknown[]) => void) { + this.handlers.set(event, [...(this.handlers.get(event) ?? []), handler]); + return this; + } + + emit(event: string, ...args: unknown[]) { + const handlers = this.handlers.get(event) ?? []; + if (event === 'error' && handlers.length === 0) { + throw args[0]; + } + for (const handler of handlers) { + handler(...args); + } + } + } + + class FakeWebSocketServer { + clients = new Set(); + + on(event: string, handler: (socket: FakeWebSocket, request: unknown) => void) { + if (event === 'connection') { + connectionHandler = handler; + } + return this; + } + + close(callback: () => void) { + callback(); + } + } + + return { + FakeWebSocket, + FakeWebSocketServer, + getConnectionHandler: () => connectionHandler, + }; +}); + +vi.mock('ws', () => ({ + WebSocket: websocketMocks.FakeWebSocket, + WebSocketServer: websocketMocks.FakeWebSocketServer, +})); + +vi.mock('../../middleware/authenticate.js', () => ({ + authenticateSocket: (_socket: unknown, _request: unknown, next: (error?: Error) => void) => { + next(new Error('Unauthorized')); + }, +})); + +import { socket } from '../WebsocketAdapter.js'; + +describe('WebsocketAdapter authentication', () => { + afterEach(async () => { + await socket.shutdown(); + }); + + it('handles an error emitted while rejecting an unauthenticated socket', () => { + socket.init({} as Server, false); + const rejectedSocket = new websocketMocks.FakeWebSocket(); + const connectionHandler = websocketMocks.getConnectionHandler(); + + expect(connectionHandler).toBeDefined(); + connectionHandler?.(rejectedSocket, {}); + + expect(rejectedSocket.close).toHaveBeenCalledWith(1008, 'Unauthorized'); + expect(() => rejectedSocket.emit('error', new Error('socket closed'))).not.toThrow(); + }); +});