From c7fd4a6dbc3d21cd8c3b0150ae810cd4d890ebea Mon Sep 17 00:00:00 2001 From: AdamMagued Date: Fri, 9 Oct 2026 12:12:41 +0300 Subject: [PATCH] feat(middleware): support rate limit metadata and response headers Signed-off-by: AdamMagued --- middleware/rate_limiter.go | 122 +++++++++++-- middleware/rate_limiter_context_test.go | 233 +++++++++++++++++++++++- 2 files changed, 333 insertions(+), 22 deletions(-) diff --git a/middleware/rate_limiter.go b/middleware/rate_limiter.go index bd70e10a9..e39afa331 100644 --- a/middleware/rate_limiter.go +++ b/middleware/rate_limiter.go @@ -15,12 +15,25 @@ import ( "golang.org/x/time/rate" ) -// Rate limit response headers set by stores that implement RateLimiterStoreContext. +// Rate limit response headers set when supported by the store. const ( HeaderXRateLimitLimit = "X-RateLimit-Limit" HeaderXRateLimitRemaining = "X-RateLimit-Remaining" + HeaderXRateLimitReset = "X-RateLimit-Reset" ) +// RateLimitMetadataContextKey is the default context key for storing RateLimitMetadata in *echo.Context. +const RateLimitMetadataContextKey = "rate_limit_metadata" + +// RateLimitMetadata contains rate limiting metadata for a visitor. +type RateLimitMetadata struct { + Limit int // Limit is the maximum number of requests allowed (burst/limit). + Remaining int // Remaining is the number of requests remaining in the current window. + Reset time.Duration // Reset is the duration until the rate limit resets to full capacity. + ResetTime time.Time // ResetTime is the time when the rate limit resets to full capacity. + RetryAfter time.Duration // RetryAfter is the duration until the next request is allowed (zero when allowed). +} + // RateLimiterStore is the interface to be implemented by custom stores. type RateLimiterStore interface { Allow(identifier string) (bool, error) @@ -34,6 +47,14 @@ type RateLimiterStoreContext interface { AllowContext(c *echo.Context, identifier string) (bool, error) } +// RateLimiterStoreWithDetails is an optional interface a RateLimiterStore may implement. +// When the configured store implements it, the rate limiter calls AllowWithDetails +// instead of Allow, and sets the standard rate limit headers (X-RateLimit-Limit, +// X-RateLimit-Remaining, X-RateLimit-Reset, and Retry-After) on the response. +type RateLimiterStoreWithDetails interface { + AllowWithDetails(identifier string) (bool, RateLimitMetadata, error) +} + // RateLimiterConfig defines the configuration for the rate limiter type RateLimiterConfig struct { Skipper Skipper @@ -46,6 +67,9 @@ type RateLimiterConfig struct { ErrorHandler func(c *echo.Context, err error) error // DenyHandler provides a handler to be called when RateLimiter denies access DenyHandler func(c *echo.Context, identifier string, err error) error + // ContextKey defines the key used to store RateLimitMetadata in *echo.Context. + // Optional. If not set, metadata is not stored in context. + ContextKey string } // Extractor is used to extract data from *echo.Context @@ -153,7 +177,14 @@ func (config RateLimiterConfig) ToMiddleware() (echo.MiddlewareFunc, error) { var allow bool var allowErr error - if sc, ok := config.Store.(RateLimiterStoreContext); ok { + if sd, ok := config.Store.(RateLimiterStoreWithDetails); ok { + var meta RateLimitMetadata + allow, meta, allowErr = sd.AllowWithDetails(identifier) + setRateLimitHeaders(c, meta, allow) + if config.ContextKey != "" { + c.Set(config.ContextKey, meta) + } + } else if sc, ok := config.Store.(RateLimiterStoreContext); ok { allow, allowErr = sc.AllowContext(c, identifier) } else { allow, allowErr = config.Store.Allow(identifier) @@ -254,19 +285,25 @@ var DefaultRateLimiterMemoryStoreConfig = RateLimiterMemoryStoreConfig{ // Allow implements RateLimiterStore.Allow func (store *RateLimiterMemoryStore) Allow(identifier string) (bool, error) { - _, allowed := store.allow(identifier) - return allowed, nil + allowed, _, err := store.allowWithDetails(identifier) + return allowed, err } // AllowContext implements RateLimiterStoreContext: it makes the allow/deny decision // and sets the X-RateLimit-* (and Retry-After when denied) response headers. func (store *RateLimiterMemoryStore) AllowContext(c *echo.Context, identifier string) (bool, error) { - limiter, allowed := store.allow(identifier) - store.setRateLimitHeaders(c, limiter, allowed) - return allowed, nil + allowed, meta, err := store.allowWithDetails(identifier) + setRateLimitHeaders(c, meta, allowed) + return allowed, err +} + +// AllowWithDetails implements RateLimiterStoreWithDetails: it makes the allow/deny decision +// and returns rate limit metadata (limit, remaining, reset, and retry-after). +func (store *RateLimiterMemoryStore) AllowWithDetails(identifier string) (bool, RateLimitMetadata, error) { + return store.allowWithDetails(identifier) } -func (store *RateLimiterMemoryStore) allow(identifier string) (*rate.Limiter, bool) { +func (store *RateLimiterMemoryStore) allowWithDetails(identifier string) (bool, RateLimitMetadata, error) { store.mutex.Lock() defer store.mutex.Unlock() @@ -281,22 +318,71 @@ func (store *RateLimiterMemoryStore) allow(identifier string) (*rate.Limiter, bo if now.Sub(store.lastCleanup) > store.expiresIn { store.cleanupStaleVisitors(now) } - return limiter.Limiter, limiter.AllowN(now, 1) -} -func (store *RateLimiterMemoryStore) setRateLimitHeaders(c *echo.Context, limiter *rate.Limiter, allowed bool) { - header := c.Response().Header() - header.Set(HeaderXRateLimitLimit, strconv.Itoa(store.burst)) + allowed := limiter.AllowN(now, 1) remaining := max(int(math.Floor(limiter.Tokens())), 0) - header.Set(HeaderXRateLimitRemaining, strconv.Itoa(remaining)) + + var retryAfter time.Duration + if !allowed { + reservation := limiter.ReserveN(now, 1) + if reservation.OK() { + if delay := reservation.DelayFrom(now); delay > 0 { + retryAfter = delay + } + reservation.CancelAt(now) + } + } + + var resetDuration time.Duration + if store.rate > 0 { + tokens := limiter.Tokens() + if tokens < 0 { + tokens = 0 + } + if tokens > float64(store.burst) { + tokens = float64(store.burst) + } + missing := float64(store.burst) - tokens + if missing > 0 { + resetDuration = time.Duration((missing / store.rate) * float64(time.Second)) + } + } + + resetTime := now.Add(resetDuration) + + meta := RateLimitMetadata{ + Limit: store.burst, + Remaining: remaining, + Reset: resetDuration, + ResetTime: resetTime, + RetryAfter: retryAfter, + } + + return allowed, meta, nil +} + +func setRateLimitHeaders(c *echo.Context, meta RateLimitMetadata, allowed bool) { + header := c.Response().Header() + header.Set(HeaderXRateLimitLimit, strconv.Itoa(meta.Limit)) + header.Set(HeaderXRateLimitRemaining, strconv.Itoa(meta.Remaining)) + + resetTime := meta.ResetTime + if resetTime.IsZero() { + if meta.Reset > 0 { + resetTime = time.Now().Add(meta.Reset) + } else { + resetTime = time.Now() + } + } + header.Set(HeaderXRateLimitReset, strconv.FormatInt(resetTime.Unix(), 10)) if !allowed { - reservation := limiter.ReserveN(store.timeNow(), 1) - if delay := reservation.Delay(); delay > 0 { - header.Set(echo.HeaderRetryAfter, strconv.Itoa(int(math.Ceil(delay.Seconds())))) + delaySec := int(math.Ceil(meta.RetryAfter.Seconds())) + if delaySec <= 0 { + delaySec = 1 } - reservation.Cancel() + header.Set(echo.HeaderRetryAfter, strconv.Itoa(delaySec)) } } diff --git a/middleware/rate_limiter_context_test.go b/middleware/rate_limiter_context_test.go index 629c01e47..9b3e1a857 100644 --- a/middleware/rate_limiter_context_test.go +++ b/middleware/rate_limiter_context_test.go @@ -8,9 +8,11 @@ import ( "net/http/httptest" "strconv" "testing" + "time" "github.com/labstack/echo/v5" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // ctxAwareStore implements both Allow and the optional AllowContext. AllowContext @@ -33,6 +35,29 @@ func (s *ctxAwareStore) AllowContext(c *echo.Context, identifier string) (bool, return s.allow, nil } +type detailsAwareStore struct { + allowCalled bool + ctxAllowCalled bool + detailsAllowCalled bool + allow bool + metadata RateLimitMetadata +} + +func (s *detailsAwareStore) Allow(identifier string) (bool, error) { + s.allowCalled = true + return s.allow, nil +} + +func (s *detailsAwareStore) AllowContext(c *echo.Context, identifier string) (bool, error) { + s.ctxAllowCalled = true + return s.allow, nil +} + +func (s *detailsAwareStore) AllowWithDetails(identifier string) (bool, RateLimitMetadata, error) { + s.detailsAllowCalled = true + return s.allow, s.metadata, nil +} + // When the store implements AllowContext, the middleware must call it instead of // Allow, so the store can set rate-limit headers on the response. func TestRateLimiter_storeAllowContextIsPreferred(t *testing.T) { @@ -54,8 +79,156 @@ func TestRateLimiter_storeAllowContextIsPreferred(t *testing.T) { assert.Equal(t, "42", rec.Header().Get("Retry-After"), "store should be able to set headers via the context") } -// The built-in memory store implements AllowContext, so it sets X-RateLimit-Limit / -// X-RateLimit-Remaining on every request and Retry-After when the limit is hit (#2961). +// When the store implements RateLimiterStoreWithDetails, the middleware must call +// AllowWithDetails and automatically emit standard response headers. +func TestRateLimiter_storeWithDetailsIsPreferred(t *testing.T) { + fixedTime := time.Date(2026, time.October, 9, 12, 0, 0, 0, time.UTC) + meta := RateLimitMetadata{ + Limit: 100, + Remaining: 95, + Reset: 30 * time.Second, + ResetTime: fixedTime.Add(30 * time.Second), + RetryAfter: 0, + } + store := &detailsAwareStore{ + allow: true, + metadata: meta, + } + + e := echo.New() + mw := RateLimiterWithConfig(RateLimiterConfig{ + Store: store, + IdentifierExtractor: func(c *echo.Context) (string, error) { return "id", nil }, + }) + handler := mw(func(c *echo.Context) error { return c.String(http.StatusOK, "ok") }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + assert.NoError(t, handler(c)) + assert.True(t, store.detailsAllowCalled) + assert.False(t, store.ctxAllowCalled) + assert.False(t, store.allowCalled) + + assert.Equal(t, "100", rec.Header().Get(HeaderXRateLimitLimit)) + assert.Equal(t, "95", rec.Header().Get(HeaderXRateLimitRemaining)) + assert.Equal(t, strconv.FormatInt(fixedTime.Add(30*time.Second).Unix(), 10), rec.Header().Get(HeaderXRateLimitReset)) + assert.Empty(t, rec.Header().Get(echo.HeaderRetryAfter)) +} + +func TestRateLimiter_storeWithDetails_DenyEmitsRetryAfter(t *testing.T) { + fixedTime := time.Date(2026, time.October, 9, 12, 0, 0, 0, time.UTC) + meta := RateLimitMetadata{ + Limit: 10, + Remaining: 0, + Reset: 5 * time.Second, + ResetTime: fixedTime.Add(5 * time.Second), + RetryAfter: 3 * time.Second, + } + store := &detailsAwareStore{ + allow: false, + metadata: meta, + } + + e := echo.New() + mw := RateLimiterWithConfig(RateLimiterConfig{ + Store: store, + IdentifierExtractor: func(c *echo.Context) (string, error) { return "id", nil }, + }) + handler := mw(func(c *echo.Context) error { return c.String(http.StatusOK, "ok") }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + err := handler(c) + assert.Error(t, err) + + assert.Equal(t, "10", rec.Header().Get(HeaderXRateLimitLimit)) + assert.Equal(t, "0", rec.Header().Get(HeaderXRateLimitRemaining)) + assert.Equal(t, strconv.FormatInt(fixedTime.Add(5*time.Second).Unix(), 10), rec.Header().Get(HeaderXRateLimitReset)) + assert.Equal(t, "3", rec.Header().Get(echo.HeaderRetryAfter)) +} + +func TestRateLimiter_storeWithDetails_ContextStorage(t *testing.T) { + meta := RateLimitMetadata{ + Limit: 50, + Remaining: 42, + Reset: 15 * time.Second, + RetryAfter: 0, + } + store := &detailsAwareStore{ + allow: true, + metadata: meta, + } + + e := echo.New() + var extractedMeta RateLimitMetadata + var foundInHandler bool + + mw := RateLimiterWithConfig(RateLimiterConfig{ + Store: store, + IdentifierExtractor: func(c *echo.Context) (string, error) { return "id", nil }, + ContextKey: "rate_limit", + }) + handler := mw(func(c *echo.Context) error { + val := c.Get("rate_limit") + if v, ok := val.(RateLimitMetadata); ok { + foundInHandler = true + extractedMeta = v + } + return c.String(http.StatusOK, "ok") + }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + assert.NoError(t, handler(c)) + assert.True(t, foundInHandler) + assert.Equal(t, 50, extractedMeta.Limit) + assert.Equal(t, 42, extractedMeta.Remaining) +} + +func TestRateLimiter_storeWithDetails_ZeroResetTimeFallback(t *testing.T) { + before := time.Now().Unix() + meta := RateLimitMetadata{ + Limit: 20, + Remaining: 5, + Reset: 10 * time.Second, + RetryAfter: 0, + } + store := &detailsAwareStore{ + allow: true, + metadata: meta, + } + + e := echo.New() + mw := RateLimiterWithConfig(RateLimiterConfig{ + Store: store, + IdentifierExtractor: func(c *echo.Context) (string, error) { return "id", nil }, + }) + handler := mw(func(c *echo.Context) error { return c.String(http.StatusOK, "ok") }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + assert.NoError(t, handler(c)) + after := time.Now().Unix() + + resetHeader := rec.Header().Get(HeaderXRateLimitReset) + require.NotEmpty(t, resetHeader) + resetVal, err := strconv.ParseInt(resetHeader, 10, 64) + assert.NoError(t, err) + assert.GreaterOrEqual(t, resetVal, before+10) + assert.LessOrEqual(t, resetVal, after+10) +} + +// The built-in memory store implements AllowContext and AllowWithDetails, so it sets +// X-RateLimit-Limit / X-RateLimit-Remaining / X-RateLimit-Reset on every request and +// Retry-After when the limit is hit (#2961). func TestRateLimiterMemoryStore_AllowContextSetsHeaders(t *testing.T) { store := NewRateLimiterMemoryStoreWithConfig(RateLimiterMemoryStoreConfig{Rate: 1, Burst: 3}) e := echo.New() @@ -72,18 +245,70 @@ func TestRateLimiterMemoryStore_AllowContextSetsHeaders(t *testing.T) { return rec } - // Burst of 3: each allowed request advertises the limit and decreasing remaining. + // Burst of 3: each allowed request advertises the limit, decreasing remaining, and reset timestamp. for i := 0; i < 3; i++ { rec := do() assert.Equal(t, http.StatusOK, rec.Code) assert.Equal(t, "3", rec.Header().Get(HeaderXRateLimitLimit)) assert.Equal(t, strconv.Itoa(2-i), rec.Header().Get(HeaderXRateLimitRemaining)) + assert.NotEmpty(t, rec.Header().Get(HeaderXRateLimitReset)) assert.Empty(t, rec.Header().Get(echo.HeaderRetryAfter)) } - // 4th request is denied: 429, remaining 0, and a Retry-After hint. + // 4th request is denied: 429, remaining 0, reset timestamp, and a Retry-After hint. rec := do() assert.Equal(t, http.StatusTooManyRequests, rec.Code) assert.Equal(t, "0", rec.Header().Get(HeaderXRateLimitRemaining)) + assert.NotEmpty(t, rec.Header().Get(HeaderXRateLimitReset)) assert.NotEmpty(t, rec.Header().Get(echo.HeaderRetryAfter)) } + +func TestRateLimiterMemoryStore_AllowWithDetails(t *testing.T) { + baseTime := time.Date(2026, time.October, 9, 12, 0, 0, 0, time.UTC) + store := NewRateLimiterMemoryStoreWithConfig(RateLimiterMemoryStoreConfig{ + Rate: 2, + Burst: 4, + ExpiresIn: 1 * time.Minute, + }) + store.timeNow = func() time.Time { + return baseTime + } + + // Request 1: 1 token consumed, 3 remaining + allowed, meta, err := store.AllowWithDetails("user1") + assert.NoError(t, err) + assert.True(t, allowed) + assert.Equal(t, 4, meta.Limit) + assert.Equal(t, 3, meta.Remaining) + assert.Equal(t, 500*time.Millisecond, meta.Reset) + assert.Equal(t, baseTime.Add(500*time.Millisecond), meta.ResetTime) + assert.Equal(t, time.Duration(0), meta.RetryAfter) + + // Request 2: 2 remaining + allowed, meta, err = store.AllowWithDetails("user1") + assert.NoError(t, err) + assert.True(t, allowed) + assert.Equal(t, 2, meta.Remaining) + assert.Equal(t, 1000*time.Millisecond, meta.Reset) + + // Request 3: 1 remaining + allowed, meta, err = store.AllowWithDetails("user1") + assert.NoError(t, err) + assert.True(t, allowed) + assert.Equal(t, 1, meta.Remaining) + + // Request 4: 0 remaining + allowed, meta, err = store.AllowWithDetails("user1") + assert.NoError(t, err) + assert.True(t, allowed) + assert.Equal(t, 0, meta.Remaining) + assert.Equal(t, 2000*time.Millisecond, meta.Reset) + + // Request 5: limit exceeded + allowed, meta, err = store.AllowWithDetails("user1") + assert.NoError(t, err) + assert.False(t, allowed) + assert.Equal(t, 0, meta.Remaining) + assert.Equal(t, 500*time.Millisecond, meta.RetryAfter) + assert.Equal(t, 2000*time.Millisecond, meta.Reset) +}