小Cの已经记不起来的博客

从零实现一个简易的限流

先吐槽一下为什么又要自己写

前几天线上服务被一个不知道哪里来的脚本疯狂刷接口,一开始想着简单点,直接 Nginx 上 limit_req 就完事了。结果配置一贴上去,发现它对单 IP 突发流量的控制还是有点粗暴,而且我们业务里真正想限的其实是 user_id,不是 IP。IP 这东西吧,前面套了 CDN、代理、网关之后,很容易一锅端,误伤自己人。

那就只能自己搞个简易的限流了。

为什么网上很多资料一打开全是 Redis + Lua、漏桶、令牌桶、滑动窗口、分布式限流、API 网关设计,真正想抄一个“最小可运行版本”的时候反而找不到呢?可能是我眼神不好吧...

下面这个版本不是拿来上生产的最终答案,就是一个很简单的起步实现。先把它看懂,后面要不要上 Redis、要不要加监控、要不要做多租户,都再说。

最简单:固定窗口计数,能跑但不太优雅

最省事的思路就是:每个 key 在一个固定时间窗口里允许访问 N 次。比如每秒 10 次。

type FixedWindowLimiter struct {
	mu     sync.Mutex
	limit  int
	window time.Duration
	hits   map[string]*countWindow
}

type countWindow struct {
	count int
	start time.Time
}

func NewFixedWindowLimiter(limit int, window time.Duration) *FixedWindowLimiter {
	return &FixedWindowLimiter{
		limit:  limit,
		window: window,
		hits:   make(map[string]*countWindow),
	}
}

func (l *FixedWindowLimiter) Allow(key string) bool {
	l.mu.Lock()
	defer l.mu.Unlock()

	now := time.Now()

	w, ok := l.hits[key]
	if !ok || now.Sub(w.start) > l.window {
		l.hits[key] = &countWindow{
			count: 1,
			start: now,
		}
		return true
	}

	w.count++
	if w.count > l.limit {
		return false
	}

	return true
}

这个写法很直观。一个 map[string]*countWindow,key 可以是 IP、用户 ID、接口路径,随便你。

但它有个老问题,就是临界突发。假设每秒允许 100 次,第一秒最后 100ms 来 100 个请求,第二秒刚开始 100ms 又进来 100 个请求,窗口一换,它又放行了。外部看起来就是短时间内被刷了 200 次。

所以如果只是练手、挡挡小脚本,固定窗口能用。稍微正式一点,还是得换个算法。

换成令牌桶,手感会好很多

令牌桶的思路也很简单:桶里最多有 capacity 个令牌,系统每隔一段时间往里面补令牌,每秒生成 rate 个。请求进来就从桶里拿一个令牌,拿不到就限流。

这个算法比固定窗口更平滑,也允许一定程度的突发。

type TokenBucket struct {
	mu       sync.Mutex
	capacity float64
	rate     float64
	tokens   float64
	last     time.Time
}

func NewTokenBucket(capacity int, ratePerSecond float64) *TokenBucket {
	return &TokenBucket{
		capacity: float64(capacity),
		rate:     ratePerSecond,
		tokens:   float64(capacity),
		last:     time.Now(),
	}
}

func (b *TokenBucket) Take(n float64) bool {
	b.mu.Lock()
	defer b.mu.Unlock()

	now := time.Now()
	elapsed := now.Sub(b.last).Seconds()
	b.last = now

	b.tokens += elapsed * b.rate
	if b.tokens > b.capacity {
		b.tokens = b.capacity
	}

	if b.tokens < n {
		return false
	}

	b.tokens -= n
	return true
}

比如容量 20,每秒补 5 个。那刚开始可以瞬间处理 20 个请求,之后平均每秒大概能过 5 个。突发允许一点,但不会一直爽。

这个版本已经比固定窗口更像回事了。

包成 HTTP 中间件,直接能用

限流这东西,写成一个类还不方便用。最自然的还是塞进 HTTP 中间件里。

下面这个例子用 Go 的 net/http,按 X-User-ID 限流,如果没有用户 ID 就按 IP 限。

func RateLimit(next http.Handler) http.Handler {
	var buckets sync.Map

	return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		key := "ip:" + r.RemoteAddr
		if uid := r.Header.Get("X-User-ID"); uid != "" {
			key = "user:" + uid
		}

		v, _ := buckets.LoadOrStore(key, NewTokenBucket(20, 5))
		bucket := v.(*TokenBucket)

		if !bucket.Take(1) {
			w.Header().Set("Retry-After", "1")
			http.Error(w, "too many requests", http.StatusTooManyRequests)
			return
		}

		next.ServeHTTP(w, r)
	})
}

然后注册路由的时候直接套上:

mux := http.NewServeMux()
mux.HandleFunc("/api/hello", func(w http.ResponseWriter, r *http.Request) {
	_, _ = fmt.Fprintln(w, "hello")
})

http.ListenAndServe(":8080", RateLimit(mux))

不出问题的话就没有问题了。你拿 curl 连打 25 次,前面 20 次会通,后面几次应该就能看到 429 Too Many Requests

for i in {1..25}; do curl -i http://localhost:8080/api/hello; done

几个很容易踩的小坑

第一个坑是 key 别随便选。

User-Agent 限流看起来花里胡哨,但 UA 可以被随便伪造,而且浏览器 UA 就那么几种,很容易变成全局限流。按完整 URL path 限流也可能炸,因为你 path 里如果带参数、带文件路径、带用户自定义 slug,buckets 里面会越来越多。

比较稳的 key 一般是:用户 ID、设备 ID、登录后的租户 ID、IP + 接口路径。

第二个坑是 sync.Map 不回收,时间一长 key 会越来越多。

如果只是单机小流量,问题不大。要是你按 IP 限流,而且来路很杂,那还是得加过期清理。简单点可以后台 goroutine 定期扫,记录 lastAccess,超过一定时间就 delete。复杂点就上 freecachegroupcache、或者 Redis。

第三个坑是 time.Now() 不要放得太随意。

如果你每秒都锁一次,其实还好。但如果是高并发,锁粒度要心里有数。上面的令牌桶每个 bucket 一把锁,问题不算大。但如果 key 特别多,锁竞争、内存占用、过期清理,都会开始变烦。

单机限流够不够?

如果你的服务只有一个实例,那上面的实现基本能用。

但如果是多实例,比如三台机器跑同一个服务,那每台机器都有自己的 TokenBucket,总限流量其实是 单机限流量 * 实例数。这个结果很迷惑,你以为限了每秒 100,实际上可能是每秒 300。

真到这一步,就得上共享存储了。最常见就是 Redis。

Redis 里存:

limit:user:{user_id}:tokens
limit:user:{user_id}:last

然后拿 Lua 脚本做“读旧令牌、补新令牌、扣减、写回”这一套操作,避免多实例并发写乱掉。

大概会长成这样,思路对,细节还要补:

local key = KEYS[1]
local capacity = tonumber(ARGV[1])
local rate = tonumber(ARGV[2])
local now = tonumber(ARGV[3])
local take = tonumber(ARGV[4])

local bucket = redis.call('HMGET', key, 'tokens', 'last')
local tokens = tonumber(bucket[1])
local last = tonumber(bucket[2])

if tokens == nil then
    tokens = capacity
    last = now
end

local elapsed = now - last
tokens = math.min(capacity, tokens + elapsed * rate)

if tokens >= take then
    tokens = tokens - take
    redis.call('HMSET', key, 'tokens', tokens, 'last', now)
    redis.call('EXPIRE', key, 60)
    return 1
else
    redis.call('HMSET', key, 'tokens', tokens, 'last', now)
    redis.call('EXPIRE', key, 60)
    return 0
end

不过说实话,如果你一开始只是想挡住无脑脚本、限制一下用户频繁调用,真没必要上来就把架构拉满。先本地令牌桶跑起来,把 key、错误码、日志、监控这些基础东西想清楚,比直接抄一个分布式限流框架更重要。

最后说点实在的

限流这件事,难的不是那几十行代码,而是几个问题你得提前想好:

谁会被限?是 IP、用户、接口、租户,还是三者组合?

限了之后给什么提示?能不能告诉客户端什么时候重试?

限流误杀了怎么快速豁免?有没有配置开关、白名单、热更新?

有没有地方能看见当前谁在被限?别等用户投诉了才发现某个接口早就 429 满天飞。

这些问题想明白了,代码其实很薄。想不明白,就算上了网关、上了 Redis、上了 Lua,最后还是会变成“为什么这个用户被限死了”,然后开始临时加白名单。

反正我这种手搓版本就图一个能跑、好看懂、改起来方便。后面业务真长大了,再换成网关限流、Sentinel、Hystrix、Resilience4j、Redis 分布式限流都行,至少基础原理已经踩过了。

评论

还没有评论。

发表评论

提交后评论将经过自动审核,审核通过后公开展示。

未在播放