import Foundation import HTTPTypes import Hummingbird import NIOCore import Synchronization /// Rejects a client's requests with `429 Too Many Requests` once they exceed a fixed-window rate limit. /// /// Added to the routes that must not be hammered — the subscription endpoint, an unauthenticated database write — it admits up to the configured limit of /// requests per client per window, and answers the excess with `429 Too Many Requests` and a `Retry-After` header naming the seconds until the /// window resets. /// /// A client is keyed by the first `X-Forwarded-For` entry when the ``Configuration`` trusts it, by the connection's remote address otherwise, and by /// one shared bucket when neither names the client. The counters live in memory with a bounded capacity, so a flood of distinct clients cannot grow the /// store without bound — and each instance of a multi-instance deployment enforces its own budget. public struct RateLimitMiddleware: Sendable { // MARK: Properties /// The fixed-window request counters, keyed by client. private let buckets: Buckets /// The limits the middleware enforces. private let configuration: Configuration // MARK: Initializers /// Creates a rate-limit middleware. /// - Parameter configuration: the limits the middleware enforces. Defaults to a budget suited to a form endpoint: a handful of requests per /// client per minute. public init( configuration: Configuration = .init() ) { self.buckets = .init( limit: configuration.limit, window: configuration.window ) self.configuration = configuration } } // MARK: - RouterMiddleware extension RateLimitMiddleware: RouterMiddleware { // MARK: Functions /// Passes the request down the chain while the client stays within its budget, and answers it with `429 Too Many Requests` and a `Retry-After` /// header once it does not. /// - Parameters: /// - request: the incoming request. /// - context: the context the request is resolved against. /// - next: the next responder in the middleware chain. /// - Returns: the downstream response, or the `429` rejection. /// - Throws: any error thrown downstream. public func handle( _ request: Request, context: Context, next: (Request, Context) async throws -> Response ) async throws -> Response { let admission = buckets.admit( client( for: request, context: context ) ) switch admission { case .admitted: return try await next( request, context ) case .limited(let retryAfter): var response = Response(status: .tooManyRequests) response.headers[.retryAfter] = String(max(1, retryAfter.components.seconds)) return response } } } // MARK: - Helpers private extension RateLimitMiddleware { // MARK: Methods /// The key identifying the requesting client: the first `X-Forwarded-For` entry when trusted, the connection's remote address otherwise, and one /// bucket shared by every unidentifiable client when neither is known. /// - Parameters: /// - request: the incoming request. /// - context: the context the request is resolved against. /// - Returns: the client key the request is counted under. func client( for request: Request, context: Context ) -> String { if configuration.trustForwardedFor, let forwarded = request.headers[.xForwardedFor]? .split(separator: ",") .first? .trimmingCharacters(in: .whitespaces), !forwarded.isEmpty { return forwarded } if let address = (context as? any RemoteAddressRequestContext)?.remoteAddress { return address.ipAddress ?? address.description } return .unidentified } } // MARK: - Buckets private extension RateLimitMiddleware { /// The outcome of asking the ``Buckets`` store to admit a request. enum Admission { /// The request is within the client's budget. case admitted /// The client exhausted its budget; the payload is the time until its window resets. case limited(retryAfter: Duration) } /// The fixed-window request counters, keyed by client. /// /// The counters sit behind a mutex rather than an actor: an admission is a handful of dictionary operations, so the lock is held only briefly and the /// calling task never suspends — requests skip the executor hop an actor would add on every pass through the middleware. final class Buckets: Sendable { // MARK: Properties /// The maximum number of clients tracked at once, bounding the store's memory. private let capacity: Int /// The per-client counters: the start of the client's current window and its request count. private let counters: Mutex<[String: (start: ContinuousClock.Instant, count: Int)]> /// The number of requests admitted per client per ``window``. private let limit: Int /// The length of the fixed window the ``limit`` applies to. private let window: Duration // MARK: Initializers /// Creates a counter store. /// - Parameters: /// - limit: the number of requests admitted per client per window. /// - window: the length of the fixed window the limit applies to. /// - capacity: the maximum number of clients tracked at once. init( limit: Int, window: Duration, capacity: Int = 10_000 ) { self.capacity = capacity self.counters = .init([:]) self.limit = limit self.window = window } // MARK: Functions /// Counts a request against the client's current window and admits it while the count stays within the limit. /// - Parameter client: the key the request is counted under. /// - Returns: the admission outcome. func admit( _ client: String ) -> Admission { let now = ContinuousClock.now return counters.withLock { counters in if let counter = counters[client], now < counter.start.advanced(by: window) { guard counter.count < limit else { return .limited(retryAfter: now.duration(to: counter.start.advanced(by: window))) } counters[client] = (counter.start, counter.count + 1) return .admitted } makeRoom( in: &counters, at: now ) counters[client] = (now, 1) return .admitted } } // MARK: Methods /// Keeps the store within its capacity before a new client is tracked: expired windows are dropped first, and when the store remains full, the /// oldest live windows are evicted in one batch — a tenth of the capacity — so the sort that finds them runs once per batch of admissions /// instead of once per request while a flood of distinct clients keeps the store full. /// - Parameters: /// - counters: the counters the room is made in. /// - now: the instant the expiry is evaluated against. private func makeRoom( in counters: inout [String: (start: ContinuousClock.Instant, count: Int)], at now: ContinuousClock.Instant ) { guard counters.count >= capacity else { return } counters = counters.filter { now < $0.value.start.advanced(by: window) } let headroom = max(1, capacity / 10) let excess = counters.count - (capacity - headroom) guard excess > 0 else { return } let oldest = counters .sorted { $0.value.start < $1.value.start } .prefix(excess) for counter in oldest { counters.removeValue(forKey: counter.key) } } } } // MARK: - Configuration extension RateLimitMiddleware { /// The limits a ``RateLimitMiddleware`` enforces. public struct Configuration: Sendable { // MARK: Properties /// The number of requests admitted per client per ``window``. public let limit: Int /// Whether a client is keyed by the first `X-Forwarded-For` entry. /// /// Enable it only behind a reverse proxy that sets the header — there, the connection's own address would name the proxy for every visitor, /// sharing one budget across all of them. On a directly reachable server the header is client-supplied, so trusting it lets a client forge fresh keys /// at will. public let trustForwardedFor: Bool /// The length of the fixed window the ``limit`` applies to. public let window: Duration // MARK: Initializers /// Creates a rate-limit configuration. /// - Parameters: /// - limit: the number of requests admitted per client per window. /// - window: the length of the fixed window the limit applies to. /// - trustForwardedFor: whether a client is keyed by the first `X-Forwarded-For` entry. public init( limit: Int = .RateLimit.limit, window: Duration = .seconds(Int.RateLimit.window), trustForwardedFor: Bool = false ) { self.limit = limit self.trustForwardedFor = trustForwardedFor self.window = window } } } // MARK: - String+Constants private extension String { /// The bucket shared by every client the middleware cannot identify. static let unidentified = "unidentified" }