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:
parent
c2c1d280d5
commit
7bc21df07c
|
|
@ -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() }) })
|
||||
},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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() })
|
||||
}
|
||||
|
|
@ -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()
|
||||
}
|
||||
})
|
||||
})
|
||||
Loading…
Reference in New Issue