Skip to content
Closed
Show file tree
Hide file tree
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
199 changes: 199 additions & 0 deletions src/__tests__/withEventRateLimit.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,199 @@
import { describe, it, expect, vi, beforeEach } from "vitest";
import { NextRequest, NextResponse } from "next/server";
import { withEventRateLimit } from "@/middleware/withEventRateLimit";
import type { RequestHandler } from "@/middleware/types";

// Mock the events module
vi.mock("@/lib/server/events", () => ({
countEventsInTimeWindow: vi.fn(),
recordEvent: vi.fn(),
}));

const { countEventsInTimeWindow, recordEvent } = await import(
"@/lib/server/events"
);
const mockedCountEventsInTimeWindow = vi.mocked(countEventsInTimeWindow);
const mockedRecordEvent = vi.mocked(recordEvent);

describe("withEventRateLimit middleware", () => {
beforeEach(() => {
vi.clearAllMocks();
});

const mockContext = {
user_id: "test-user-123",
params: Promise.resolve({}),
};

const mockRequest = new NextRequest("http://localhost:3000/api/test");

const mockHandler: RequestHandler<typeof mockContext> = vi.fn(async () =>
NextResponse.json({ success: true }),
);

it("should allow request when under rate limit", async () => {
// Mock that user has made 5 calls (under the limit of 10)
mockedCountEventsInTimeWindow.mockResolvedValue(5);
mockedRecordEvent.mockResolvedValue(undefined);

const rateLimitedHandler = withEventRateLimit(
{
event: "fetch_reviews",
maxCalls: 10,
timeWindowHours: 24,
},
mockHandler,
);

const response = await rateLimitedHandler(mockRequest, mockContext);

expect(response.status).toBe(200);
expect(mockedCountEventsInTimeWindow).toHaveBeenCalledWith(
"fetch_reviews",
"test-user-123",
24,
);
expect(mockHandler).toHaveBeenCalledWith(mockRequest, mockContext);
expect(mockedRecordEvent).toHaveBeenCalledWith(
"fetch_reviews",
"test-user-123",
{},
);
});

it("should reject request when rate limit is exceeded", async () => {
// Mock that user has made 10 calls (at the limit)
mockedCountEventsInTimeWindow.mockResolvedValue(10);

const rateLimitedHandler = withEventRateLimit(
{
event: "fetch_reviews",
maxCalls: 10,
timeWindowHours: 24,
},
mockHandler,
);

const response = await rateLimitedHandler(mockRequest, mockContext);

expect(response.status).toBe(429);
const responseData = await response.json();
expect(responseData).toEqual({
error: "Rate limit exceeded",
details: {
event: "fetch_reviews",
maxCalls: 10,
timeWindowHours: 24,
currentCount: 10,
},
});

expect(mockedCountEventsInTimeWindow).toHaveBeenCalledWith(
"fetch_reviews",
"test-user-123",
24,
);
expect(mockHandler).not.toHaveBeenCalled();
expect(mockedRecordEvent).not.toHaveBeenCalled();
});

it("should not record event when handler returns error status", async () => {
// Mock that user has made 2 calls (under the limit)
mockedCountEventsInTimeWindow.mockResolvedValue(2);

// Create a handler that returns an error
const errorHandler: RequestHandler<typeof mockContext> = vi.fn(async () =>
NextResponse.json({ error: "Bad request" }, { status: 400 }),
);

const rateLimitedHandler = withEventRateLimit(
{
event: "fetch_reviews",
maxCalls: 10,
timeWindowHours: 24,
},
errorHandler,
);

const response = await rateLimitedHandler(mockRequest, mockContext);

expect(response.status).toBe(400);
expect(mockedCountEventsInTimeWindow).toHaveBeenCalled();
expect(errorHandler).toHaveBeenCalled();
expect(mockedRecordEvent).not.toHaveBeenCalled();
});

it("should record event with custom metadata", async () => {
mockedCountEventsInTimeWindow.mockResolvedValue(1);
mockedRecordEvent.mockResolvedValue(undefined);

const customMetadata = { business_id: 123 };
const rateLimitedHandler = withEventRateLimit(
{
event: "update_reviews",
maxCalls: 5,
timeWindowHours: 12,
metadata: customMetadata,
},
mockHandler,
);

const response = await rateLimitedHandler(mockRequest, mockContext);

expect(response.status).toBe(200);
expect(mockedRecordEvent).toHaveBeenCalledWith(
"update_reviews",
"test-user-123",
customMetadata,
);
});

it("should handle database errors gracefully", async () => {
// Mock database error
mockedCountEventsInTimeWindow.mockRejectedValue(
new Error("Database error"),
);

const rateLimitedHandler = withEventRateLimit(
{
event: "fetch_reviews",
maxCalls: 10,
timeWindowHours: 24,
},
mockHandler,
);

const response = await rateLimitedHandler(mockRequest, mockContext);

expect(response.status).toBe(500);
const responseData = await response.json();
expect(responseData).toEqual({
error: "Internal Server Error",
});

expect(mockHandler).not.toHaveBeenCalled();
});

it("should work with different time windows and limits", async () => {
mockedCountEventsInTimeWindow.mockResolvedValue(50);
mockedRecordEvent.mockResolvedValue(undefined);

const rateLimitedHandler = withEventRateLimit(
{
event: "fetch_reviews",
maxCalls: 100,
timeWindowHours: 1, // 1 hour window
},
mockHandler,
);

const response = await rateLimitedHandler(mockRequest, mockContext);

expect(response.status).toBe(200);
expect(mockedCountEventsInTimeWindow).toHaveBeenCalledWith(
"fetch_reviews",
"test-user-123",
1,
);
});
});
105 changes: 56 additions & 49 deletions src/app/api/(endpoints)/v1/fetch-updated-data/route.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import schema from "./schema";
import { NextRouteContext, RequestHandler } from "@/middleware/types";
import { withBody } from "@/middleware/withBody";
import { withApiKey } from "@/middleware/withApiKey";
import { withEventRateLimit } from "@/middleware/withEventRateLimit";
import { withPaidAccess } from "@/middleware/withPaidAccess";

import { getLastEvent } from "@/lib/server/events";
Expand All @@ -23,65 +24,71 @@ import { businesses } from "@/schema/schema";
* Checks if reviews/stats need updating, updates if needed, then returns latest data. This endpoint is called by 11ty in the clients
* website to ensure that their google reviews are updated any time the clients site is rebuilt.
*
* Rate limits:
* - fetch_reviews: 100 calls per 24 hours (for this endpoint usage)
*
* @param { business_id: number } - The database ID of the business
* @returns Latest reviews and stats for the business
*/
export const POST: RequestHandler<NextRouteContext> = withApiKey(
withPaidAccess(
withBody(schema, async (_, context) => {
try {
const { business_id } = context.body;
withEventRateLimit(
{ event: "fetch_reviews", maxCalls: 100, timeWindowHours: 24 },
withBody(schema, async (_, context) => {
try {
const { business_id } = context.body;

await userHasOwnership(context.user_id, business_id, businesses);
const oneDayAgo = new Date();
oneDayAgo.setDate(oneDayAgo.getDate() - 1);
await userHasOwnership(context.user_id, business_id, businesses);
const oneDayAgo = new Date();
oneDayAgo.setDate(oneDayAgo.getDate() - 1);

// Check last update times
const lastUpdateReviews = await getLastEvent(
"update_reviews",
context.user_id,
);
const lastUpdateStats = await getLastEvent(
"update_stats",
context.user_id,
);
// Check last update times
const lastUpdateReviews = await getLastEvent(
"update_reviews",
context.user_id,
);
const lastUpdateStats = await getLastEvent(
"update_stats",
context.user_id,
);

// If data is out of date, fetch and update
if (
!lastUpdateReviews?.timestamp ||
lastUpdateReviews.timestamp < oneDayAgo
) {
await updateBusinessReviews(business_id);
}
// If data is out of date, fetch and update
if (
!lastUpdateReviews?.timestamp ||
lastUpdateReviews.timestamp < oneDayAgo
) {
await updateBusinessReviews(business_id);
}

if (
!lastUpdateStats?.timestamp ||
lastUpdateStats.timestamp < oneDayAgo
) {
await updateBusinessStats(business_id);
}
if (
!lastUpdateStats?.timestamp ||
lastUpdateStats.timestamp < oneDayAgo
) {
await updateBusinessStats(business_id);
}

// Get the data
const [reviews, stats] = await Promise.all([
selectBusinessReviews(business_id),
selectBusinessStats(business_id),
]);
// Get the data
const [reviews, stats] = await Promise.all([
selectBusinessReviews(business_id),
selectBusinessStats(business_id),
]);

const response = schema.response.parse({
reviews: reviews.map((review) => ({
...review,
datetime: review.datetime ? review.datetime.toISOString() : null,
})),
stats,
});
return NextResponse.json(response);
} catch (error) {
console.error("Error processing request:", error);
return NextResponse.json(
{ error: "Internal Server Error" },
{ status: 500 },
);
}
}),
const response = schema.response.parse({
reviews: reviews.map((review) => ({
...review,
datetime: review.datetime ? review.datetime.toISOString() : null,
})),
stats,
});
return NextResponse.json(response);
} catch (error) {
console.error("Error processing request:", error);
return NextResponse.json(
{ error: "Internal Server Error" },
{ status: 500 },
);
}
}),
),
),
);
52 changes: 52 additions & 0 deletions src/app/api/demo/rate-limit-test/route.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
import { NextResponse } from "next/server";

import schema from "./schema";
import { NextRouteContext, RequestHandler } from "@/middleware/types";
import { withApiKey } from "@/middleware/withApiKey";
import { withEventRateLimit } from "@/middleware/withEventRateLimit";
import { countEventsInTimeWindow } from "@/lib/server/events";

/**
* Demo endpoint showing multiple rate limits on a single endpoint.
* This demonstrates the capability to chain multiple withEventRateLimit middlewares.
*
* Rate limits:
* - fetch_reviews: 10 calls per 24 hours
* - update_reviews: 3 calls per 24 hours
*
* This endpoint will be rejected if either rate limit is exceeded.
*/
export const GET: RequestHandler<NextRouteContext> = withApiKey(
withEventRateLimit(
{ event: "fetch_reviews", maxCalls: 10, timeWindowHours: 24 },
withEventRateLimit(
{ event: "update_reviews", maxCalls: 3, timeWindowHours: 24 },
async (_, context) => {
try {
// Get current event counts for demonstration
const [fetchReviewsCount, updateReviewsCount] = await Promise.all([
countEventsInTimeWindow("fetch_reviews", context.user_id, 24),
countEventsInTimeWindow("update_reviews", context.user_id, 24),
]);

const response = schema.response.parse({
message: "Rate limit test endpoint accessed successfully!",
userId: context.user_id,
eventCounts: {
fetchReviews: fetchReviewsCount,
updateReviews: updateReviewsCount,
},
});

return NextResponse.json(response);
} catch (error) {
console.error("Error in rate limit test endpoint:", error);
return NextResponse.json(
{ error: "Internal Server Error" },
{ status: 500 },
);
}
},
),
),
);
Loading