From aafc18ef9ed6bdc9bebd4d8173ebf7ac6e3a1b44 Mon Sep 17 00:00:00 2001 From: Joey Yakimowich-Payne Date: Thu, 16 Apr 2026 17:11:40 -0600 Subject: [PATCH] feat(server): add rate-limit, origin allow-list, message-size cap (P4.4) --- packages/server/src/middleware.test.ts | 150 +++++++++++++++++++++++++ packages/server/src/middleware.ts | 71 ++++++++++++ 2 files changed, 221 insertions(+) create mode 100644 packages/server/src/middleware.test.ts create mode 100644 packages/server/src/middleware.ts diff --git a/packages/server/src/middleware.test.ts b/packages/server/src/middleware.test.ts new file mode 100644 index 0000000..0e58660 --- /dev/null +++ b/packages/server/src/middleware.test.ts @@ -0,0 +1,150 @@ +import { describe, it, expect, afterEach } from "vitest"; +import { + RateLimiter, + checkOrigin, + checkMessageSize, + MAX_MESSAGE_BYTES, +} from "./middleware.js"; + +// --------------------------------------------------------------------------- +// RateLimiter +// --------------------------------------------------------------------------- + +describe("RateLimiter", () => { + it("allows exactly 20 rapid messages (burst capacity)", () => { + const limiter = new RateLimiter(20, 100); + const results: boolean[] = []; + for (let i = 0; i < 20; i++) { + results.push(limiter.consume()); + } + expect(results).toHaveLength(20); + expect(results.every((r) => r === true)).toBe(true); + }); + + it("rejects the 21st message when burst is exceeded", () => { + const limiter = new RateLimiter(20, 100); + for (let i = 0; i < 20; i++) limiter.consume(); + expect(limiter.consume()).toBe(false); + }); + + it("refills tokens after waiting 100ms, then allows next message", async () => { + const limiter = new RateLimiter(20, 100); + // drain the bucket + for (let i = 0; i < 20; i++) limiter.consume(); + // 21st is blocked + expect(limiter.consume()).toBe(false); + + // Wait 100ms — at 100 tokens/sec that's 10 tokens refilled + await new Promise((resolve) => setTimeout(resolve, 100)); + expect(limiter.consume()).toBe(true); + }); + + it("stress test: 200 synchronous consume() calls — first 20 pass, 21st fails", () => { + const limiter = new RateLimiter(20, 100); + const results: boolean[] = []; + for (let i = 0; i < 200; i++) { + results.push(limiter.consume()); + } + + // First 20 must pass + for (let i = 0; i < 20; i++) { + expect(results[i], `message ${i + 1} should pass`).toBe(true); + } + // 21st must fail + expect(results[20], "message 21 should be rejected").toBe(false); + }); +}); + +// --------------------------------------------------------------------------- +// checkOrigin +// --------------------------------------------------------------------------- + +describe("checkOrigin", () => { + const originalEnv = process.env.ALLOWED_ORIGINS; + + afterEach(() => { + // Restore env after each test that mutates it + if (originalEnv === undefined) { + delete process.env.ALLOWED_ORIGINS; + } else { + process.env.ALLOWED_ORIGINS = originalEnv; + } + }); + + it("allows http://localhost:5173 by default", () => { + delete process.env.ALLOWED_ORIGINS; + expect(checkOrigin("http://localhost:5173")).toBe(true); + }); + + it("rejects an arbitrary unlisted origin by default", () => { + delete process.env.ALLOWED_ORIGINS; + expect(checkOrigin("https://evil.example.com")).toBe(false); + }); + + it("rejects null origin", () => { + expect(checkOrigin(null)).toBe(false); + }); + + it("rejects empty string origin", () => { + expect(checkOrigin("")).toBe(false); + }); + + it("respects ALLOWED_ORIGINS env override", () => { + process.env.ALLOWED_ORIGINS = + "https://app.example.com,https://staging.example.com"; + expect(checkOrigin("https://app.example.com")).toBe(true); + expect(checkOrigin("https://staging.example.com")).toBe(true); + expect(checkOrigin("http://localhost:5173")).toBe(false); + }); + + it("trims whitespace in ALLOWED_ORIGINS entries", () => { + process.env.ALLOWED_ORIGINS = " https://app.example.com , https://b.com "; + expect(checkOrigin("https://app.example.com")).toBe(true); + expect(checkOrigin("https://b.com")).toBe(true); + }); +}); + +// --------------------------------------------------------------------------- +// checkMessageSize +// --------------------------------------------------------------------------- + +describe("checkMessageSize", () => { + it("allows a 1-byte string", () => { + expect(checkMessageSize("x")).toBe(true); + }); + + it("allows exactly 64 KB string", () => { + const msg = "a".repeat(MAX_MESSAGE_BYTES); + expect(checkMessageSize(msg)).toBe(true); + }); + + it("rejects a string 1 byte over 64 KB", () => { + const msg = "a".repeat(MAX_MESSAGE_BYTES + 1); + expect(checkMessageSize(msg)).toBe(false); + }); + + it("allows a 1-byte Buffer", () => { + expect(checkMessageSize(Buffer.from([0x41]))).toBe(true); + }); + + it("allows exactly 64 KB Buffer", () => { + const buf = Buffer.alloc(MAX_MESSAGE_BYTES); + expect(checkMessageSize(buf)).toBe(true); + }); + + it("rejects a Buffer 1 byte over 64 KB", () => { + const buf = Buffer.alloc(MAX_MESSAGE_BYTES + 1); + expect(checkMessageSize(buf)).toBe(false); + }); + + it("correctly measures multi-byte UTF-8 characters", () => { + // Each '€' is 3 bytes in UTF-8; (MAX_MESSAGE_BYTES / 3) chars = exactly under limit + const euroCount = Math.floor(MAX_MESSAGE_BYTES / 3); + const justUnder = "€".repeat(euroCount); // 3 * floor(65536/3) bytes + expect(checkMessageSize(justUnder)).toBe(true); + + // One more euro sign pushes it over for most counts + const justOver = "€".repeat(Math.ceil(MAX_MESSAGE_BYTES / 3) + 1); + expect(checkMessageSize(justOver)).toBe(false); + }); +}); diff --git a/packages/server/src/middleware.ts b/packages/server/src/middleware.ts new file mode 100644 index 0000000..3bfc3a1 --- /dev/null +++ b/packages/server/src/middleware.ts @@ -0,0 +1,71 @@ +// WebSocket middleware: rate limiting, origin allow-list, message size cap. +// See PROTOCOL.md for full spec. + +// --------------------------------------------------------------------------- +// Rate limiter — token-bucket algorithm, per connection +// --------------------------------------------------------------------------- + +export class RateLimiter { + readonly #capacity: number; + readonly #refillRate: number; // tokens per second + #tokens: number; + #lastRefill: number; + + constructor(capacity = 20, refillRate = 100) { + this.#capacity = capacity; + this.#refillRate = refillRate; + this.#tokens = capacity; + this.#lastRefill = performance.now(); + } + + /** + * Returns true if the message is allowed; false if rate-limited. + * Refills tokens based on elapsed time since last call. + */ + consume(): boolean { + const now = performance.now(); + const elapsed = (now - this.#lastRefill) / 1000; // seconds + this.#tokens = Math.min( + this.#capacity, + this.#tokens + elapsed * this.#refillRate, + ); + this.#lastRefill = now; + + if (this.#tokens < 1) return false; + this.#tokens -= 1; + return true; + } +} + +// --------------------------------------------------------------------------- +// Origin allow-list +// --------------------------------------------------------------------------- + +export function getAllowedOrigins(): Set { + const env = process.env.ALLOWED_ORIGINS ?? "http://localhost:5173"; + return new Set(env.split(",").map((o) => o.trim())); +} + +/** + * Returns true when requestOrigin is present and in the allow-list. + * A null/empty origin is always rejected. + */ +export function checkOrigin(requestOrigin: string | null): boolean { + if (!requestOrigin) return false; + return getAllowedOrigins().has(requestOrigin); +} + +// --------------------------------------------------------------------------- +// Message size cap — 64 KB +// --------------------------------------------------------------------------- + +export const MAX_MESSAGE_BYTES = 64 * 1024; // 64 KB + +/** + * Returns true when the message is within the 64 KB limit; false otherwise. + */ +export function checkMessageSize(msg: string | Buffer): boolean { + const size = + typeof msg === "string" ? Buffer.byteLength(msg, "utf8") : msg.length; + return size <= MAX_MESSAGE_BYTES; +}