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