diff --git a/lib/share/rate-limit-core.test.ts b/lib/share/rate-limit-core.test.ts new file mode 100644 index 0000000..0d5d8e8 --- /dev/null +++ b/lib/share/rate-limit-core.test.ts @@ -0,0 +1,53 @@ +import { describe, expect, test } from "bun:test"; +import { enforceFixedWindowRateLimit } from "@/lib/share/rate-limit-core"; + +function memoryIncrement() { + const counts = new Map(); + return async (key: string) => { + const count = (counts.get(key) ?? 0) + 1; + counts.set(key, count); + return count; + }; +} + +describe("share fixed-window rate limits", () => { + for (const [name, limit] of [ + ["sender shares", 10], + ["bot requests", 60], + ["comments", 30], + ] as const) { + test(`allows ${limit} ${name} and rejects request ${limit + 1}`, async () => { + const increment = memoryIncrement(); + const options = { + key: name, + limit, + windowSeconds: 3_600, + increment, + now: 3_601_000, + }; + + for (let request = 1; request <= limit; request += 1) { + await expect(enforceFixedWindowRateLimit(options)).resolves.toBeUndefined(); + } + await expect(enforceFixedWindowRateLimit(options)).rejects.toMatchObject({ + status: 429, + retryAfter: 3_599, + }); + }); + } + + test("fails closed when Redis is unavailable", async () => { + await expect( + enforceFixedWindowRateLimit({ + key: "sender:1", + limit: 10, + windowSeconds: 3_600, + increment: async () => { + throw new Error("offline"); + }, + }), + ).rejects.toEqual( + expect.objectContaining({ status: 503 }), + ); + }); +}); diff --git a/lib/share/rate-limit-core.ts b/lib/share/rate-limit-core.ts new file mode 100644 index 0000000..c3b4e01 --- /dev/null +++ b/lib/share/rate-limit-core.ts @@ -0,0 +1,41 @@ +import { ShareHttpError } from "@/lib/share/http-error"; + +export type FixedWindowIncrement = ( + key: string, + expiresInSeconds: number, +) => Promise; + +export async function enforceFixedWindowRateLimit({ + key, + limit, + windowSeconds, + increment, + now = Date.now(), +}: { + key: string; + limit: number; + windowSeconds: number; + increment: FixedWindowIncrement; + now?: number; +}) { + const nowSeconds = Math.floor(now / 1000); + const window = Math.floor(nowSeconds / windowSeconds); + const retryAfter = windowSeconds - (nowSeconds % windowSeconds); + + let count: number; + try { + count = await increment( + `erika:share:${key}:${window}`, + windowSeconds + 1, + ); + } catch { + throw new ShareHttpError("Rate limiting is temporarily unavailable", 503); + } + + if (!Number.isFinite(count)) { + throw new ShareHttpError("Rate limiting is temporarily unavailable", 503); + } + if (count > limit) { + throw new ShareHttpError("Rate limit exceeded", 429, retryAfter); + } +} diff --git a/lib/share/rate-limit.ts b/lib/share/rate-limit.ts index 643f3de..2d57cdb 100644 --- a/lib/share/rate-limit.ts +++ b/lib/share/rate-limit.ts @@ -1,7 +1,7 @@ import "server-only"; import { getRedisClient } from "@/lib/redis"; -import { ShareHttpError } from "@/lib/share/http-error"; +import { enforceFixedWindowRateLimit } from "@/lib/share/rate-limit-core"; const FIXED_WINDOW_SCRIPT = ` local count = redis.call("INCR", KEYS[1]) @@ -20,25 +20,20 @@ export async function enforceShareRateLimit({ limit: number; windowSeconds: number; }) { - const nowSeconds = Math.floor(Date.now() / 1000); - const window = Math.floor(nowSeconds / windowSeconds); - const retryAfter = windowSeconds - (nowSeconds % windowSeconds); - - try { - const redis = await getRedisClient(); - const count = Number( - await redis.send("EVAL", [ - FIXED_WINDOW_SCRIPT, - "1", - `erika:share:${key}:${window}`, - String(windowSeconds + 1), - ]) - ); - if (count > limit) { - throw new ShareHttpError("Rate limit exceeded", 429, retryAfter); - } - } catch (error) { - if (error instanceof ShareHttpError) throw error; - throw new ShareHttpError("Rate limiting is temporarily unavailable", 503); - } + return enforceFixedWindowRateLimit({ + key, + limit, + windowSeconds, + increment: async (redisKey, expiresInSeconds) => { + const redis = await getRedisClient(); + return Number( + await redis.send("EVAL", [ + FIXED_WINDOW_SCRIPT, + "1", + redisKey, + String(expiresInSeconds), + ]), + ); + }, + }); }