package httpmw import ( "fmt" "net/http" "sync" "time" "github.com/labstack/echo/v4" "golang.org/x/time/rate" ) type RateLimiterConfig struct { Enabled bool RPS float64 Burst int BucketTTL time.Duration } type limiterEntry struct { limiter *rate.Limiter lastSeen time.Time } type PrincipalRateLimiter struct { enabled bool rps rate.Limit burst int bucketTTL time.Duration mu sync.Mutex entries map[string]*limiterEntry lastSweep time.Time sweepAfter time.Duration } func NewPrincipalRateLimiter(cfg RateLimiterConfig) (*PrincipalRateLimiter, error) { if cfg.BucketTTL <= 0 { cfg.BucketTTL = 5 * time.Minute } if cfg.Enabled { if cfg.RPS <= 0 { return nil, fmt.Errorf("rate limit rps must be positive") } if cfg.Burst <= 0 { return nil, fmt.Errorf("rate limit burst must be positive") } } return &PrincipalRateLimiter{ enabled: cfg.Enabled, rps: rate.Limit(cfg.RPS), burst: cfg.Burst, bucketTTL: cfg.BucketTTL, entries: map[string]*limiterEntry{}, lastSweep: time.Now(), sweepAfter: time.Minute, }, nil } func (rl *PrincipalRateLimiter) Middleware(next echo.HandlerFunc) echo.HandlerFunc { if rl == nil || !rl.enabled { return next } return func(c echo.Context) error { if c.Request().Method == http.MethodOptions || c.Request().URL.Path == "/_health" { return next(c) } key := rl.identifier(c) now := time.Now() if !rl.allow(key, now) { c.Response().Header().Set("Retry-After", "1") return c.JSON(http.StatusTooManyRequests, map[string]string{ "error": "RateLimited", "message": "rate limit exceeded", }) } return next(c) } } func (rl *PrincipalRateLimiter) identifier(c echo.Context) string { if principal, ok := PrincipalFromContext(c); ok { if principal.Subject != "" { return "sub:" + principal.Subject } } if ip := c.RealIP(); ip != "" { return "ip:" + ip } return "ip:unknown" } func (rl *PrincipalRateLimiter) allow(key string, now time.Time) bool { rl.mu.Lock() defer rl.mu.Unlock() if now.Sub(rl.lastSweep) >= rl.sweepAfter { rl.sweepStaleLocked(now) rl.lastSweep = now } entry, ok := rl.entries[key] if !ok { entry = &limiterEntry{limiter: rate.NewLimiter(rl.rps, rl.burst)} rl.entries[key] = entry } entry.lastSeen = now return entry.limiter.Allow() } func (rl *PrincipalRateLimiter) sweepStaleLocked(now time.Time) { for key, entry := range rl.entries { if now.Sub(entry.lastSeen) > rl.bucketTTL { delete(rl.entries, key) } } }