← back to Handbag Auth Nextjs

src/lib/rateLimit.ts

118 lines

// Rate Limiting Middleware
// Prevents API abuse and DDoS attacks

interface RateLimitEntry {
  count: number
  resetTime: number
}

const rateLimitStore = new Map<string, RateLimitEntry>()

export interface RateLimitConfig {
  windowMs: number  // Time window in milliseconds
  maxRequests: number  // Max requests per window
}

const DEFAULT_CONFIG: RateLimitConfig = {
  windowMs: 60 * 1000,  // 1 minute
  maxRequests: 100  // 100 requests per minute
}

/**
 * Rate limiter using in-memory store
 * For production, consider using Redis for distributed rate limiting
 */
export function rateLimit(
  identifier: string,
  config: RateLimitConfig = DEFAULT_CONFIG
): { allowed: boolean; remaining: number; resetTime: number } {
  const now = Date.now()
  const entry = rateLimitStore.get(identifier)

  // Clean up expired entries periodically
  if (Math.random() < 0.01) {  // 1% chance on each call
    cleanupExpiredEntries(now)
  }

  if (!entry || now > entry.resetTime) {
    // Create new entry or reset expired entry
    const newEntry: RateLimitEntry = {
      count: 1,
      resetTime: now + config.windowMs
    }
    rateLimitStore.set(identifier, newEntry)

    return {
      allowed: true,
      remaining: config.maxRequests - 1,
      resetTime: newEntry.resetTime
    }
  }

  // Increment existing entry
  entry.count++

  if (entry.count > config.maxRequests) {
    return {
      allowed: false,
      remaining: 0,
      resetTime: entry.resetTime
    }
  }

  return {
    allowed: true,
    remaining: config.maxRequests - entry.count,
    resetTime: entry.resetTime
  }
}

/**
 * Clean up expired entries to prevent memory leaks
 */
function cleanupExpiredEntries(now: number): void {
  for (const [key, entry] of rateLimitStore.entries()) {
    if (now > entry.resetTime) {
      rateLimitStore.delete(key)
    }
  }
}

/**
 * Get client IP from request headers
 */
export function getClientIp(request: Request): string {
  const forwarded = request.headers.get('x-forwarded-for')
  const realIp = request.headers.get('x-real-ip')

  if (forwarded) {
    return forwarded.split(',')[0].trim()
  }

  if (realIp) {
    return realIp
  }

  return 'unknown'
}

/**
 * Get rate limit stats (for monitoring/debugging)
 */
export function getRateLimitStats(): {
  totalEntries: number
  entries: Array<{ ip: string; count: number; expiresIn: number }>
} {
  const now = Date.now()
  const entries = Array.from(rateLimitStore.entries()).map(([ip, entry]) => ({
    ip,
    count: entry.count,
    expiresIn: Math.max(0, entry.resetTime - now)
  }))

  return {
    totalEntries: rateLimitStore.size,
    entries
  }
}