package middleware import ( "net/http" "net/netip" "slices" "strings" "time" "github.com/go-chi/httprate" ) const ( // loginRateLimit is the maximum number of login attempts // per interval. loginRateLimit = 5 // loginRateInterval is the time window for the rate limit. loginRateInterval = 1 * time.Minute // passwordChangeRateLimit is the maximum number of password // change attempts per interval. Each attempt verifies the // current password, so the endpoint must be rate-limited // like any other password-based authentication endpoint. passwordChangeRateLimit = 5 // passwordChangeRateInterval is the time window for the // password change rate limit. passwordChangeRateInterval = 1 * time.Minute // receiverRateInterval is the time window for the webhook // receiver rate limit. The configured limit is expressed in // requests per minute. receiverRateInterval = 1 * time.Minute ) // normalizeAddr strips the IPv4-in-IPv6 wrapper and any zone from // addr so that comparisons and bucket keys are canonical. func normalizeAddr(addr netip.Addr) netip.Addr { return addr.Unmap().WithZone("") } // isTrustedProxy reports whether addr belongs to a network the // operator listed in TRUSTED_PROXIES. The list is empty by default, // so by default nothing is trusted. func (m *Middleware) isTrustedProxy(addr netip.Addr) bool { for _, prefix := range m.params.Config.TrustedProxies { if prefix.Contains(addr) { return true } } return false } // forwardedClientAddr returns the client address named by this // request's forwarded headers. It is consulted only for requests // whose direct peer is a trusted proxy. // // True-Client-IP and X-Real-IP are single-valued, and a trusted // proxy is expected to overwrite whatever the client sent, so they // are taken as given. X-Forwarded-For is a chain the client can // prepend to, so it is walked right to left and the first hop that // is not itself a trusted proxy wins: entries the client inserted // sit to the left of the proxies' own appends and cannot be picked // while the chain is intact. func (m *Middleware) forwardedClientAddr( r *http.Request, ) (netip.Addr, bool) { for _, header := range []string{"True-Client-IP", "X-Real-IP"} { addr, err := netip.ParseAddr( strings.TrimSpace(r.Header.Get(header)), ) if err == nil { return normalizeAddr(addr), true } } hops := strings.Split( strings.Join(r.Header.Values("X-Forwarded-For"), ","), ",", ) for _, hop := range slices.Backward(hops) { addr, err := netip.ParseAddr(strings.TrimSpace(hop)) if err != nil { continue } if addr = normalizeAddr(addr); !m.isTrustedProxy(addr) { return addr, true } } return netip.Addr{}, false } // rateLimitKey is the client identity every rate limiter in this // package buckets on. Forwarded headers are honoured only when the // direct peer (RemoteAddr) is inside the configured trusted-proxy // set; otherwise the peer address itself is the key. Without that // gate any client could mint a fresh bucket per request, or starve // another client's bucket, by picking an X-Forwarded-For value — // which makes every limit here decorative against a deliberate // attacker. func (m *Middleware) rateLimitKey(r *http.Request) (string, error) { return m.clientKey(r), nil } // clientKey computes the bucket key described on rateLimitKey. func (m *Middleware) clientKey(r *http.Request) string { peer, err := netip.ParseAddr(ipFromHostPort(r.RemoteAddr)) if err != nil { // Not an address we can reason about; key on the raw // value rather than collapsing such peers into one // shared bucket. return r.RemoteAddr } peer = normalizeAddr(peer) if !m.isTrustedProxy(peer) { return peer.String() } if addr, ok := m.forwardedClientAddr(r); ok { return addr.String() } return peer.String() } // tooManyRequests returns the 429 handler shared by every limiter: // it logs the rejection with logMessage and answers with // responseMessage. httprate adds the Retry-After header (RFC 6585). func (m *Middleware) tooManyRequests( logMessage, responseMessage string, ) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { m.log.Warn(logMessage, "path", r.URL.Path) http.Error(w, responseMessage, http.StatusTooManyRequests) } } // LoginRateLimit returns middleware that enforces per-IP rate // limiting on login attempts using go-chi/httprate. Only POST // requests are rate-limited; GET requests (rendering the login // form) pass through unaffected. When the rate limit is exceeded, // a 429 Too Many Requests response is returned. Clients are // identified by rateLimitKey. func (m *Middleware) LoginRateLimit() func(http.Handler) http.Handler { return m.postRateLimit( loginRateLimit, loginRateInterval, "login rate limit exceeded", "Too many login attempts. Please try again later.", ) } // PasswordChangeRateLimit returns middleware that enforces // per-IP rate limiting on password change attempts. The change // endpoint verifies the current password, so without a limit a // stolen session could be used to brute-force it; the limit // matches the login endpoint's. func (m *Middleware) PasswordChangeRateLimit() func(http.Handler) http.Handler { return m.postRateLimit( passwordChangeRateLimit, passwordChangeRateInterval, "password change rate limit exceeded", "Too many password change attempts. "+ "Please try again later.", ) } // postRateLimit builds middleware that enforces a per-IP rate // limit on POST requests only; all other methods pass through // unaffected. Requests over the limit receive a 429 with the // given response message, and each rejection is logged with the // given log message. Clients are identified by rateLimitKey. func (m *Middleware) postRateLimit( limit int, interval time.Duration, logMessage, responseMessage string, ) func(http.Handler) http.Handler { limiter := httprate.Limit( limit, interval, httprate.WithKeyFuncs(m.rateLimitKey), httprate.WithLimitHandler( m.tooManyRequests(logMessage, responseMessage), ), ) return func(next http.Handler) http.Handler { limited := limiter(next) return http.HandlerFunc(func( w http.ResponseWriter, r *http.Request, ) { // Only rate-limit POST requests. if r.Method != http.MethodPost { next.ServeHTTP(w, r) return } limited.ServeHTTP(w, r) }) } } // ReceiverRateLimit returns middleware that rate-limits the // public webhook receiver endpoint per client IP per request // path (the path contains the entrypoint UUID, so each sender // is limited per entrypoint without affecting other senders or // other entrypoints). The limit is Config.ReceiverRateLimit // requests per minute. Requests over the limit receive a 429. // Clients are identified by rateLimitKey. func (m *Middleware) ReceiverRateLimit() func(http.Handler) http.Handler { return httprate.Limit( m.params.Config.ReceiverRateLimit, receiverRateInterval, httprate.WithKeyFuncs( m.rateLimitKey, httprate.KeyByEndpoint, ), httprate.WithLimitHandler(m.tooManyRequests( "webhook receiver rate limit exceeded", "Too many requests. Please slow down.", )), ) }