60 lines
1.0 KiB
Go
60 lines
1.0 KiB
Go
package message
|
|
|
|
import (
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
// rateLimiter 是每端一个令牌桶:速率 rps、容量 burst。
|
|
// rps<=0 表示不限速。
|
|
type rateLimiter struct {
|
|
rps float64
|
|
burst float64
|
|
|
|
mu sync.Mutex
|
|
m map[string]*tokenBucket
|
|
}
|
|
|
|
type tokenBucket struct {
|
|
tokens float64
|
|
last time.Time
|
|
}
|
|
|
|
func newRateLimiter(rps float64, burst int) *rateLimiter {
|
|
if burst <= 0 {
|
|
burst = defaultRequestBurst
|
|
}
|
|
return &rateLimiter{
|
|
rps: rps,
|
|
burst: float64(burst),
|
|
m: make(map[string]*tokenBucket),
|
|
}
|
|
}
|
|
|
|
// allow 消耗 1 个令牌;允许则 true。
|
|
func (r *rateLimiter) allow(endpointID string, now time.Time) bool {
|
|
if r == nil || r.rps <= 0 {
|
|
return true
|
|
}
|
|
r.mu.Lock()
|
|
defer r.mu.Unlock()
|
|
b := r.m[endpointID]
|
|
if b == nil {
|
|
b = &tokenBucket{tokens: r.burst, last: now}
|
|
r.m[endpointID] = b
|
|
}
|
|
elapsed := now.Sub(b.last).Seconds()
|
|
if elapsed > 0 {
|
|
b.tokens += elapsed * r.rps
|
|
if b.tokens > r.burst {
|
|
b.tokens = r.burst
|
|
}
|
|
b.last = now
|
|
}
|
|
if b.tokens < 1 {
|
|
return false
|
|
}
|
|
b.tokens--
|
|
return true
|
|
}
|