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
This commit is contained in:
troytt 2026-09-29 21:01:47 +00:00
parent c2c1d280d5
commit 7bc21df07c
3 changed files with 923 additions and 80 deletions

View File

@ -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<void>
}
/**
* 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<ArrayBufferLike>
readonly consumed: number
readonly error?: FrameError | undefined
}
interface ClientState {
readonly socket: Socket
readonly roomKey: string
readonly peerId: number
buffer: Uint8Array<ArrayBufferLike>
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<ArrayBufferLike>): { frames: { opcode: number; payload: Uint8Array }[]; rest: Uint8Array<ArrayBufferLike> } {
const frames: { opcode: number; payload: Uint8Array }[] = []
export function decodeFrames(
buffer: Uint8Array<ArrayBufferLike>,
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<ArrayBufferLike>,
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<Relay> {
const log = (message: string): void => { if (options.verbose === true) console.log(`[relay] ${message}`) }
const sockets = new Set<Socket>()
export async function startRelay(port = 0, options: RelayOptions = {}): Promise<Relay> {
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<Socket, ClientState>()
const rooms = new Map<string, Set<ClientState>>()
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<ArrayBufferLike> = 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<number>()
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<void>(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<void> => {
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<void>(resolve => { server.close(() => { resolve() }) })
},
}

21
server/relay-main.ts Normal file
View File

@ -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<void> => {
await relay.close()
process.exit(0)
}
process.on('SIGINT', () => { void stop() })
process.on('SIGTERM', () => { void stop() })
}

View File

@ -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<number | null>
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<number | null> => {
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<void>(r => { wsGoodA.addEventListener('open', () => { r() }) }),
new Promise<void>(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<void>(r => { wsA.addEventListener('open', () => { r() }) }),
new Promise<void>(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<void>(r => { ws0.addEventListener('open', () => { r() }) }),
new Promise<void>(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<boolean>(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()
}
})
})