From f33d5536323f73566c782ac31531928166b1c69b Mon Sep 17 00:00:00 2001 From: Nathaniel Nanle Date: Tue, 29 Sep 2026 10:46:50 +0100 Subject: [PATCH] feat(backend): implement payload sanitization and strict validation for WebSocket Relay Service - Add websocket-relay-validation: deep payload sanitization (prototype-pollution keys, control/bidi chars, depth/size limits), strict zod schemas for join:merchant / join:checkout / leave:checkout (UUID-only, unknown keys rejected), and a per-socket guard that reports errors via ack or relay:error, logs without raw payloads, and disconnects repeat offenders - Extract relay server setup from app.js into websocket-relay-server; cap maxHttpBufferSize at 16 KiB (configurable via WS_MAX_HTTP_BUFFER_SIZE) - Sanitize outbound checkout:presence and Horizon poller payment events - Deep-sanitize field values in sanitizeRelayMessage - Fix TypeError crash on join:* events sent without a payload - Add unit + real Socket.IO integration tests and docs Closes #1452 Co-Authored-By: Claude Opus 5.5 --- backend/.env.example | 4 + backend/docs/WEBSOCKET_RELAY_VALIDATION.md | 78 +++ backend/src/app.js | 67 +-- backend/src/lib/horizon-poller.js | 10 +- backend/src/lib/websocket-relay-security.js | 10 +- backend/src/lib/websocket-relay-server.js | 108 ++++ .../src/lib/websocket-relay-server.test.js | 235 +++++++++ backend/src/lib/websocket-relay-validation.js | 386 ++++++++++++++ .../lib/websocket-relay-validation.test.js | 470 ++++++++++++++++++ 9 files changed, 1305 insertions(+), 63 deletions(-) create mode 100644 backend/docs/WEBSOCKET_RELAY_VALIDATION.md create mode 100644 backend/src/lib/websocket-relay-server.js create mode 100644 backend/src/lib/websocket-relay-server.test.js create mode 100644 backend/src/lib/websocket-relay-validation.js create mode 100644 backend/src/lib/websocket-relay-validation.test.js diff --git a/backend/.env.example b/backend/.env.example index 6295ebb9..ee68433f 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -97,3 +97,7 @@ RESEND_API_KEY=your_resend_api_key # Sender address for receipt emails (must be a verified domain in Resend) # Defaults to: receipts@notifications.dripsnetwork.com RECEIPT_FROM_EMAIL=receipts@yourdomain.com + +# WebSocket relay limits (issue #1452) +WS_MAX_HTTP_BUFFER_SIZE=16384 +WS_MAX_INVALID_EVENTS=10 diff --git a/backend/docs/WEBSOCKET_RELAY_VALIDATION.md b/backend/docs/WEBSOCKET_RELAY_VALIDATION.md new file mode 100644 index 00000000..19d2a5e4 --- /dev/null +++ b/backend/docs/WEBSOCKET_RELAY_VALIDATION.md @@ -0,0 +1,78 @@ +# WebSocket Relay: Payload Sanitization & Strict Validation + +Issue: #1452 + +The Socket.IO relay (`src/lib/websocket-relay-server.js`) lets merchant dashboards +join `merchant:` rooms and checkout pages join `checkout:` rooms. Every +inbound event now passes through a guard (`src/lib/websocket-relay-validation.js`) +before any handler runs. + +## Inbound pipeline + +1. **Engine limit.** `maxHttpBufferSize` is 16 KiB (the Socket.IO default is 1 MB). Larger frames close the connection. +2. **Event allow-list.** Only `join:merchant`, `join:checkout` and `leave:checkout` are accepted. Any other event gets `UNKNOWN_EVENT`. +3. **Shape check.** The payload must be a plain JSON object. A missing payload, `null`, an array or a primitive gets `INVALID_PAYLOAD`. +4. **Byte limit.** The serialised payload must be 4 KiB or less (`PAYLOAD_TOO_LARGE`). +5. **Deep sanitization** (`sanitizeRelayPayload`): + - drops the `__proto__`, `constructor` and `prototype` keys at any depth + - NFC-normalises strings and strips C0/C1 control characters and bidi overrides (`\t \n \r` are kept) + - rejects non-finite numbers, BigInt, binary data and circular references + - bounds depth (8), keys per object (64), array length (256) and string length (2048) +6. **Strict schema** (zod `.strict()`). Unknown keys are rejected. IDs must be UUIDs and are trimmed and lower-cased, so room names are canonical and cannot be injected (for example `":admin"`). + +## Failure handling + +| Situation | Behaviour | +| --- | --- | +| Invalid event, client sent an ack callback | `ack({ ok: false, event, error: { code, message, issues? } })` | +| Invalid event, no ack callback | `socket.emit("relay:error", { ok: false, event, error })` | +| Every rejection | `logger.warn` records the socket id, event, error code and violation count. The raw payload is never logged or echoed. | +| `WS_MAX_INVALID_EVENTS` rejections on one socket | The socket is disconnected. | +| A handler throws | The error is logged, the client gets `INTERNAL_ERROR`, and the process keeps running. | + +Error codes: `UNKNOWN_EVENT`, `INVALID_PAYLOAD`, `PAYLOAD_TOO_LARGE`, +`VALIDATION_FAILED`, `PAYLOAD_TOO_DEEP`, `TOO_MANY_KEYS`, `ARRAY_TOO_LONG`, +`STRING_TOO_LONG`, `INVALID_NUMBER`, `INVALID_DATE`, `UNSUPPORTED_TYPE`, +`CIRCULAR_REFERENCE`, `UNSERIALIZABLE`, `INTERNAL_ERROR`. + +## Outbound sanitization + +`sanitizeOutboundPayload()` is applied to `checkout:presence` emits and to all +payment events from the Horizon poller (`notifyPaymentEvent`). It uses the same +rules with looser size limits. Data that came from Horizon or from merchant +metadata cannot push prototype keys, control characters or non-JSON values to +dashboards. If an outbound payload cannot be sanitized, it is dropped and a +warning is logged; nothing crashes. + +`sanitizeRelayMessage()` in `websocket-relay-security.js` now also deep-sanitizes +the value of each allowed field. + +## Configuration + +| Variable | Default | Purpose | +| --- | --- | --- | +| `WS_MAX_HTTP_BUFFER_SIZE` | `16384` | Maximum inbound frame size in bytes | +| `WS_MAX_INVALID_EVENTS` | `10` | Invalid events allowed per socket before it is disconnected | + +## Adding a new inbound event + +Add a strict zod schema to `INBOUND_EVENT_SCHEMAS`, then register the handler with +`relay.on(event, handler)`. `createRelayGuard().on()` refuses events that have no schema. + +## Security notes + +- **Fixed:** a `join:*` event with no payload used to throw a destructuring `TypeError` inside the listener. +- **Fixed:** any non-empty string was accepted as a room ID. Arbitrary room names could be created, and the relay's memory could grow without bound. +- **Fixed:** frames up to 1 MB were accepted for events that need fewer than 100 bytes. +- **Out of scope:** `join:merchant` is still unauthenticated. Anyone who knows a merchant UUID can listen to that merchant's room. The follow-up is to require a JWT at handshake (`verifyRelayToken`) and check that the token's merchant matches the room. + +## Tests + +``` +npx vitest run src/lib/websocket-relay-validation.test.js \ + src/lib/websocket-relay-server.test.js \ + src/lib/websocket-relay-security.test.js +``` + +`websocket-relay-server.test.js` starts a real Socket.IO server and drives it with +a raw WebSocket client that speaks the Engine.IO v4 / Socket.IO v5 wire protocol. diff --git a/backend/src/app.js b/backend/src/app.js index 503fc3cd..a5f3103f 100644 --- a/backend/src/app.js +++ b/backend/src/app.js @@ -1,7 +1,6 @@ import cors from "cors"; import helmet from "helmet"; import express from "express"; -import { Server as SocketIOServer } from "socket.io"; import swaggerUi from "swagger-ui-express"; import { ZodError } from "zod"; import path from "node:path"; @@ -49,6 +48,7 @@ import { versionDeprecationMiddleware } from "./lib/version-deprecation.js"; import oracleRouter from "./routes/oracle.js"; import { getPaymentSessionValidatorHealth } from "./lib/payment-session-validator.js"; import { configureExchangeRateCoordination } from "./services/exchangeRateService.js"; +import { createRelayServer } from "./lib/websocket-relay-server.js"; export async function createApp({ redisClient }) { const app = express(); @@ -62,64 +62,13 @@ export async function createApp({ redisClient }) { const __dirname = path.dirname(__filename); const publicDir = path.join(__dirname, "..", "public"); - // Create socket.io instance (attached to HTTP server in server.js) - const io = new SocketIOServer({ - cors: { - origin: process.env.CORS_ALLOWED_ORIGINS - ? process.env.CORS_ALLOWED_ORIGINS.split(",").map((o) => o.trim()) - : ["http://localhost:3000"], - credentials: true, - }, - }); - - const checkoutRoomName = (paymentId) => `checkout:${paymentId}`; - const emitCheckoutPresence = (paymentId) => { - const room = checkoutRoomName(paymentId); - const activeViewers = io.sockets.adapter.rooms.get(room)?.size ?? 0; - - io.to(room).emit("checkout:presence", { - payment_id: paymentId, - active_viewers: activeViewers, - }); - }; - - // Socket.io room management: clients join their merchant-specific room - io.on("connection", (socket) => { - const joinedCheckoutRooms = new Set(); - - socket.on("join:merchant", ({ merchant_id }) => { - if (typeof merchant_id === "string" && merchant_id.length > 0) { - socket.join(`merchant:${merchant_id}`); - } - }); - - socket.on("join:checkout", ({ payment_id }) => { - if (typeof payment_id !== "string" || payment_id.length === 0) { - return; - } - - const room = checkoutRoomName(payment_id); - joinedCheckoutRooms.add(payment_id); - socket.join(room); - emitCheckoutPresence(payment_id); - }); - - socket.on("leave:checkout", ({ payment_id }) => { - if (typeof payment_id !== "string" || payment_id.length === 0) { - return; - } - - joinedCheckoutRooms.delete(payment_id); - socket.leave(checkoutRoomName(payment_id)); - emitCheckoutPresence(payment_id); - }); - - socket.on("disconnect", () => { - for (const paymentId of joinedCheckoutRooms) { - emitCheckoutPresence(paymentId); - } - joinedCheckoutRooms.clear(); - }); + // Create socket.io relay (attached to HTTP server in server.js). Inbound + // events are sanitized and strictly validated before use (issue #1452). + const io = createRelayServer({ + corsOrigins: process.env.CORS_ALLOWED_ORIGINS + ? process.env.CORS_ALLOWED_ORIGINS.split(",").map((o) => o.trim()) + : ["http://localhost:3000"], + logger, }); // Make DB pool and io accessible on every request diff --git a/backend/src/lib/horizon-poller.js b/backend/src/lib/horizon-poller.js index e454f5f8..0fc26260 100644 --- a/backend/src/lib/horizon-poller.js +++ b/backend/src/lib/horizon-poller.js @@ -46,6 +46,7 @@ import { sendReceiptEmail } from "./email.js"; import { renderReceiptEmail } from "./email-templates.js"; import { getPayloadForVersion } from "../webhooks/resolver.js"; import { streamManager } from "./stream-manager.js"; +import { sanitizeOutboundPayload } from "./websocket-relay-validation.js"; import { connectRedisClient, invalidatePaymentCache } from "./redis.js"; import { logger } from "./logger.js"; import { @@ -790,7 +791,14 @@ function sleep(ms) { function notifyPaymentEvent(payment, { sseEvent, sseData, socketEvent, socketData }) { streamManager.notify(payment.id, sseEvent, sseData); if (_io && payment.merchant_id) { - _io.to(`merchant:${payment.merchant_id}`).emit(socketEvent, socketData); + let payload; + try { + payload = sanitizeOutboundPayload(socketData); + } catch (err) { + logger.warn({ err, paymentId: payment.id, socketEvent }, "Horizon poller: dropped unsafe socket payload"); + return; + } + _io.to(`merchant:${payment.merchant_id}`).emit(socketEvent, payload); } } diff --git a/backend/src/lib/websocket-relay-security.js b/backend/src/lib/websocket-relay-security.js index e0bd6406..399d8edb 100644 --- a/backend/src/lib/websocket-relay-security.js +++ b/backend/src/lib/websocket-relay-security.js @@ -10,6 +10,7 @@ */ import jwt from "jsonwebtoken"; +import { sanitizeRelayPayload } from "./websocket-relay-validation.js"; // ─── Allowed message fields ─────────────────────────────────────────────────── @@ -135,7 +136,8 @@ function verifyRelayToken(token, secret) { * * @param {any} msg - The parsed WebSocket message object * @returns {{ sanitized: object, warnings: string[] }} - * @throws {Error} When `msg` is not a non-null object, or when required fields are missing + * @throws {Error} When `msg` is not a non-null object, when required fields are missing, + * or (RelayValidationError) when a field value breaks the relay payload limits */ function sanitizeRelayMessage(msg) { if (msg === null || typeof msg !== "object" || Array.isArray(msg)) { @@ -145,10 +147,12 @@ function sanitizeRelayMessage(msg) { const warnings = []; const sanitized = {}; - // Copy only allowed fields + // Copy only allowed fields, deep-sanitizing each value so nested payloads + // cannot carry prototype-pollution keys or control characters (issue #1452) for (const [key, value] of Object.entries(msg)) { if (ALLOWED_MESSAGE_FIELDS.has(key)) { - sanitized[key] = value; + const clean = sanitizeRelayPayload(value); + if (clean !== undefined) sanitized[key] = clean; } else { warnings.push(`Unknown field stripped: '${key}'`); } diff --git a/backend/src/lib/websocket-relay-server.js b/backend/src/lib/websocket-relay-server.js new file mode 100644 index 00000000..cb24ad68 --- /dev/null +++ b/backend/src/lib/websocket-relay-server.js @@ -0,0 +1,108 @@ +/** + * websocket-relay-server.js + * + * Builds the Socket.IO relay server used for merchant dashboards and checkout + * presence. Every inbound event goes through the relay guard, which sanitizes + * and strictly validates the payload before any room is joined (issue #1452). + */ + +import { Server as SocketIOServer } from "socket.io"; +import { + createRelayGuard, + sanitizeOutboundPayload, + DEFAULT_MAX_VIOLATIONS, +} from "./websocket-relay-validation.js"; + +/** Inbound relay events are tiny room joins; the socket.io default is 1 MB. */ +export const DEFAULT_WS_MAX_HTTP_BUFFER_SIZE = 16 * 1024; + +function parsePositiveInt(value, fallback) { + const parsed = Number.parseInt(value, 10); + return Number.isFinite(parsed) && parsed > 0 ? parsed : fallback; +} + +export const checkoutRoomName = (paymentId) => `checkout:${paymentId}`; + +/** + * Register the relay's connection handlers on an existing Socket.IO server. + * + * @param {import("socket.io").Server} io + * @param {object} [opts] + * @param {{ warn: Function, error: Function }} [opts.logger] + * @param {number} [opts.maxViolations] - Invalid events tolerated per socket + */ +export function attachRelayHandlers(io, opts = {}) { + const emitCheckoutPresence = (paymentId) => { + const room = checkoutRoomName(paymentId); + const activeViewers = io.sockets.adapter.rooms.get(room)?.size ?? 0; + + io.to(room).emit( + "checkout:presence", + sanitizeOutboundPayload({ + payment_id: paymentId, + active_viewers: activeViewers, + }), + ); + }; + + io.on("connection", (socket) => { + const joinedCheckoutRooms = new Set(); + const relay = createRelayGuard(socket, { + logger: opts.logger, + maxViolations: opts.maxViolations, + }); + + relay.on("join:merchant", ({ merchant_id }) => { + socket.join(`merchant:${merchant_id}`); + }); + + relay.on("join:checkout", ({ payment_id }) => { + joinedCheckoutRooms.add(payment_id); + socket.join(checkoutRoomName(payment_id)); + emitCheckoutPresence(payment_id); + }); + + relay.on("leave:checkout", ({ payment_id }) => { + joinedCheckoutRooms.delete(payment_id); + socket.leave(checkoutRoomName(payment_id)); + emitCheckoutPresence(payment_id); + }); + + socket.on("disconnect", () => { + for (const paymentId of joinedCheckoutRooms) { + emitCheckoutPresence(paymentId); + } + joinedCheckoutRooms.clear(); + }); + }); + + return io; +} + +/** + * Create the relay Socket.IO server (attached to the HTTP server in server.js). + * + * Environment: + * WS_MAX_HTTP_BUFFER_SIZE - max inbound frame size in bytes (default 16 KiB) + * WS_MAX_INVALID_EVENTS - invalid events before a socket is disconnected (default 10) + * + * @param {object} opts + * @param {string[]} opts.corsOrigins + * @param {{ warn: Function, error: Function }} [opts.logger] + * @param {NodeJS.ProcessEnv} [opts.env] + * @returns {import("socket.io").Server} + */ +export function createRelayServer({ corsOrigins, logger, env = process.env }) { + const io = new SocketIOServer({ + cors: { origin: corsOrigins, credentials: true }, + maxHttpBufferSize: parsePositiveInt( + env.WS_MAX_HTTP_BUFFER_SIZE, + DEFAULT_WS_MAX_HTTP_BUFFER_SIZE, + ), + }); + + return attachRelayHandlers(io, { + logger, + maxViolations: parsePositiveInt(env.WS_MAX_INVALID_EVENTS, DEFAULT_MAX_VIOLATIONS), + }); +} diff --git a/backend/src/lib/websocket-relay-server.test.js b/backend/src/lib/websocket-relay-server.test.js new file mode 100644 index 00000000..eb9a41b2 --- /dev/null +++ b/backend/src/lib/websocket-relay-server.test.js @@ -0,0 +1,235 @@ +/** + * Integration tests for the relay server: a real Socket.IO server on an + * ephemeral port, driven by a raw WebSocket speaking the Engine.IO v4 / + * Socket.IO v5 wire protocol (so no socket.io-client dependency is needed). + */ +import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import http from "node:http"; +import WebSocket from "ws"; +import { createRelayServer, DEFAULT_WS_MAX_HTTP_BUFFER_SIZE } from "./websocket-relay-server.js"; +import { RELAY_ERROR_EVENT } from "./websocket-relay-validation.js"; + +const MERCHANT_ID = "3f1c2b4a-8d6e-4f7a-9b0c-1d2e3f4a5b6c"; +const PAYMENT_ID = "a1b2c3d4-e5f6-4a7b-8c9d-0e1f2a3b4c5d"; + +/** Minimal Socket.IO client over a raw WebSocket. */ +function connectClient(port) { + return new Promise((resolve, reject) => { + const ws = new WebSocket(`ws://127.0.0.1:${port}/socket.io/?EIO=4&transport=websocket`); + const events = []; + const acks = new Map(); + const waiters = []; + let nextAckId = 0; + let closed = false; + const closeWaiters = []; + + const flush = () => { + for (let i = waiters.length - 1; i >= 0; i--) { + const w = waiters[i]; + const found = events.find(w.match); + if (found) { + waiters.splice(i, 1); + w.resolve(found); + } + } + }; + + const client = { + events, + emit(event, ...args) { + ws.send(`42${JSON.stringify([event, ...args])}`); + }, + emitWithAck(event, ...args) { + const id = nextAckId++; + return new Promise((res) => { + acks.set(id, res); + ws.send(`42${id}${JSON.stringify([event, ...args])}`); + }); + }, + sendRaw(frame) { + ws.send(frame); + }, + waitFor(event, timeout = 2000) { + return new Promise((res, rej) => { + const w = { match: ([name]) => name === event, resolve: res }; + waiters.push(w); + flush(); + setTimeout(() => rej(new Error(`Timed out waiting for ${event}`)), timeout).unref(); + }); + }, + waitForClose(timeout = 2000) { + if (closed) return Promise.resolve(); + return new Promise((res, rej) => { + closeWaiters.push(res); + setTimeout(() => rej(new Error("Timed out waiting for close")), timeout).unref(); + }); + }, + get closed() { + return closed; + }, + close() { + ws.close(); + }, + }; + + ws.on("message", (raw) => { + const msg = raw.toString(); + if (msg.startsWith("0")) { + ws.send("40"); // Socket.IO CONNECT to the main namespace + } else if (msg === "2") { + ws.send("3"); // pong + } else if (msg.startsWith("40")) { + resolve(client); + } else if (msg.startsWith("42")) { + events.push(JSON.parse(msg.slice(2))); + flush(); + } else if (msg.startsWith("43")) { + const m = /^43(\d+)(.*)$/s.exec(msg); + acks.get(Number(m[1]))?.(JSON.parse(m[2])); + } else if (msg.startsWith("41")) { + closed = true; + closeWaiters.splice(0).forEach((r) => r()); + } + }); + ws.on("close", () => { + closed = true; + closeWaiters.splice(0).forEach((r) => r()); + }); + ws.on("error", reject); + }); +} + +describe("websocket relay server (integration)", () => { + let httpServer; + let io; + let port; + let logger; + const clients = []; + + const start = async (env = {}) => { + logger = { warn: vi.fn(), error: vi.fn() }; + httpServer = http.createServer(); + io = createRelayServer({ corsOrigins: ["http://localhost:3000"], logger, env }); + io.attach(httpServer); + await new Promise((r) => httpServer.listen(0, "127.0.0.1", r)); + port = httpServer.address().port; + }; + + const client = async () => { + const c = await connectClient(port); + clients.push(c); + return c; + }; + + const roomsOf = async () => { + const [socket] = await io.fetchSockets(); + return socket ? [...socket.rooms] : []; + }; + + beforeEach(async () => { + await start(); + }); + + afterEach(async () => { + clients.splice(0).forEach((c) => c.close()); + io.close(); + await new Promise((r) => httpServer.close(r)); + }); + + it("sets a tight maxHttpBufferSize by default and honours the env override", async () => { + expect(io.engine.opts.maxHttpBufferSize).toBe(DEFAULT_WS_MAX_HTTP_BUFFER_SIZE); + const custom = createRelayServer({ + corsOrigins: [], + logger, + env: { WS_MAX_HTTP_BUFFER_SIZE: "2048" }, + }); + custom.attach(http.createServer()); + expect(custom.engine.opts.maxHttpBufferSize).toBe(2048); + custom.close(); + }); + + it("joins the merchant room for a valid UUID (case-normalised)", async () => { + const c = await client(); + c.emit("join:merchant", { merchant_id: MERCHANT_ID.toUpperCase() }); + await vi.waitFor(async () => { + expect(await roomsOf()).toContain(`merchant:${MERCHANT_ID}`); + }); + + io.to(`merchant:${MERCHANT_ID}`).emit("payment:confirmed", { id: PAYMENT_ID }); + const [, payload] = await c.waitFor("payment:confirmed"); + expect(payload).toEqual({ id: PAYMENT_ID }); + }); + + it("emits checkout presence on join and leave", async () => { + const c = await client(); + c.emit("join:checkout", { payment_id: PAYMENT_ID }); + const [, joined] = await c.waitFor("checkout:presence"); + expect(joined).toEqual({ payment_id: PAYMENT_ID, active_viewers: 1 }); + + c.emit("leave:checkout", { payment_id: PAYMENT_ID }); + await vi.waitFor(async () => { + expect(await roomsOf()).not.toContain(`checkout:${PAYMENT_ID}`); + }); + }); + + it("does not crash the server on an event with no payload", async () => { + const c = await client(); + c.emit("join:merchant"); + const [, err] = await c.waitFor(RELAY_ERROR_EVENT); + expect(err).toMatchObject({ event: "join:merchant", error: { code: "INVALID_PAYLOAD" } }); + + // server still healthy + c.emit("join:checkout", { payment_id: PAYMENT_ID }); + await c.waitFor("checkout:presence"); + }); + + it("rejects room-name injection and does not join any room", async () => { + const c = await client(); + c.emit("join:merchant", { merchant_id: `${MERCHANT_ID}` + "x" }); + const [, err] = await c.waitFor(RELAY_ERROR_EVENT); + expect(err.error.code).toBe("VALIDATION_FAILED"); + expect((await roomsOf()).some((r) => r.startsWith("merchant:"))).toBe(false); + }); + + it("rejects unknown keys and unknown events", async () => { + const c = await client(); + const reply = await c.emitWithAck("join:checkout", { payment_id: PAYMENT_ID, extra: 1 }); + expect(reply).toEqual([ + expect.objectContaining({ ok: false, error: expect.objectContaining({ code: "VALIDATION_FAILED" }) }), + ]); + + c.emit("admin:broadcast", { msg: "hi" }); + const [, err] = await c.waitFor(RELAY_ERROR_EVENT); + expect(err).toMatchObject({ event: "admin:broadcast", error: { code: "UNKNOWN_EVENT" } }); + }); + + it("strips prototype-pollution keys sent over the wire", async () => { + const c = await client(); + c.sendRaw(`42["join:checkout",{"payment_id":"${PAYMENT_ID}","__proto__":{"polluted":true}}]`); + await c.waitFor("checkout:presence"); + expect({}.polluted).toBeUndefined(); + }); + + it("disconnects a socket after repeated invalid events", async () => { + await new Promise((r) => { + io.close(); + httpServer.close(r); + }); + await start({ WS_MAX_INVALID_EVENTS: "3" }); + + const c = await client(); + for (let i = 0; i < 3; i++) c.emit("join:checkout", { payment_id: "nope" }); + await c.waitForClose(); + expect(logger.warn).toHaveBeenCalledWith( + expect.objectContaining({ violations: 3 }), + expect.stringContaining("disconnecting"), + ); + }); + + it("drops the connection when a frame exceeds maxHttpBufferSize", async () => { + const c = await client(); + c.emit("join:checkout", { payment_id: PAYMENT_ID, pad: "x".repeat(DEFAULT_WS_MAX_HTTP_BUFFER_SIZE) }); + await c.waitForClose(); + expect(c.closed).toBe(true); + }); +}); diff --git a/backend/src/lib/websocket-relay-validation.js b/backend/src/lib/websocket-relay-validation.js new file mode 100644 index 00000000..272b22f1 --- /dev/null +++ b/backend/src/lib/websocket-relay-validation.js @@ -0,0 +1,386 @@ +/** + * websocket-relay-validation.js + * + * Payload sanitization and strict validation for the Socket.IO relay (issue #1452). + * + * - sanitizeRelayPayload(): deep, JSON-safe clone that strips prototype-pollution + * keys and control characters and enforces depth / size limits + * - validateInboundRelayEvent(): allow-listed inbound events, each checked + * against a strict zod schema (unknown keys rejected, IDs must be UUIDs) + * - sanitizeOutboundPayload(): same deep sanitization for server → client emits + * - createRelayGuard(): per-socket wrapper that validates every inbound event, + * reports errors to the client, logs them, and disconnects repeat offenders + */ + +import { z } from "zod"; + +// ─── Limits ─────────────────────────────────────────────────────────────────── + +export const RELAY_PAYLOAD_LIMITS = Object.freeze({ + maxBytes: 4096, + maxDepth: 8, + maxKeys: 64, + maxArrayLength: 256, + maxStringLength: 2048, +}); + +/** Invalid events tolerated per socket before it is disconnected. */ +export const DEFAULT_MAX_VIOLATIONS = 10; + +/** Event name used to report validation failures back to the client. */ +export const RELAY_ERROR_EVENT = "relay:error"; + +const FORBIDDEN_KEYS = new Set(["__proto__", "constructor", "prototype"]); + +// C0/C1 control characters (except \t \n \r) and Unicode bidi overrides, which +// can be used to spoof text rendered in dashboards and logs. +// eslint-disable-next-line no-control-regex +const UNSAFE_CHARS_RE = /[\u0000-\u0008\u000B\u000C\u000E-\u001F\u007F-\u009F‪-‮⁦-⁩]/g; + +// ─── Errors ─────────────────────────────────────────────────────────────────── + +export class RelayValidationError extends Error { + constructor(code, message, details) { + super(message); + this.name = "RelayValidationError"; + this.code = code; + if (details !== undefined) this.details = details; + } +} + +// ─── Deep sanitization ──────────────────────────────────────────────────────── + +/** + * Strip unsafe characters from a string and normalise it to NFC. + * + * @param {string} value + * @returns {string} + */ +export function sanitizeRelayString(value) { + return value.normalize("NFC").replace(UNSAFE_CHARS_RE, ""); +} + +/** + * Return a sanitized, JSON-safe deep copy of `value`. + * + * - Keys `__proto__`, `constructor` and `prototype` are dropped. + * - Strings are NFC-normalised and stripped of control / bidi characters. + * - `undefined`, functions and symbols are dropped (as JSON.stringify would). + * - Dates become ISO strings; non-finite numbers, BigInts, binary data and + * circular references are rejected. + * - Depth, key count, array length and string length are bounded. + * + * @param {any} value + * @param {Partial} [limits] + * @returns {any} + * @throws {RelayValidationError} + */ +export function sanitizeRelayPayload(value, limits = {}) { + const opts = { ...RELAY_PAYLOAD_LIMITS, ...limits }; + return sanitizeNode(value, opts, 0, new WeakSet(), "$"); +} + +function sanitizeNode(value, opts, depth, seen, path) { + if (value === null) return null; + + switch (typeof value) { + case "string": + if (value.length > opts.maxStringLength) { + throw new RelayValidationError( + "STRING_TOO_LONG", + `String at ${path} exceeds ${opts.maxStringLength} characters`, + ); + } + return sanitizeRelayString(value); + case "number": + if (!Number.isFinite(value)) { + throw new RelayValidationError("INVALID_NUMBER", `Non-finite number at ${path}`); + } + return value; + case "boolean": + return value; + case "undefined": + case "function": + case "symbol": + return undefined; + case "bigint": + throw new RelayValidationError("UNSUPPORTED_TYPE", `BigInt at ${path} is not allowed`); + default: + break; + } + + if (value instanceof Date) { + if (Number.isNaN(value.getTime())) { + throw new RelayValidationError("INVALID_DATE", `Invalid date at ${path}`); + } + return value.toISOString(); + } + + if (ArrayBuffer.isView(value) || value instanceof ArrayBuffer) { + throw new RelayValidationError("UNSUPPORTED_TYPE", `Binary data at ${path} is not allowed`); + } + + if (depth >= opts.maxDepth) { + throw new RelayValidationError( + "PAYLOAD_TOO_DEEP", + `Payload exceeds maximum nesting depth of ${opts.maxDepth}`, + ); + } + + if (seen.has(value)) { + throw new RelayValidationError("CIRCULAR_REFERENCE", `Circular reference at ${path}`); + } + seen.add(value); + + try { + if (Array.isArray(value)) { + if (value.length > opts.maxArrayLength) { + throw new RelayValidationError( + "ARRAY_TOO_LONG", + `Array at ${path} exceeds ${opts.maxArrayLength} items`, + ); + } + return value.map((item, i) => { + const clean = sanitizeNode(item, opts, depth + 1, seen, `${path}[${i}]`); + return clean === undefined ? null : clean; + }); + } + + const keys = Object.keys(value); + if (keys.length > opts.maxKeys) { + throw new RelayValidationError( + "TOO_MANY_KEYS", + `Object at ${path} exceeds ${opts.maxKeys} keys`, + ); + } + + const out = {}; + for (const rawKey of keys) { + const key = sanitizeRelayString(rawKey); + if (FORBIDDEN_KEYS.has(key) || key === "") continue; + const clean = sanitizeNode(value[rawKey], opts, depth + 1, seen, `${path}.${key}`); + if (clean !== undefined) out[key] = clean; + } + return out; + } finally { + seen.delete(value); + } +} + +/** + * Byte length of `value` once serialised for the wire. + * + * @param {any} value + * @returns {number} + * @throws {RelayValidationError} When the value cannot be serialised + */ +export function relayPayloadByteLength(value) { + try { + return Buffer.byteLength(JSON.stringify(value) ?? "", "utf8"); + } catch { + throw new RelayValidationError("UNSERIALIZABLE", "Payload cannot be serialised"); + } +} + +// ─── Inbound event schemas ──────────────────────────────────────────────────── + +const uuid = (field) => + z + .string({ + required_error: `${field} is required`, + invalid_type_error: `${field} must be a string`, + }) + .trim() + .uuid(`${field} must be a valid UUID`) + .transform((v) => v.toLowerCase()); + +/** + * Strict schemas for every inbound event the relay accepts. Any event not + * listed here is rejected, and unknown keys inside a payload are rejected + * rather than silently ignored. + */ +export const INBOUND_EVENT_SCHEMAS = Object.freeze({ + "join:merchant": z.object({ merchant_id: uuid("merchant_id") }).strict(), + "join:checkout": z.object({ payment_id: uuid("payment_id") }).strict(), + "leave:checkout": z.object({ payment_id: uuid("payment_id") }).strict(), +}); + +/** + * Sanitize and strictly validate an inbound relay event. + * + * @param {string} event - Socket.IO event name + * @param {any} payload - First argument sent with the event + * @param {object} [opts] + * @param {Record} [opts.schemas] + * @param {Partial} [opts.limits] + * @returns {{ ok: true, data: object } | { ok: false, error: { code: string, message: string, issues?: object[] } }} + */ +export function validateInboundRelayEvent(event, payload, opts = {}) { + const schemas = opts.schemas ?? INBOUND_EVENT_SCHEMAS; + const limits = { ...RELAY_PAYLOAD_LIMITS, ...opts.limits }; + + if (typeof event !== "string" || !Object.hasOwn(schemas, event)) { + return fail("UNKNOWN_EVENT", "Event is not supported by the relay"); + } + + if (payload === null || typeof payload !== "object" || Array.isArray(payload)) { + return fail("INVALID_PAYLOAD", "Payload must be a JSON object"); + } + + let sanitized; + try { + const byteLength = relayPayloadByteLength(payload); + if (byteLength > limits.maxBytes) { + return fail( + "PAYLOAD_TOO_LARGE", + `Payload size ${byteLength} bytes exceeds limit of ${limits.maxBytes} bytes`, + ); + } + sanitized = sanitizeRelayPayload(payload, limits); + } catch (err) { + if (err instanceof RelayValidationError) return fail(err.code, err.message); + throw err; + } + + const parsed = schemas[event].safeParse(sanitized); + if (!parsed.success) { + return fail( + "VALIDATION_FAILED", + "Payload failed validation", + parsed.error.issues.map((issue) => ({ + path: issue.path.join("."), + message: issue.message, + })), + ); + } + + return { ok: true, data: parsed.data }; +} + +function fail(code, message, issues) { + return { ok: false, error: issues ? { code, message, issues } : { code, message } }; +} + +// ─── Outbound sanitization ──────────────────────────────────────────────────── + +const OUTBOUND_LIMITS = Object.freeze({ + maxDepth: 10, + maxKeys: 256, + maxArrayLength: 1000, + maxStringLength: 16384, +}); + +/** + * Sanitize a payload before it is emitted to clients, so that data sourced + * from Horizon or merchant metadata cannot inject prototype keys, control + * characters or non-JSON values into dashboards. + * + * @param {object} payload + * @param {Partial} [limits] + * @returns {object} + * @throws {RelayValidationError} + */ +export function sanitizeOutboundPayload(payload, limits = {}) { + return sanitizeRelayPayload(payload, { ...OUTBOUND_LIMITS, ...limits }); +} + +// ─── Socket guard ───────────────────────────────────────────────────────────── + +/** + * Create a per-socket guard that validates every inbound event before its + * handler runs. + * + * On a validation failure the guard: + * 1. logs a warning (event name, error code, socket id — never the raw payload), + * 2. replies via the ack callback if one was supplied, otherwise emits + * `relay:error` to the socket, + * 3. counts a violation and disconnects the socket once `maxViolations` is reached. + * + * Events that have no registered handler are also counted as violations. + * + * @param {import("socket.io").Socket} socket + * @param {object} [opts] + * @param {{ warn: Function, error: Function }} [opts.logger] + * @param {number} [opts.maxViolations] + * @param {Record} [opts.schemas] + * @param {Partial} [opts.limits] + * @returns {{ on: (event: string, handler: (data: object, ack?: Function) => void) => void, violations: () => number }} + */ +export function createRelayGuard(socket, opts = {}) { + const logger = opts.logger ?? console; + const maxViolations = opts.maxViolations ?? DEFAULT_MAX_VIOLATIONS; + const schemas = opts.schemas ?? INBOUND_EVENT_SCHEMAS; + const handled = new Set(); + let violations = 0; + + const reject = (event, error, ack) => { + violations += 1; + logger.warn( + { socketId: socket.id, event, code: error.code, violations }, + "WebSocket relay: rejected inbound event", + ); + + const body = { ok: false, event, error }; + if (typeof ack === "function") { + ack(body); + } else { + socket.emit(RELAY_ERROR_EVENT, body); + } + + if (violations >= maxViolations) { + logger.warn( + { socketId: socket.id, violations }, + "WebSocket relay: disconnecting socket after repeated invalid events", + ); + socket.disconnect(true); + } + }; + + if (typeof socket.onAny === "function") { + socket.onAny((event, ...args) => { + if (handled.has(event)) return; + const ack = typeof args[args.length - 1] === "function" ? args[args.length - 1] : undefined; + reject( + typeof event === "string" ? event : String(event), + { code: "UNKNOWN_EVENT", message: "Event is not supported by the relay" }, + ack, + ); + }); + } + + return { + on(event, handler) { + if (!Object.hasOwn(schemas, event)) { + throw new Error(`No relay schema registered for event '${event}'`); + } + handled.add(event); + + socket.on(event, (...args) => { + const ack = typeof args[args.length - 1] === "function" ? args.pop() : undefined; + const result = validateInboundRelayEvent(event, args[0], { + schemas, + limits: opts.limits, + }); + + if (!result.ok) { + reject(event, result.error, ack); + return; + } + + try { + handler(result.data, ack); + } catch (err) { + logger.error({ err, socketId: socket.id, event }, "WebSocket relay: handler failed"); + const body = { + ok: false, + event, + error: { code: "INTERNAL_ERROR", message: "Failed to process event" }, + }; + if (typeof ack === "function") ack(body); + else socket.emit(RELAY_ERROR_EVENT, body); + } + }); + }, + violations: () => violations, + }; +} diff --git a/backend/src/lib/websocket-relay-validation.test.js b/backend/src/lib/websocket-relay-validation.test.js new file mode 100644 index 00000000..5736940b --- /dev/null +++ b/backend/src/lib/websocket-relay-validation.test.js @@ -0,0 +1,470 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { EventEmitter } from "node:events"; +import { + sanitizeRelayString, + sanitizeRelayPayload, + sanitizeOutboundPayload, + relayPayloadByteLength, + validateInboundRelayEvent, + createRelayGuard, + RelayValidationError, + RELAY_PAYLOAD_LIMITS, + RELAY_ERROR_EVENT, + DEFAULT_MAX_VIOLATIONS, +} from "./websocket-relay-validation.js"; + +const MERCHANT_ID = "3f1c2b4a-8d6e-4f7a-9b0c-1d2e3f4a5b6c"; +const PAYMENT_ID = "a1b2c3d4-e5f6-4a7b-8c9d-0e1f2a3b4c5d"; + +// ─── sanitizeRelayString ────────────────────────────────────────────────────── + +describe("sanitizeRelayString", () => { + it("removes C0/C1 control characters but keeps tab, newline and CR", () => { + expect(sanitizeRelayString("a\u0000b\u0007c\u001Fd\u007Fe\u0085f")).toBe("abcdef"); + expect(sanitizeRelayString("line1\nline2\tx\r")).toBe("line1\nline2\tx\r"); + }); + + it("removes bidi override characters used for text spoofing", () => { + expect(sanitizeRelayString("pay‮moc.live")).toBe("paymoc.live"); + expect(sanitizeRelayString("⁦x⁩")).toBe("x"); + }); + + it("normalises to NFC", () => { + expect(sanitizeRelayString("é")).toBe("é"); + }); +}); + +// ─── sanitizeRelayPayload ───────────────────────────────────────────────────── + +describe("sanitizeRelayPayload", () => { + it("returns an equal deep copy for clean input", () => { + const input = { a: 1, b: "x", c: [true, null, { d: 2.5 }] }; + const out = sanitizeRelayPayload(input); + expect(out).toEqual(input); + expect(out).not.toBe(input); + expect(out.c).not.toBe(input.c); + }); + + it("strips prototype-pollution keys at every depth", () => { + const input = JSON.parse( + '{"__proto__":{"polluted":true},"nested":{"constructor":{"prototype":{"x":1}},"prototype":1,"ok":1}}', + ); + const out = sanitizeRelayPayload(input); + expect(Object.keys(out)).toEqual(["nested"]); + expect(out.nested).toEqual({ ok: 1 }); + expect({}.polluted).toBeUndefined(); + }); + + it("strips keys that only become forbidden after sanitization", () => { + const out = sanitizeRelayPayload({ ["__pro\u0000to__"]: { x: 1 }, ["\u0000"]: 1 }); + expect(out).toEqual({}); + }); + + it("sanitizes strings inside objects and arrays", () => { + expect(sanitizeRelayPayload({ s: "a\u0000b", arr: ["‮c"] })).toEqual({ + s: "ab", + arr: ["c"], + }); + }); + + it("drops undefined, function and symbol values; arrays get null", () => { + const out = sanitizeRelayPayload({ + u: undefined, + f: () => {}, + s: Symbol("x"), + arr: [undefined, () => {}], + keep: 0, + }); + expect(out).toEqual({ keep: 0, arr: [null, null] }); + }); + + it("converts dates to ISO strings and rejects invalid dates", () => { + const d = new Date("2026-01-01T00:00:00.000Z"); + expect(sanitizeRelayPayload({ d })).toEqual({ d: "2026-01-01T00:00:00.000Z" }); + expect(() => sanitizeRelayPayload({ d: new Date("nope") })).toThrow(RelayValidationError); + }); + + it.each([ + ["NaN", { n: NaN }, "INVALID_NUMBER"], + ["Infinity", { n: Infinity }, "INVALID_NUMBER"], + ["BigInt", { n: 1n }, "UNSUPPORTED_TYPE"], + ["Buffer", { b: Buffer.from("x") }, "UNSUPPORTED_TYPE"], + ["ArrayBuffer", { b: new ArrayBuffer(2) }, "UNSUPPORTED_TYPE"], + ])("rejects %s values", (_label, input, code) => { + expect(() => sanitizeRelayPayload(input)).toThrow(expect.objectContaining({ code })); + }); + + it("rejects circular references", () => { + const a = { b: {} }; + a.b.a = a; + expect(() => sanitizeRelayPayload(a)).toThrow( + expect.objectContaining({ code: "CIRCULAR_REFERENCE" }), + ); + }); + + it("allows the same object to appear twice in sibling positions", () => { + const shared = { x: 1 }; + expect(sanitizeRelayPayload({ a: shared, b: shared })).toEqual({ a: { x: 1 }, b: { x: 1 } }); + }); + + it("enforces maximum depth", () => { + let deep = {}; + const root = deep; + for (let i = 0; i < RELAY_PAYLOAD_LIMITS.maxDepth + 1; i++) { + deep.next = {}; + deep = deep.next; + } + expect(() => sanitizeRelayPayload(root)).toThrow( + expect.objectContaining({ code: "PAYLOAD_TOO_DEEP" }), + ); + expect(() => sanitizeRelayPayload({ a: { b: {} } }, { maxDepth: 3 })).not.toThrow(); + }); + + it("enforces key, array and string limits", () => { + const manyKeys = Object.fromEntries( + Array.from({ length: RELAY_PAYLOAD_LIMITS.maxKeys + 1 }, (_, i) => [`k${i}`, i]), + ); + expect(() => sanitizeRelayPayload(manyKeys)).toThrow( + expect.objectContaining({ code: "TOO_MANY_KEYS" }), + ); + expect(() => + sanitizeRelayPayload({ a: new Array(RELAY_PAYLOAD_LIMITS.maxArrayLength + 1).fill(0) }), + ).toThrow(expect.objectContaining({ code: "ARRAY_TOO_LONG" })); + expect(() => + sanitizeRelayPayload({ s: "x".repeat(RELAY_PAYLOAD_LIMITS.maxStringLength + 1) }), + ).toThrow(expect.objectContaining({ code: "STRING_TOO_LONG" })); + }); + + it("returns an object without an inherited prototype-pollution surface", () => { + const out = sanitizeRelayPayload(JSON.parse('{"__proto__":{"admin":true}}')); + expect(out.admin).toBeUndefined(); + expect(Object.getPrototypeOf(out)).toBe(Object.prototype); + }); +}); + +// ─── relayPayloadByteLength ─────────────────────────────────────────────────── + +describe("relayPayloadByteLength", () => { + it("measures UTF-8 serialised bytes", () => { + expect(relayPayloadByteLength({ a: "é" })).toBe(Buffer.byteLength('{"a":"é"}')); + }); + + it("throws RelayValidationError on unserialisable input", () => { + const a = {}; + a.self = a; + expect(() => relayPayloadByteLength(a)).toThrow( + expect.objectContaining({ code: "UNSERIALIZABLE" }), + ); + }); +}); + +// ─── validateInboundRelayEvent ──────────────────────────────────────────────── + +describe("validateInboundRelayEvent", () => { + it.each([ + ["join:merchant", { merchant_id: MERCHANT_ID }], + ["join:checkout", { payment_id: PAYMENT_ID }], + ["leave:checkout", { payment_id: PAYMENT_ID }], + ])("accepts a valid %s payload", (event, payload) => { + expect(validateInboundRelayEvent(event, payload)).toEqual({ ok: true, data: payload }); + }); + + it("lower-cases and trims UUIDs so room names are canonical", () => { + const result = validateInboundRelayEvent("join:merchant", { + merchant_id: ` ${MERCHANT_ID.toUpperCase()} `, + }); + expect(result).toEqual({ ok: true, data: { merchant_id: MERCHANT_ID } }); + }); + + it("rejects unknown events", () => { + const result = validateInboundRelayEvent("admin:broadcast", {}); + expect(result.ok).toBe(false); + expect(result.error.code).toBe("UNKNOWN_EVENT"); + }); + + it("does not treat Object.prototype keys as registered events", () => { + expect(validateInboundRelayEvent("toString", {}).error.code).toBe("UNKNOWN_EVENT"); + expect(validateInboundRelayEvent("__proto__", {}).error.code).toBe("UNKNOWN_EVENT"); + }); + + it.each([ + ["undefined", undefined], + ["null", null], + ["a string", PAYMENT_ID], + ["a number", 42], + ["an array", [PAYMENT_ID]], + ])("rejects %s payload", (_label, payload) => { + const result = validateInboundRelayEvent("join:checkout", payload); + expect(result.ok).toBe(false); + expect(result.error.code).toBe("INVALID_PAYLOAD"); + }); + + it("rejects unknown keys (strict schema)", () => { + const result = validateInboundRelayEvent("join:checkout", { + payment_id: PAYMENT_ID, + room: "merchant:someone-else", + }); + expect(result.ok).toBe(false); + expect(result.error.code).toBe("VALIDATION_FAILED"); + }); + + it.each([ + ["missing", {}], + ["empty", { merchant_id: "" }], + ["non-string", { merchant_id: 123 }], + ["non-UUID", { merchant_id: "not-a-uuid" }], + ["room injection", { merchant_id: `${MERCHANT_ID}:admin` }], + ["nested object", { merchant_id: { $ne: null } }], + ])("rejects a %s merchant_id", (_label, payload) => { + const result = validateInboundRelayEvent("join:merchant", payload); + expect(result.ok).toBe(false); + expect(result.error.code).toBe("VALIDATION_FAILED"); + expect(result.error.issues[0].path).toBe("merchant_id"); + }); + + it("strips control characters before validating", () => { + const result = validateInboundRelayEvent("join:checkout", { + payment_id: `${PAYMENT_ID}\u0000`, + }); + expect(result).toEqual({ ok: true, data: { payment_id: PAYMENT_ID } }); + }); + + it("ignores __proto__ smuggling rather than failing the strict check", () => { + const payload = JSON.parse(`{"payment_id":"${PAYMENT_ID}","__proto__":{"x":1}}`); + expect(validateInboundRelayEvent("join:checkout", payload)).toEqual({ + ok: true, + data: { payment_id: PAYMENT_ID }, + }); + }); + + it("rejects payloads larger than the byte limit before deep inspection", () => { + const result = validateInboundRelayEvent("join:checkout", { + payment_id: PAYMENT_ID, + pad: "x".repeat(RELAY_PAYLOAD_LIMITS.maxBytes), + }); + expect(result.ok).toBe(false); + expect(result.error.code).toBe("PAYLOAD_TOO_LARGE"); + }); + + it("maps sanitization failures to their error code", () => { + const circular = { payment_id: PAYMENT_ID }; + circular.self = circular; + expect(validateInboundRelayEvent("join:checkout", circular).error.code).toBe( + "UNSERIALIZABLE", + ); + expect( + validateInboundRelayEvent("join:checkout", { payment_id: PAYMENT_ID, n: 1n }).error.code, + ).toBe("UNSERIALIZABLE"); + }); + + it("does not echo the raw payload back in the error", () => { + const secret = "sk_live_should_not_be_echoed"; + const result = validateInboundRelayEvent("join:merchant", { merchant_id: secret }); + expect(JSON.stringify(result)).not.toContain(secret); + }); +}); + +// ─── sanitizeOutboundPayload ────────────────────────────────────────────────── + +describe("sanitizeOutboundPayload", () => { + it("sanitizes emitted payloads and serialises dates", () => { + const out = sanitizeOutboundPayload({ + id: PAYMENT_ID, + memo: "hello‮evil", + confirmed_at: new Date("2026-01-01T00:00:00.000Z"), + extra: undefined, + }); + expect(out).toEqual({ + id: PAYMENT_ID, + memo: "helloevil", + confirmed_at: "2026-01-01T00:00:00.000Z", + }); + }); + + it("uses more generous limits than inbound validation", () => { + const long = "x".repeat(RELAY_PAYLOAD_LIMITS.maxStringLength + 1); + expect(sanitizeOutboundPayload({ s: long }).s).toBe(long); + }); +}); + +// ─── createRelayGuard ───────────────────────────────────────────────────────── + +class FakeSocket extends EventEmitter { + constructor() { + super(); + this.id = "socket-1"; + this.sent = []; + this.disconnected = false; + this.anyListeners = []; + } + + onAny(listener) { + this.anyListeners.push(listener); + } + + // Mirrors socket.io: catch-all listeners run before the event's own listeners + receive(event, ...args) { + for (const listener of this.anyListeners) listener(event, ...args); + super.emit(event, ...args); + } + + emit(event, ...args) { + this.sent.push([event, ...args]); + return true; + } + + disconnect(close) { + this.disconnected = close; + } +} + +describe("createRelayGuard", () => { + let socket; + let logger; + + beforeEach(() => { + socket = new FakeSocket(); + logger = { warn: vi.fn(), error: vi.fn() }; + }); + + it("passes validated data to the handler", () => { + const guard = createRelayGuard(socket, { logger }); + const handler = vi.fn(); + guard.on("join:merchant", handler); + + socket.receive("join:merchant", { merchant_id: MERCHANT_ID.toUpperCase() }); + + expect(handler).toHaveBeenCalledWith({ merchant_id: MERCHANT_ID }, undefined); + expect(socket.sent).toEqual([]); + expect(guard.violations()).toBe(0); + }); + + it("forwards the ack callback to the handler", () => { + const guard = createRelayGuard(socket, { logger }); + const handler = vi.fn((_data, ack) => ack({ ok: true })); + const ack = vi.fn(); + guard.on("join:checkout", handler); + + socket.receive("join:checkout", { payment_id: PAYMENT_ID }, ack); + + expect(ack).toHaveBeenCalledWith({ ok: true }); + }); + + it("does not crash on a missing payload (previous destructuring TypeError)", () => { + const guard = createRelayGuard(socket, { logger }); + const handler = vi.fn(); + guard.on("join:merchant", handler); + + expect(() => socket.receive("join:merchant")).not.toThrow(); + expect(() => socket.receive("join:merchant", null)).not.toThrow(); + expect(handler).not.toHaveBeenCalled(); + }); + + it("emits relay:error and logs without the raw payload on invalid input", () => { + const guard = createRelayGuard(socket, { logger }); + const handler = vi.fn(); + guard.on("join:merchant", handler); + + socket.receive("join:merchant", { merchant_id: "secret-token-value" }); + + expect(handler).not.toHaveBeenCalled(); + expect(socket.sent).toHaveLength(1); + const [event, body] = socket.sent[0]; + expect(event).toBe(RELAY_ERROR_EVENT); + expect(body).toMatchObject({ + ok: false, + event: "join:merchant", + error: { code: "VALIDATION_FAILED" }, + }); + expect(logger.warn).toHaveBeenCalledTimes(1); + expect(JSON.stringify(logger.warn.mock.calls)).not.toContain("secret-token-value"); + }); + + it("replies through the ack callback instead of relay:error when provided", () => { + const guard = createRelayGuard(socket, { logger }); + guard.on("join:checkout", vi.fn()); + const ack = vi.fn(); + + socket.receive("join:checkout", { payment_id: 1 }, ack); + + expect(ack).toHaveBeenCalledWith( + expect.objectContaining({ ok: false, error: expect.objectContaining({ code: "VALIDATION_FAILED" }) }), + ); + expect(socket.sent).toEqual([]); + }); + + it("rejects events that have no registered handler", () => { + const guard = createRelayGuard(socket, { logger }); + guard.on("join:merchant", vi.fn()); + + socket.receive("admin:broadcast", { anything: true }); + + expect(socket.sent[0][1]).toMatchObject({ + event: "admin:broadcast", + error: { code: "UNKNOWN_EVENT" }, + }); + expect(guard.violations()).toBe(1); + }); + + it("does not double-count registered events via the catch-all listener", () => { + const guard = createRelayGuard(socket, { logger }); + guard.on("join:merchant", vi.fn()); + + socket.receive("join:merchant", { merchant_id: "bad" }); + + expect(guard.violations()).toBe(1); + expect(socket.sent).toHaveLength(1); + }); + + it("disconnects the socket after maxViolations invalid events", () => { + const guard = createRelayGuard(socket, { logger, maxViolations: 3 }); + guard.on("join:checkout", vi.fn()); + + socket.receive("join:checkout", {}); + socket.receive("join:checkout", {}); + expect(socket.disconnected).toBe(false); + socket.receive("join:checkout", {}); + + expect(socket.disconnected).toBe(true); + expect(guard.violations()).toBe(3); + }); + + it("defaults maxViolations to DEFAULT_MAX_VIOLATIONS", () => { + const guard = createRelayGuard(socket, { logger }); + guard.on("join:checkout", vi.fn()); + + for (let i = 0; i < DEFAULT_MAX_VIOLATIONS - 1; i++) socket.receive("join:checkout", {}); + expect(socket.disconnected).toBe(false); + socket.receive("join:checkout", {}); + expect(socket.disconnected).toBe(true); + }); + + it("contains handler exceptions and reports INTERNAL_ERROR", () => { + const guard = createRelayGuard(socket, { logger }); + guard.on("join:checkout", () => { + throw new Error("boom"); + }); + + expect(() => socket.receive("join:checkout", { payment_id: PAYMENT_ID })).not.toThrow(); + expect(logger.error).toHaveBeenCalledTimes(1); + expect(socket.sent[0][1]).toMatchObject({ error: { code: "INTERNAL_ERROR" } }); + expect(guard.violations()).toBe(0); + }); + + it("refuses to register a handler for an event without a schema", () => { + const guard = createRelayGuard(socket, { logger }); + expect(() => guard.on("custom:event", vi.fn())).toThrow(/No relay schema/); + }); + + it("works on sockets without onAny", () => { + const bare = new FakeSocket(); + bare.onAny = undefined; + const guard = createRelayGuard(bare, { logger }); + const handler = vi.fn(); + guard.on("join:checkout", handler); + bare.receive = (event, ...args) => EventEmitter.prototype.emit.call(bare, event, ...args); + + bare.receive("join:checkout", { payment_id: PAYMENT_ID }); + expect(handler).toHaveBeenCalled(); + }); +});