package middleware import ( "net/http" "sync" "golang.org/x/time/rate" ) // RateLimiter 基于调用方标识的令牌桶限流。 type RateLimiter struct { limiters map[string]*rate.Limiter rate rate.Limit burst int mu sync.Mutex } func NewRateLimiter(r rate.Limit, burst int) *RateLimiter { return &RateLimiter{ limiters: make(map[string]*rate.Limiter), rate: r, burst: burst, } } func (rl *RateLimiter) Middleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { subject := "anonymous" if sub, ok := r.Context().Value(KeySubject).(string); ok { subject = sub } limiter := rl.getLimiter(subject) if !limiter.Allow() { http.Error(w, "rate limit exceeded", http.StatusTooManyRequests) return } next.ServeHTTP(w, r) }) } func (rl *RateLimiter) getLimiter(key string) *rate.Limiter { rl.mu.Lock() defer rl.mu.Unlock() limiter, ok := rl.limiters[key] if !ok { limiter = rate.NewLimiter(rl.rate, rl.burst) rl.limiters[key] = limiter } return limiter }