Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
136 changes: 136 additions & 0 deletions packages/server/src/middleware/rate-limit.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,136 @@
import { Hono } from "hono";
import { createMiddleware } from "hono/factory";
import { describe, expect, it } from "vitest";

import { createMemoryApiRateLimitStorage } from "../lib/rate-limit-storage";
import type { StorageVariables } from "../storage";
import type { BetterAuthSessionVariables } from "./better-auth-session";
import { apiRateLimiter } from "./rate-limit";

const RATE_LIMIT = 100;
const EXTRA_REQUESTS = 5;

function createMockStorage(): StorageVariables["storage"] {
return {
apiRateLimitStorage: createMemoryApiRateLimitStorage(),
} as unknown as StorageVariables["storage"];
}

function createSession(userId: string): BetterAuthSessionVariables["session"] {
return { user: { id: userId } } as BetterAuthSessionVariables["session"];
}

function createApp(input: {
enabled: boolean;
getSession?: () => BetterAuthSessionVariables["session"];
}) {
const storage = createMockStorage();
return new Hono<{
Variables: BetterAuthSessionVariables;
}>()
.use(
"*",
createMiddleware<{ Variables: BetterAuthSessionVariables }>(
async (c, next) => {
c.set("storage", storage);
c.set("session", input.getSession?.() ?? null);
await next();
}
)
)
.use("*", apiRateLimiter({ enabled: input.enabled }))
.get("/api/auth", (c) => c.json({ ok: true }))
.get("/api/auth/sign-in", (c) => c.json({ ok: true }))
.get("/api/data", (c) => c.json({ ok: true }))
.get("/api/health", (c) => c.json({ ok: true }))
.get("/api/webhook", (c) => c.json({ ok: true }))
.get("/api/webhooks/stripe", (c) => c.json({ ok: true }));
}

type TestApp = ReturnType<typeof createApp>;

async function requestPath(
app: TestApp,
path: string,
headers?: Record<string, string>
) {
return app.request(path, { headers });
}

async function expectAllOk(
app: TestApp,
path: string,
count: number,
headers?: Record<string, string>
) {
for (let index = 0; index < count; index += 1) {
const response = await requestPath(app, path, headers);
expect(response.status).toBe(200);
}
}

describe("api rate limiter", () => {
it("limits anonymous requests by client IP after 100 requests per minute", async () => {
const app = createApp({ enabled: true });
const headers = { "cf-connecting-ip": "203.0.113.10" };

await expectAllOk(app, "/api/data", RATE_LIMIT, headers);

const response = await requestPath(app, "/api/data", headers);
expect(response.status).toBe(429);
await expect(response.json()).resolves.toEqual({
error: {
code: "rate_limited",
message: "Too many requests, please try again later.",
},
});
});

it("tracks separate buckets for different client IPs", async () => {
const app = createApp({ enabled: true });
const firstIp = { "cf-connecting-ip": "203.0.113.10" };
const secondIp = { "cf-connecting-ip": "203.0.113.11" };

await expectAllOk(app, "/api/data", RATE_LIMIT, firstIp);

const otherIpResponse = await requestPath(app, "/api/data", secondIp);
expect(otherIpResponse.status).toBe(200);

const limitedResponse = await requestPath(app, "/api/data", firstIp);
expect(limitedResponse.status).toBe(429);
});

it("keys authenticated requests by user id instead of client IP", async () => {
let session = createSession("user-a");
const app = createApp({ enabled: true, getSession: () => session });
const headers = { "cf-connecting-ip": "203.0.113.10" };

await expectAllOk(app, "/api/data", RATE_LIMIT, headers);

session = createSession("user-b");
const otherUserResponse = await requestPath(app, "/api/data", headers);
expect(otherUserResponse.status).toBe(200);

session = createSession("user-a");
const limitedResponse = await requestPath(app, "/api/data", headers);
expect(limitedResponse.status).toBe(429);
});

it.each([
"/api/health",
"/api/webhook",
"/api/webhooks/stripe",
"/api/auth",
"/api/auth/sign-in",
])("skips %s even past the limit", async (path) => {
const app = createApp({ enabled: true });

await expectAllOk(app, path, RATE_LIMIT + EXTRA_REQUESTS);
});

it("does not limit any request when disabled", async () => {
const app = createApp({ enabled: false });

await expectAllOk(app, "/api/data", RATE_LIMIT + EXTRA_REQUESTS);
});
});