feat(server): add rate-limit, origin allow-list, message-size cap (P4.4)

This commit is contained in:
Joey Yakimowich-Payne 2026-04-16 17:11:40 -06:00
commit aafc18ef9e
No known key found for this signature in database
2 changed files with 221 additions and 0 deletions

View file

@ -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);
});
});

View file

@ -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<string> {
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;
}