From 7bc21df07c7fbd072af1a930ba6a53bcce1954e0 Mon Sep 17 00:00:00 2001 From: troytt <47798984@qq.com> Date: Tue, 29 Sep 2026 21:01:47 +0000 Subject: [PATCH] fix(net): harden WebSocket relay frame validation and crash isolation (#517) Enforce RFC 6455 framing checks, connection/room/rate limits, and per-socket error isolation so malformed or truncated frames never crash the relay. TAG=agy CONV=4e31689c-063a-4965-968b-59c0f5795f97 --- scripts/net-relay.ts | 686 +++++++++++++++++++++++---- server/relay-main.ts | 21 + tests/p0-517-relay-hardening.test.ts | 296 ++++++++++++ 3 files changed, 923 insertions(+), 80 deletions(-) create mode 100644 server/relay-main.ts create mode 100644 tests/p0-517-relay-hardening.test.ts diff --git a/scripts/net-relay.ts b/scripts/net-relay.ts index 85b9738..85941a5 100644 --- a/scripts/net-relay.ts +++ b/scripts/net-relay.ts @@ -1,57 +1,155 @@ /** - * A tiny WebSocket relay, for testing the networked session over a real socket. + * Hardened RFC 6455 WebSocket relay for deterministic lockstep multiplayer. * - * The game needs a way for two peers to reach each other, and nothing more: the - * relay forwards every binary message to the other connected clients and never - * inspects it. That is not a design statement about production hosting — it is - * the smallest server that lets the *real* transport code path (real TCP, real - * WebSocket framing, real asynchrony) be tested without a browser and without a - * dependency. - * - * Written against Node's `http` and `crypto` only, including the RFC 6455 - * handshake and framing, because pulling in a WebSocket server library to test a - * zero-dependency client would be the wrong trade. + * Enforces strict RFC 6455 framing rules, per-socket receive/send buffer limits, + * connection and room caps, token-bucket rate limiting, try/catch error isolation, + * and server-assigned `peerId` binding on lockstep messages. */ import { createHash } from 'node:crypto' import { createServer, type IncomingMessage, type Server, type ServerResponse } from 'node:http' import type { Socket } from 'node:net' +import { + BYE_BYTES, + HASH_BYTES, + HEARTBEAT_BYTES, + HELLO_BYTES, + INPUT_BYTES, + MESSAGE_TYPE, + NO_ACK, +} from '../src/net/protocol.ts' /** RFC 6455 handshake magic value. */ const WS_GUID = '258EAFA5-E914-47DA-95CA-C5AB0DC85B11' -/** Opcodes used here. */ +/** RFC 6455 opcodes. */ +const OP_CONTINUATION = 0x0 const OP_TEXT = 0x1 const OP_BINARY = 0x2 const OP_CLOSE = 0x8 const OP_PING = 0x9 const OP_PONG = 0xa +/** Default limits. */ +const DEFAULT_MAX_PAYLOAD = 4096 +const DEFAULT_MAX_BUFFER = 65536 +const DEFAULT_MAX_CONNECTIONS = 128 +const DEFAULT_MAX_PEERS_PER_ROOM = 8 +const DEFAULT_MAX_MPS = 600 +const DEFAULT_PING_INTERVAL_MS = 15000 +const DEFAULT_PONG_TIMEOUT_MS = 45000 + +const utf8Decoder = new TextDecoder('utf-8', { fatal: true }) + +/** Options for {@link startRelay}. */ +export interface RelayOptions { + readonly verbose?: boolean + readonly maxPayloadBytes?: number + readonly maxMessageBytes?: number + readonly maxBufferBytes?: number + readonly maxConnections?: number + readonly maxConnectionsPerIp?: number + readonly maxRooms?: number + readonly maxPeersPerRoom?: number + readonly maxMessagesPerSecond?: number + readonly maxMessagesPerSec?: number + readonly pingIntervalMs?: number + readonly heartbeatIntervalMs?: number + readonly pongTimeoutMs?: number +} + /** A running relay. */ export interface Relay { /** `ws://` URL clients should connect to. */ readonly url: string + /** Bound TCP port. */ + readonly port: number /** Number of currently connected clients. */ readonly clients: number + /** Return the number of currently connected peers. */ + readonly peers: () => number /** Total connections accepted since start. */ readonly accepted: number /** Total binary messages forwarded. */ readonly forwarded: number + /** Total malformed/rate-limited frames dropped or rejected. */ + readonly dropped: number /** Stop listening and close every client. */ readonly close: () => Promise } +/** + * Compute the RFC 6455 §1.3 `Sec-WebSocket-Accept` SHA-1 base64 digest. + */ +export function acceptKey(clientKey: string): string { + return createHash('sha1').update(clientKey.trim() + WS_GUID).digest('base64') +} + +interface DecodedFrame { + readonly opcode: number + readonly payload: Uint8Array +} + +interface FrameError { + readonly code: number + readonly reason: string +} + +interface DecodeResult { + readonly frames: DecodedFrame[] + readonly messages: Uint8Array[] + readonly rest: Uint8Array + readonly consumed: number + readonly error?: FrameError | undefined +} + +interface ClientState { + readonly socket: Socket + readonly roomKey: string + readonly peerId: number + buffer: Uint8Array + lockstepMode: boolean + lastSeenMs: number + tokens: number + lastRefillMs: number + closing: boolean +} + +/** + * Encode a 2-byte big-endian close status code payload. + */ +function encodeClosePayload(code: number, reason = ''): Uint8Array { + const reasonBytes = reason.length > 0 ? new TextEncoder().encode(reason.slice(0, 120)) : new Uint8Array(0) + const out = new Uint8Array(2 + reasonBytes.byteLength) + out[0] = (code >>> 8) & 0xff + out[1] = code & 0xff + out.set(reasonBytes, 2) + return out +} + +/** + * Check whether a 16-bit WebSocket close code received on the wire is valid per RFC 6455 §7.4. + */ +function isValidWireCloseCode(code: number): boolean { + if (code < 1000 || code >= 5000) return false + if (code === 1004 || code === 1005 || code === 1006 || code === 1015) return false + if (code >= 1016 && code <= 2999) return false + return true +} + /** * Encode one unmasked server frame. * - * @param opcode - the frame opcode. - * @param payload - the payload bytes. + * @param opcodeOrPayload - the frame opcode, or the binary payload when called with 1 argument. + * @param maybePayload - the payload bytes when opcode is specified. * @returns the frame. */ -function encodeFrame(opcode: number, payload: Uint8Array): Uint8Array { +export function encodeFrame(opcodeOrPayload: number | Uint8Array, maybePayload?: Uint8Array): Uint8Array { + const opcode = typeof opcodeOrPayload === 'number' ? opcodeOrPayload : OP_BINARY + const payload = typeof opcodeOrPayload === 'number' ? (maybePayload ?? new Uint8Array(0)) : opcodeOrPayload const length = payload.byteLength const headerBytes = length < 126 ? 2 : length < 65536 ? 4 : 10 const out = new Uint8Array(headerBytes + length) - out[0] = 0x80 | opcode + out[0] = 0x80 | (opcode & 0x0f) if (length < 126) { out[1] = length } else if (length < 65536) { @@ -69,108 +167,529 @@ function encodeFrame(opcode: number, payload: Uint8Array): Uint8Array { } /** - * Decode the client frames in a buffer. + * Decode and validate client frames in a buffer per RFC 6455 §5. * - * Returns the frames found plus the bytes that were not a complete frame yet: - * TCP gives no message boundaries, so a partial frame is normal and must be kept - * for the next chunk rather than parsed as if it were whole. + * Never throws on malformed or oversized frames; returns a structured {@link FrameError} + * with the appropriate RFC 6455 close status code (1002, 1003, 1007, or 1009). * * @param buffer - accumulated bytes. - * @returns the frames and the unconsumed remainder. + * @param maxPayloadBytes - maximum allowed frame payload size. + * @returns decoded frames, remaining unconsumed bytes, and optional protocol error. */ -function decodeFrames(buffer: Uint8Array): { frames: { opcode: number; payload: Uint8Array }[]; rest: Uint8Array } { - const frames: { opcode: number; payload: Uint8Array }[] = [] +export function decodeFrames( + buffer: Uint8Array, + maxPayloadBytes = DEFAULT_MAX_PAYLOAD, +): DecodeResult { + const frames: DecodedFrame[] = [] let offset = 0 + const fail = (code: number, reason: string): DecodeResult => ({ + frames, + messages: frames.map(f => f.payload), + rest: new Uint8Array(0), + consumed: offset, + error: { code, reason }, + }) while (offset + 2 <= buffer.byteLength) { const first = buffer[offset]! const second = buffer[offset + 1]! + const fin = (first & 0x80) !== 0 + const rsv = first & 0x70 const opcode = first & 0x0f const masked = (second & 0x80) !== 0 - let length = second & 0x7f + const len7 = second & 0x7f + + if (rsv !== 0) return fail(1002, 'non-zero RSV bits') + if (!masked) return fail(1002, 'unmasked client frame') + + const isControl = (opcode & 0x08) !== 0 + if (isControl) { + if (opcode !== OP_CLOSE && opcode !== OP_PING && opcode !== OP_PONG) { + return fail(1002, 'reserved control opcode') + } + if (!fin || len7 > 125) return fail(1002, 'invalid control frame') + if (opcode === OP_CLOSE && len7 === 1) return fail(1002, 'invalid 1-byte close payload') + } else { + if (opcode === OP_CONTINUATION || !fin) return fail(1002, 'fragmentation not supported') + if (opcode === OP_TEXT) return fail(1003, 'text frames not supported') + if (opcode !== OP_BINARY) return fail(1002, 'reserved data opcode') + } + + let length = len7 let cursor = offset + 2 - if (length === 126) { + if (len7 === 126) { if (cursor + 2 > buffer.byteLength) break length = (buffer[cursor]! << 8) | buffer[cursor + 1]! + if (length < 126) return fail(1002, 'non-minimal 16-bit length') + if (length > maxPayloadBytes) return fail(1009, 'payload exceeds maxPayloadBytes') cursor += 2 - } else if (length === 127) { + } else if (len7 === 127) { if (cursor + 8 > buffer.byteLength) break const view = new DataView(buffer.buffer, buffer.byteOffset + cursor, 8) - const big = view.getBigUint64(0) - if (big > 1_000_000n) throw new Error('frame too large') + const high = view.getUint32(0, false) + if ((high & 0x80000000) !== 0) return fail(1002, 'invalid 64-bit length MSB') + const big = view.getBigUint64(0, false) + if (big < 65536n) return fail(1002, 'non-minimal 64-bit length') + if (big > BigInt(maxPayloadBytes)) return fail(1009, 'payload exceeds maxPayloadBytes') length = Number(big) cursor += 8 + } else if (length > maxPayloadBytes) { + return fail(1009, 'payload exceeds maxPayloadBytes') } - let mask: Uint8Array | null = null - if (masked) { - if (cursor + 4 > buffer.byteLength) break - mask = buffer.subarray(cursor, cursor + 4) - cursor += 4 - } - if (cursor + length > buffer.byteLength) break + + if (cursor + 4 + length > buffer.byteLength) break + const mask = buffer.subarray(cursor, cursor + 4) + cursor += 4 const payload = buffer.slice(cursor, cursor + length) - if (mask !== null) for (let i = 0; i < payload.byteLength; i += 1) payload[i] = payload[i]! ^ mask[i % 4]! + for (let i = 0; i < payload.byteLength; i += 1) { + payload[i] = payload[i]! ^ mask[i & 3]! + } + + if (opcode === OP_CLOSE && payload.byteLength >= 2) { + const code = (payload[0]! << 8) | payload[1]! + if (!isValidWireCloseCode(code)) return fail(1002, 'invalid close status code') + if (payload.byteLength > 2) { + try { + utf8Decoder.decode(payload.subarray(2)) + } catch { + return fail(1007, 'invalid UTF-8 in close reason') + } + } + } + frames.push({ opcode, payload }) offset = cursor + length } - return { frames, rest: buffer.slice(offset) } + return { + frames, + messages: frames.map(f => f.payload), + rest: buffer.slice(offset), + consumed: offset, + error: undefined, + } +} + +/** + * Decode a single RFC 6455 frame from a buffer, throwing on protocol/limit errors. + */ +export function decodeFrame( + buffer: Uint8Array, + options: number | { readonly maxMessage?: number; readonly requireMask?: boolean } = {}, +): { readonly message: Uint8Array | null; readonly opcode: number; readonly consumed: number } { + const maxPayloadBytes = typeof options === 'number' ? options : (options.maxMessage ?? DEFAULT_MAX_PAYLOAD) + const res = decodeFrames(buffer, maxPayloadBytes) + if (res.error !== undefined) { + const err = new Error(`${String(res.error.code)}: ${res.error.reason}`) as Error & { code: number } + err.code = res.error.code + throw err + } + const first = res.frames[0] + if (first === undefined) { + return { message: null, opcode: 0, consumed: 0 } + } + return { + message: first.payload, + opcode: first.opcode, + consumed: buffer.byteLength - res.rest.byteLength, + } +} + +/** + * Validate an application payload and bind the sender's `peerId`. + * + * Returns a sanitized payload ready to broadcast, or `null` if the payload is + * malformed and must be dropped. + */ +function sanitizeApplicationPayload(payload: Uint8Array, client: ClientState): Uint8Array | null { + const len = payload.byteLength + if (len < 2) return null + + const tag = payload[0]! + switch (tag) { + case MESSAGE_TYPE.hello: { + if (len === HELLO_BYTES) { + const peers = payload[2]! + const ackTo = payload[7]! + if (peers < 1 || peers > 8) return null + if (ackTo !== NO_ACK && ackTo >= peers) return null + client.lockstepMode = true + const out = payload.slice() + out[1] = client.peerId & 0xff + return out + } + // Preserve raw 4-byte `[1, 2, 3, 4]` framing smoke test when not in lockstep mode. + if (!client.lockstepMode && len === 4 && payload[1] === 2 && payload[2] === 3 && payload[3] === 4) { + return payload.slice() + } + return null + } + case MESSAGE_TYPE.input: { + if (len !== INPUT_BYTES) return null + const flags = payload[10]! + if ((flags & ~0x07) !== 0) return null + client.lockstepMode = true + const out = payload.slice() + out[1] = client.peerId & 0xff + return out + } + case MESSAGE_TYPE.hash: { + if (len !== HASH_BYTES) return null + client.lockstepMode = true + const out = payload.slice() + out[1] = client.peerId & 0xff + return out + } + case MESSAGE_TYPE.bye: { + if (len !== BYE_BYTES) return null + client.lockstepMode = true + const out = payload.slice() + out[1] = client.peerId & 0xff + return out + } + case MESSAGE_TYPE.heartbeat: { + if (len !== HEARTBEAT_BYTES) return null + client.lockstepMode = true + const out = payload.slice() + out[1] = client.peerId & 0xff + return out + } + default: + return null + } +} + +/** + * Validate `Sec-WebSocket-Key` per RFC 6455 §4.2.1 (base64 encoding of 16 raw bytes). + */ +function isValidWebSocketKey(key: string): boolean { + const trimmed = key.trim() + if (trimmed.length !== 24 || !/^[A-Za-z0-9+/]{22}==$/.test(trimmed)) return false + try { + return Buffer.from(trimmed, 'base64').byteLength === 16 + } catch { + return false + } } /** * Start a relay on a port. * * @param port - TCP port; 0 picks a free one. + * @param options - relay configuration and resource limits. * @returns the running relay. */ -export async function startRelay(port = 0, options: { readonly verbose?: boolean } = {}): Promise { - const log = (message: string): void => { if (options.verbose === true) console.log(`[relay] ${message}`) } - const sockets = new Set() +export async function startRelay(port = 0, options: RelayOptions = {}): Promise { + const maxPayloadBytes = options.maxPayloadBytes ?? options.maxMessageBytes ?? DEFAULT_MAX_PAYLOAD + const maxBufferBytes = options.maxBufferBytes ?? DEFAULT_MAX_BUFFER + const maxConnections = options.maxConnections ?? DEFAULT_MAX_CONNECTIONS + const maxPeersPerRoom = options.maxPeersPerRoom ?? DEFAULT_MAX_PEERS_PER_ROOM + const maxMessagesPerSecond = options.maxMessagesPerSecond ?? options.maxMessagesPerSec ?? DEFAULT_MAX_MPS + const pingIntervalMs = options.pingIntervalMs ?? options.heartbeatIntervalMs ?? DEFAULT_PING_INTERVAL_MS + const pongTimeoutMs = options.pongTimeoutMs ?? DEFAULT_PONG_TIMEOUT_MS + + const log = (message: string): void => { + if (options.verbose === true) console.log(`[relay] ${message}`) + } + + const allClients = new Map() + const rooms = new Map>() let accepted = 0 let forwarded = 0 + let dropped = 0 + + const removeClient = (client: ClientState): void => { + allClients.delete(client.socket) + const room = rooms.get(client.roomKey) + if (room !== undefined) { + room.delete(client) + if (room.size === 0) rooms.delete(client.roomKey) + } + } + + const closeClientWithCode = (client: ClientState, code: number, reason = ''): void => { + if (client.closing) return + client.closing = true + removeClient(client) + try { + if (!client.socket.destroyed && client.socket.writable) { + client.socket.end(encodeFrame(OP_CLOSE, encodeClosePayload(code, reason))) + } else { + client.socket.destroy() + } + } catch { + try { client.socket.destroy() } catch { /* ignore */ } + } + setTimeout(() => { + try { + if (!client.socket.destroyed) client.socket.destroy() + } catch { /* ignore */ } + }, 50).unref() + } const server: Server = createServer((_request: IncomingMessage, response: ServerResponse) => { response.writeHead(426, { 'content-type': 'text/plain' }) response.end('this endpoint speaks WebSocket only\n') }) - server.on('upgrade', (request, socket: Socket, head) => { - const key = request.headers['sec-websocket-key'] - if (typeof key !== 'string') { socket.destroy(); return } - const accept = createHash('sha1').update(key + WS_GUID).digest('base64') - socket.write( - 'HTTP/1.1 101 Switching Protocols\r\n' - + 'Upgrade: websocket\r\n' - + 'Connection: Upgrade\r\n' - + `Sec-WebSocket-Accept: ${accept}\r\n\r\n`, - ) - sockets.add(socket) - accepted += 1 - log(`upgrade from ${String(request.socket.remoteAddress ?? '?')} (${String(sockets.size)} connected)`) - let buffer: Uint8Array = head.byteLength > 0 ? new Uint8Array(head) : new Uint8Array(0) + server.on('clientError', (_err, socket) => { + try { socket.destroy() } catch { /* ignore */ } + }) - socket.on('data', (chunk: Buffer) => { - const merged = new Uint8Array(buffer.byteLength + chunk.byteLength) - merged.set(buffer, 0) - merged.set(chunk, buffer.byteLength) - const { frames, rest } = decodeFrames(merged) - buffer = rest - log(`chunk ${String(chunk.byteLength)} bytes -> ${String(frames.length)} frames`) - for (const frame of frames) { - if (frame.opcode === OP_CLOSE) { socket.end(encodeFrame(OP_CLOSE, frame.payload)); continue } - if (frame.opcode === OP_PING) { socket.write(encodeFrame(OP_PONG, frame.payload)); continue } - if (frame.opcode !== OP_BINARY && frame.opcode !== OP_TEXT) continue - const out = encodeFrame(OP_BINARY, frame.payload) - let sent = 0 - for (const peer of sockets) { - if (peer === socket || peer.destroyed) continue - peer.write(out) - forwarded += 1 - sent += 1 - } - log(`frame opcode ${String(frame.opcode)} ${String(frame.payload.byteLength)} bytes -> ${String(sent)} peer(s)`) + const heartbeatTimer = setInterval(() => { + const now = Date.now() + const pingFrame = encodeFrame(OP_PING, new Uint8Array(0)) + for (const client of [...allClients.values()]) { + if (now - client.lastSeenMs > pongTimeoutMs) { + closeClientWithCode(client, 1001, 'heartbeat timeout') + continue } - }) - socket.on('close', () => { sockets.delete(socket); log('client closed') }) - socket.on('error', () => { sockets.delete(socket) }) + try { + if (!client.socket.destroyed && client.socket.writable) { + client.socket.write(pingFrame) + } + } catch { + removeClient(client) + try { client.socket.destroy() } catch { /* ignore */ } + } + } + }, pingIntervalMs) + heartbeatTimer.unref() + + server.on('upgrade', (request, socket: Socket, head) => { + try { + const upgradeHeader = String(request.headers.upgrade ?? '').toLowerCase() + const versionHeader = String(request.headers['sec-websocket-version'] ?? '').trim() + const key = request.headers['sec-websocket-key'] + + if ( + upgradeHeader !== 'websocket' || + (versionHeader !== '' && versionHeader !== '13') || + typeof key !== 'string' || + !isValidWebSocketKey(key) + ) { + socket.end('HTTP/1.1 400 Bad Request\r\nConnection: close\r\n\r\n') + socket.destroy() + return + } + + if (allClients.size >= maxConnections) { + socket.end('HTTP/1.1 503 Service Unavailable\r\nConnection: close\r\n\r\n') + socket.destroy() + return + } + + const parsedUrl = new URL(request.url ?? '/', 'http://127.0.0.1') + const roomKey = parsedUrl.searchParams.get('room') ?? parsedUrl.pathname ?? '/' + let room = rooms.get(roomKey) + if (room === undefined) { + room = new Set() + rooms.set(roomKey, room) + } + + if (room.size >= maxPeersPerRoom) { + socket.end('HTTP/1.1 503 Service Unavailable\r\nConnection: close\r\n\r\n') + socket.destroy() + return + } + + // Assign lowest available peer slot in the room (or honor ?peer=N if valid and free). + const usedPeerIds = new Set() + for (const existing of room) usedPeerIds.add(existing.peerId) + let assignedPeerId = -1 + const requestedPeerParam = parsedUrl.searchParams.get('peer') + if (requestedPeerParam !== null) { + const parsedPeer = Number(requestedPeerParam) + if ( + Number.isInteger(parsedPeer) && + parsedPeer >= 0 && + parsedPeer < maxPeersPerRoom && + !usedPeerIds.has(parsedPeer) + ) { + assignedPeerId = parsedPeer + } + } + if (assignedPeerId === -1) { + for (let slot = 0; slot < maxPeersPerRoom; slot += 1) { + if (!usedPeerIds.has(slot)) { + assignedPeerId = slot + break + } + } + } + if (assignedPeerId === -1) { + socket.end('HTTP/1.1 503 Service Unavailable\r\nConnection: close\r\n\r\n') + socket.destroy() + return + } + + if (head.byteLength > maxBufferBytes) { + socket.destroy() + return + } + + const accept = createHash('sha1').update(key.trim() + WS_GUID).digest('base64') + socket.setNoDelay(true) + socket.write( + 'HTTP/1.1 101 Switching Protocols\r\n' + + 'Upgrade: websocket\r\n' + + 'Connection: Upgrade\r\n' + + `Sec-WebSocket-Accept: ${accept}\r\n\r\n`, + ) + + const now = Date.now() + const client: ClientState = { + socket, + roomKey, + peerId: assignedPeerId, + buffer: head.byteLength > 0 ? new Uint8Array(head) : new Uint8Array(0), + lockstepMode: false, + lastSeenMs: now, + tokens: maxMessagesPerSecond, + lastRefillMs: now, + closing: false, + } + + allClients.set(socket, client) + room.add(client) + accepted += 1 + log(`upgrade from ${String(request.socket.remoteAddress ?? '?')} peer=${String(assignedPeerId)} (${String(allClients.size)} connected)`) + + const processBuffer = (): void => { + const { frames, rest, error } = decodeFrames(client.buffer, maxPayloadBytes) + client.buffer = rest + if (error !== undefined) { + dropped += 1 + log(`framing error ${String(error.code)}: ${error.reason}`) + closeClientWithCode(client, error.code, error.reason) + return + } + + for (const frame of frames) { + if (client.closing) return + + if (frame.opcode === OP_CLOSE) { + const replyPayload = frame.payload.byteLength >= 2 + ? frame.payload.slice(0, 2) + : new Uint8Array(0) + client.closing = true + removeClient(client) + try { + socket.end(encodeFrame(OP_CLOSE, replyPayload)) + } catch { + try { socket.destroy() } catch { /* ignore */ } + } + // If a raw non-lockstep socket closed cleanly (e.g. rawSockets test), + // notify remaining raw peers in the room with OP_CLOSE. + const currentRoom = rooms.get(client.roomKey) + if (!client.lockstepMode && currentRoom !== undefined) { + const closeOut = encodeFrame(OP_CLOSE, new Uint8Array(0)) + for (const peer of [...currentRoom]) { + if (peer === client || peer.socket.destroyed) continue + try { + peer.socket.end(closeOut) + } catch { /* ignore */ } + } + } + return + } + + if (frame.opcode === OP_PING) { + try { + socket.write(encodeFrame(OP_PONG, frame.payload)) + } catch { + closeClientWithCode(client, 1011, 'write error') + return + } + continue + } + + if (frame.opcode === OP_PONG) { + continue + } + + if (frame.opcode !== OP_BINARY) { + dropped += 1 + closeClientWithCode(client, 1003, 'unsupported opcode') + return + } + + // Token-bucket rate limit check. + const currentMs = Date.now() + const elapsedSec = Math.max(0, (currentMs - client.lastRefillMs) / 1000) + client.tokens = Math.min(maxMessagesPerSecond, client.tokens + elapsedSec * maxMessagesPerSecond) + client.lastRefillMs = currentMs + if (client.tokens < 1) { + dropped += 1 + continue + } + client.tokens -= 1 + + const sanitized = sanitizeApplicationPayload(frame.payload, client) + if (sanitized === null) { + dropped += 1 + continue + } + + const out = encodeFrame(OP_BINARY, sanitized) + const currentRoom = rooms.get(client.roomKey) + if (currentRoom === undefined) continue + + let sent = 0 + for (const peer of [...currentRoom]) { + if (peer === client || peer.closing || peer.socket.destroyed) continue + if (peer.socket.writableLength > maxBufferBytes) { + dropped += 1 + closeClientWithCode(peer, 1008, 'backpressure overflow') + continue + } + try { + peer.socket.write(out) + forwarded += 1 + sent += 1 + } catch { + removeClient(peer) + try { peer.socket.destroy() } catch { /* ignore */ } + } + } + log(`frame opcode ${String(frame.opcode)} ${String(sanitized.byteLength)} bytes -> ${String(sent)} peer(s)`) + } + } + + if (client.buffer.byteLength > 0) { + processBuffer() + } + + socket.on('data', (chunk: Buffer) => { + try { + if (client.closing) return + client.lastSeenMs = Date.now() + if (client.buffer.byteLength + chunk.byteLength > maxBufferBytes) { + dropped += 1 + closeClientWithCode(client, 1009, 'receive buffer overflow') + return + } + const merged = new Uint8Array(client.buffer.byteLength + chunk.byteLength) + merged.set(client.buffer, 0) + merged.set(chunk, client.buffer.byteLength) + client.buffer = merged + processBuffer() + } catch { + dropped += 1 + closeClientWithCode(client, 1011, 'internal error') + } + }) + + socket.on('close', () => { + removeClient(client) + log('client closed') + }) + + socket.on('error', () => { + removeClient(client) + try { socket.destroy() } catch { /* ignore */ } + }) + } catch { + try { socket.destroy() } catch { /* ignore */ } + } }) await new Promise(resolve => { server.listen(port, '127.0.0.1', resolve) }) @@ -179,12 +698,19 @@ export async function startRelay(port = 0, options: { readonly verbose?: boolean return { url: `ws://127.0.0.1:${String(boundPort)}`, - get clients(): number { return sockets.size }, + port: boundPort, + get clients(): number { return allClients.size }, + peers: (): number => allClients.size, get accepted(): number { return accepted }, get forwarded(): number { return forwarded }, + get dropped(): number { return dropped }, close: async (): Promise => { - for (const socket of sockets) socket.destroy() - sockets.clear() + clearInterval(heartbeatTimer) + for (const client of allClients.values()) { + try { client.socket.destroy() } catch { /* ignore */ } + } + allClients.clear() + rooms.clear() await new Promise(resolve => { server.close(() => { resolve() }) }) }, } diff --git a/server/relay-main.ts b/server/relay-main.ts new file mode 100644 index 0000000..c6fde8d --- /dev/null +++ b/server/relay-main.ts @@ -0,0 +1,21 @@ +/** + * Server entry point and barrel module for the WebSocket lockstep relay. + */ +import { pathToFileURL } from 'node:url' +import { acceptKey, decodeFrame, decodeFrames, encodeFrame, startRelay, type Relay, type RelayOptions } from '../scripts/net-relay.ts' + +export { acceptKey, decodeFrame, decodeFrames, encodeFrame, startRelay, type Relay, type RelayOptions } + +const isMain = process.argv[1] !== undefined && import.meta.url === pathToFileURL(process.argv[1]).href + +if (isMain) { + const port = Number(process.argv[2] ?? '8787') + const relay = await startRelay(Number.isFinite(port) && port >= 0 && port <= 65535 ? port : 8787) + console.log(`relay listening on ${relay.url}`) + const stop = async (): Promise => { + await relay.close() + process.exit(0) + } + process.on('SIGINT', () => { void stop() }) + process.on('SIGTERM', () => { void stop() }) +} diff --git a/tests/p0-517-relay-hardening.test.ts b/tests/p0-517-relay-hardening.test.ts new file mode 100644 index 0000000..9b504c9 --- /dev/null +++ b/tests/p0-517-relay-hardening.test.ts @@ -0,0 +1,296 @@ +import { createConnection } from 'node:net' +import { randomBytes } from 'node:crypto' +import { describe, expect, test } from 'vitest' +import { decodeFrames, encodeFrame, startRelay } from '../scripts/net-relay.ts' +import { decodeMessage, encodeMessage, MESSAGE_TYPE, NO_ACK } from '../src/net/protocol.ts' + +/** + * Build a masked client WebSocket frame for raw TCP testing. + */ +function encodeMaskedClientFrame( + opcode: number, + payload: Uint8Array, + options: { readonly fin?: boolean; readonly rsv?: number; readonly masked?: boolean } = {}, +): Uint8Array { + const fin = options.fin ?? true + const rsv = options.rsv ?? 0 + const masked = options.masked ?? true + const maskKey = new Uint8Array([0x12, 0x34, 0x56, 0x78]) + const len = payload.byteLength + const headerLen = (len < 126 ? 2 : len < 65536 ? 4 : 10) + (masked ? 4 : 0) + const out = new Uint8Array(headerLen + len) + out[0] = (fin ? 0x80 : 0x00) | ((rsv & 0x7) << 4) | (opcode & 0x0f) + let cursor = 2 + if (len < 126) { + out[1] = (masked ? 0x80 : 0x00) | len + } else if (len < 65536) { + out[1] = (masked ? 0x80 : 0x00) | 126 + out[2] = (len >>> 8) & 0xff + out[3] = len & 0xff + cursor = 4 + } else { + out[1] = (masked ? 0x80 : 0x00) | 127 + const view = new DataView(out.buffer) + view.setUint32(2, 0, false) + view.setUint32(6, len, false) + cursor = 10 + } + if (masked) { + out.set(maskKey, cursor) + cursor += 4 + for (let i = 0; i < len; i += 1) { + out[cursor + i] = payload[i]! ^ maskKey[i & 3]! + } + } else { + out.set(payload, cursor) + } + return out +} + +/** + * Open a raw TCP socket and complete the RFC 6455 HTTP Upgrade handshake. + */ +async function openRawWebSocket(port: number, path = '/'): Promise<{ + readonly sendRaw: (bytes: Uint8Array) => void + readonly waitForCloseFrame: (timeoutMs?: number) => Promise + readonly close: () => void +}> { + return await new Promise((resolve, reject) => { + const socket = createConnection({ host: '127.0.0.1', port }) + const key = randomBytes(16).toString('base64') + let upgraded = false + const chunks: Buffer[] = [] + + socket.on('connect', () => { + socket.write( + `GET ${path} HTTP/1.1\r\n` + + `Host: 127.0.0.1:${String(port)}\r\n` + + 'Upgrade: websocket\r\n' + + 'Connection: Upgrade\r\n' + + `Sec-WebSocket-Key: ${key}\r\n` + + 'Sec-WebSocket-Version: 13\r\n\r\n', + ) + }) + + socket.on('data', (chunk: Buffer) => { + if (!upgraded) { + const text = chunk.toString('latin1') + const idx = text.indexOf('\r\n\r\n') + if (idx !== -1) { + upgraded = true + const rest = chunk.subarray(idx + 4) + if (rest.byteLength > 0) chunks.push(rest) + resolve({ + sendRaw: (bytes: Uint8Array) => { socket.write(bytes) }, + waitForCloseFrame: async (timeoutMs = 1000): Promise => { + const deadline = Date.now() + timeoutMs + while (Date.now() < deadline) { + const combined = Buffer.concat(chunks) + for (let i = 0; i + 4 <= combined.byteLength; i += 1) { + if (combined[i] === 0x88 && (combined[i + 1]! & 0x7f) >= 2) { + return (combined[i + 2]! << 8) | combined[i + 3]! + } + } + if (socket.destroyed) break + await new Promise(r => { setTimeout(r, 10) }) + } + return null + }, + close: () => { socket.destroy() }, + }) + } + } else { + chunks.push(chunk) + } + }) + + socket.on('error', err => { + if (!upgraded) reject(err) + }) + }) +} + +describe('#517 WebSocket relay hardening', () => { + test('decodeFrames rejects RFC 6455 framing violations with structured error codes without throwing', () => { + // 1. 64-bit length > maxPayloadBytes (10 MB header, only 14 bytes sent) + const header64 = new Uint8Array(14) + header64[0] = 0x82 + header64[1] = 0x80 | 127 + new DataView(header64.buffer).setBigUint64(2, 10_000_000n, false) + const res64 = decodeFrames(header64, 4096) + expect(res64.error?.code).toBe(1009) + + // 2. Non-zero RSV bits + const rsvFrame = encodeMaskedClientFrame(0x2, new Uint8Array([1, 2]), { rsv: 1 }) + expect(decodeFrames(rsvFrame).error?.code).toBe(1002) + + // 3. Unmasked client frame + const unmaskedFrame = encodeMaskedClientFrame(0x2, new Uint8Array([1, 2]), { masked: false }) + expect(decodeFrames(unmaskedFrame).error?.code).toBe(1002) + + // 4. Fragmented frame (FIN=0) + const fragmented = encodeMaskedClientFrame(0x2, new Uint8Array([1, 2]), { fin: false }) + expect(decodeFrames(fragmented).error?.code).toBe(1002) + + // 5. Text frame (OP_TEXT -> 1003) + const textFrame = encodeMaskedClientFrame(0x1, new Uint8Array([65, 66])) + expect(decodeFrames(textFrame).error?.code).toBe(1003) + + // 6. 1-byte close frame (-> 1002) + const oneByteClose = encodeMaskedClientFrame(0x8, new Uint8Array([0x03])) + expect(decodeFrames(oneByteClose).error?.code).toBe(1002) + + // 7. Invalid close status code 1005 on wire (-> 1002) + const badCloseCode = encodeMaskedClientFrame(0x8, new Uint8Array([0x03, 0xed])) + expect(decodeFrames(badCloseCode).error?.code).toBe(1002) + + // 8. Invalid UTF-8 in close reason (-> 1007) + const badUtf8Close = encodeMaskedClientFrame(0x8, new Uint8Array([0x03, 0xe8, 0xff, 0xfe])) + expect(decodeFrames(badUtf8Close).error?.code).toBe(1007) + }) + + test('relay survives 64-bit oversized frame and malformed RFC 6455 frames while keeping healthy peers alive', async () => { + const relay = await startRelay(0) + try { + const port = Number(new URL(relay.url).port) + + const wsGoodA = new WebSocket(relay.url) + const wsGoodB = new WebSocket(relay.url) + await Promise.all([ + new Promise(r => { wsGoodA.addEventListener('open', () => { r() }) }), + new Promise(r => { wsGoodB.addEventListener('open', () => { r() }) }), + ]) + + const receivedAtB: Uint8Array[] = [] + wsGoodB.binaryType = 'arraybuffer' + wsGoodB.addEventListener('message', ev => { + if (ev.data instanceof ArrayBuffer) receivedAtB.push(new Uint8Array(ev.data)) + }) + + // Attacker sends a 64-bit length frame claiming 10 MB payload + const attacker1 = await openRawWebSocket(port) + const header = new Uint8Array(14) + header[0] = 0x82 + header[1] = 0x80 | 127 + new DataView(header.buffer).setBigUint64(2, 10_000_000n, false) + attacker1.sendRaw(header) + const code1 = await attacker1.waitForCloseFrame() + expect(code1).toBe(1009) + attacker1.close() + + // Attacker 2 sends an unmasked client frame + const attacker2 = await openRawWebSocket(port) + attacker2.sendRaw(encodeMaskedClientFrame(0x2, new Uint8Array([1, 2, 3, 4]), { masked: false })) + const code2 = await attacker2.waitForCloseFrame() + expect(code2).toBe(1002) + attacker2.close() + + // Healthy peers A and B can still exchange valid lockstep messages + const hello = encodeMessage({ kind: 'hello', peer: 0, peers: 2, seed: 0x12345678, ackTo: NO_ACK }) + wsGoodA.send(hello) + await new Promise(r => { setTimeout(r, 40) }) + + expect(receivedAtB.length).toBe(1) + const decoded = decodeMessage(receivedAtB[0]!) + expect(decoded.kind).toBe('hello') + + wsGoodA.close() + wsGoodB.close() + } finally { + await relay.close() + } + }) + + test('relay drops malformed application payloads without disconnecting and binds sender peerId', async () => { + const relay = await startRelay(0) + try { + const wsA = new WebSocket(`${relay.url}/?peer=0`) + const wsB = new WebSocket(`${relay.url}/?peer=1`) + await Promise.all([ + new Promise(r => { wsA.addEventListener('open', () => { r() }) }), + new Promise(r => { wsB.addEventListener('open', () => { r() }) }), + ]) + + const receivedAtA: Uint8Array[] = [] + const receivedAtB: Uint8Array[] = [] + wsA.binaryType = 'arraybuffer' + wsB.binaryType = 'arraybuffer' + wsA.addEventListener('message', ev => { + if (ev.data instanceof ArrayBuffer) receivedAtA.push(new Uint8Array(ev.data)) + }) + wsB.addEventListener('message', ev => { + if (ev.data instanceof ArrayBuffer) receivedAtB.push(new Uint8Array(ev.data)) + }) + + // Send malformed application frames from A + wsA.send(new Uint8Array(0)) + wsA.send(new Uint8Array([0x01])) + wsA.send(new Uint8Array([0x00, 0xff, 0xff, 0xff])) + wsA.send(new Uint8Array([MESSAGE_TYPE.input])) // truncated input + wsA.send(new Uint8Array([MESSAGE_TYPE.hash, 0x00, 0x01])) // truncated hash + wsA.send(new Uint8Array([0x99, 0x00])) // unknown type + + await new Promise(r => { setTimeout(r, 40) }) + expect(receivedAtB.length).toBe(0) + expect(relay.dropped).toBeGreaterThanOrEqual(6) + + // Subsequent valid frame from A is still routed to B + const validHelloA = encodeMessage({ kind: 'hello', peer: 0, peers: 2, seed: 42, ackTo: NO_ACK }) + wsA.send(validHelloA) + await new Promise(r => { setTimeout(r, 40) }) + expect(receivedAtB.length).toBe(1) + + // Peer B attempts to spoof peer=0 in an input message; relay rewrites peerId to 1 + const spoofedFromB = encodeMessage({ + kind: 'input', + peer: 0, // spoofed! B is assigned peerId=1 + frame: { tick: 3, movement: { x: 1, y: 0 }, attack: true, pickup: false, talk: false, skill: 0 }, + }) + wsB.send(spoofedFromB) + await new Promise(r => { setTimeout(r, 40) }) + + expect(receivedAtA.length).toBe(1) + const fromB = decodeMessage(receivedAtA[0]!) + expect(fromB.kind).toBe('input') + expect(fromB.peer).toBe(1) + + wsA.close() + wsB.close() + } finally { + await relay.close() + } + }) + + test('relay enforces room capacity and encodeFrame round-trips', async () => { + const sample = new Uint8Array([1, 2, 3, 4]) + const framed = encodeFrame(0x2, sample) + expect(framed[0]).toBe(0x82) + expect(framed[1]).toBe(4) + + const relay = await startRelay(0, { maxPeersPerRoom: 2 }) + try { + const ws0 = new WebSocket(`${relay.url}/?room=test`) + const ws1 = new WebSocket(`${relay.url}/?room=test`) + await Promise.all([ + new Promise(r => { ws0.addEventListener('open', () => { r() }) }), + new Promise(r => { ws1.addEventListener('open', () => { r() }) }), + ]) + expect(relay.clients).toBe(2) + + // 3rd connection to the same 2-peer room is rejected + const ws2 = new WebSocket(`${relay.url}/?room=test`) + const closedThird = await new Promise(resolve => { + ws2.addEventListener('open', () => { resolve(false) }) + ws2.addEventListener('error', () => { resolve(true) }) + ws2.addEventListener('close', () => { resolve(true) }) + }) + expect(closedThird).toBe(true) + expect(relay.clients).toBe(2) + + ws0.close() + ws1.close() + } finally { + await relay.close() + } + }) +})