feat(server): add rate-limit, origin allow-list, message-size cap (P4.4)
This commit is contained in:
parent
7d07bb78ba
commit
aafc18ef9e
2 changed files with 221 additions and 0 deletions
150
packages/server/src/middleware.test.ts
Normal file
150
packages/server/src/middleware.test.ts
Normal 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);
|
||||||
|
});
|
||||||
|
});
|
||||||
71
packages/server/src/middleware.ts
Normal file
71
packages/server/src/middleware.ts
Normal 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;
|
||||||
|
}
|
||||||
Loading…
Add table
Add a link
Reference in a new issue