84 lines
1.7 KiB
Go
84 lines
1.7 KiB
Go
//go:build linux || freebsd
|
|
// +build linux freebsd
|
|
|
|
package transport
|
|
|
|
import (
|
|
"io"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
// TokenBucket implements the token bucket algorithm
|
|
type TokenBucket struct {
|
|
mu sync.Mutex
|
|
capacity int64
|
|
tokens float64
|
|
refillRate float64
|
|
lastRefill time.Time
|
|
}
|
|
|
|
// NewTokenBucket creates new rate limiter
|
|
func NewTokenBucket(maxBytesPerSec int64, burstSize int64) *TokenBucket {
|
|
if maxBytesPerSec <= 0 {
|
|
return nil
|
|
}
|
|
return &TokenBucket{
|
|
capacity: burstSize,
|
|
tokens: float64(burstSize),
|
|
refillRate: float64(maxBytesPerSec),
|
|
lastRefill: time.Now(),
|
|
}
|
|
}
|
|
|
|
// Allow waits for enough tokens
|
|
func (tb *TokenBucket) Allow(n int64) {
|
|
if tb == nil {
|
|
return
|
|
}
|
|
tb.mu.Lock()
|
|
defer tb.mu.Unlock()
|
|
|
|
now := time.Now()
|
|
elapsed := now.Sub(tb.lastRefill).Seconds()
|
|
tb.tokens += elapsed * tb.refillRate
|
|
if tb.tokens > float64(tb.capacity) {
|
|
tb.tokens = float64(tb.capacity)
|
|
}
|
|
tb.lastRefill = now
|
|
|
|
for tb.tokens < float64(n) {
|
|
needed := float64(n) - tb.tokens
|
|
waitTime := time.Duration(needed/tb.refillRate*1e9) * time.Nanosecond
|
|
tb.mu.Unlock()
|
|
time.Sleep(waitTime)
|
|
tb.mu.Lock()
|
|
now := time.Now()
|
|
elapsed := now.Sub(tb.lastRefill).Seconds()
|
|
tb.tokens += elapsed * tb.refillRate
|
|
if tb.tokens > float64(tb.capacity) {
|
|
tb.tokens = float64(tb.capacity)
|
|
}
|
|
tb.lastRefill = now
|
|
}
|
|
tb.tokens -= float64(n)
|
|
}
|
|
|
|
// WrapReader returns a reader with rate limiting
|
|
func (tb *TokenBucket) WrapReader(r io.Reader) io.Reader {
|
|
if tb == nil {
|
|
return r
|
|
}
|
|
return &limitedReader{r: r, limiter: tb}
|
|
}
|
|
|
|
type limitedReader struct {
|
|
r io.Reader
|
|
limiter *TokenBucket
|
|
}
|
|
|
|
func (lr *limitedReader) Read(p []byte) (int, error) {
|
|
lr.limiter.Allow(int64(len(p)))
|
|
return lr.r.Read(p)
|
|
}
|