// Package ratelimit counts each client's requests over a minute, an hour // and a day, as the "Counting method" section of SPEC.md describes, and // tells when a request takes a client over a rate limit. The counts are // kept in memory only, for at most 20,000 clients. package ratelimit import ( "net/netip" "sync" "time" "github.com/hashicorp/golang-lru/v2/simplelru" ) // maxClients is how many clients are kept. Past it, the least recently // seen client is dropped, and starts afresh if it comes back. const maxClients = 20000 const day = 24 * time.Hour // Limits are the most requests a client may make in a minute, an hour and // a day. Zero is no limit. type Limits struct { PerMinute int64 PerHour int64 PerDay int64 } // Limiter counts each client's requests against the limits. It is safe // for concurrent use. type Limiter struct { windows [3]window mu sync.Mutex // clients holds each client's buckets, one pair for each of windows, // in the same order. clients *simplelru.LRU[netip.Prefix, *[3]buckets] } // New returns a Limiter for limits, with no client counted yet. func New(limits Limits) *Limiter { clients, err := simplelru.NewLRU[netip.Prefix, *[3]buckets](maxClients, nil) if err != nil { panic(err) // NewLRU fails only for a size below one } return &Limiter{ windows: [3]window{ {name: "minute", length: time.Minute, limit: limits.PerMinute}, {name: "hour", length: time.Hour, limit: limits.PerHour}, {name: "day", length: day, limit: limits.PerDay}, }, clients: clients, } } // Count counts a request from client at now, in every window, whether or // not it is refused. It returns the window whose limit the request takes // the client over, "minute", "hour" or "day", the shortest if it is over // several, or "" if it is within every limit. func (l *Limiter) Count(client netip.Prefix, now time.Time) string { l.mu.Lock() defer l.mu.Unlock() counts, seen := l.clients.Get(client) if !seen { counts = &[3]buckets{} l.clients.Add(client, counts) } limitHit := "" for i, w := range l.windows { requests := counts[i].add(now, w.length) if limitHit == "" && w.limit > 0 && requests > float64(w.limit) { limitHit = w.name } } return limitHit } // window is a length of time over which requests are counted, and the // most requests a client may make in it. type window struct { name string length time.Duration limit int64 } // buckets are a client's two buckets in one window: the requests in the // bucket under way, which began at start, and in the bucket before it. type buckets struct { start time.Time current int64 previous int64 } // add counts a request at now in a window of length, and returns the // client's requests in the window that ends at now: those in the bucket // under way, and those in the bucket before it weighted by how much of // that bucket the window still covers. // // Concurrent requests can be counted out of order, so now can be before // the bucket under way began; such a request is counted in that bucket. func (b *buckets) add(now time.Time, length time.Duration) float64 { start := now.Truncate(length) if start.After(b.start) { if start.Equal(b.start.Add(length)) { b.previous = b.current } else { b.previous = 0 } b.start = start b.current = 0 } b.current++ elapsed := max(now.Sub(b.start), 0) covered := 1 - float64(elapsed)/float64(length) return float64(b.previous)*covered + float64(b.current) }