Files

84 lines
1.7 KiB
Go
Raw Permalink Normal View History

//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)
}