package ratelimit import ( "testing" "time" "github.com/alicebob/miniredis/v2" "github.com/redis/go-redis/v9" "golang.org/x/time/rate" ) func newTestRedis(t *testing.T) (*miniredis.Miniredis, *redisStore) { t.Helper() mr := miniredis.RunT(t) rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) return mr, newRedisStore(rdb) } func TestRedisLimiterBurstThenDeny(t *testing.T) { _, s := newTestRedis(t) l := s.newLimiter(rate.Limit(1.0/60.0), 3) // 每分钟 1 个,突发 3 for i := 0; i < 3; i++ { res, err := l.allow("ip:1.1.1.1") if err != nil { t.Fatalf("allow 报错:%v", err) } if !res.Allowed { t.Fatalf("burst 内第 %d 个应放行", i+1) } } res, err := l.allow("ip:1.1.1.1") if err != nil { t.Fatalf("allow 报错:%v", err) } if res.Allowed { t.Fatal("超出 burst 应拒绝") } if res.RetryAfter <= 0 { t.Fatal("拒绝时应给正的 RetryAfter") } // 不同 key 独立 if res, _ := l.allow("ip:2.2.2.2"); !res.Allowed { t.Fatal("不同 key 应独立计") } } func TestRedisFailLockerLockExpireReset(t *testing.T) { mr, s := newTestRedis(t) fl := &redisFailLocker{rdb: s.rdb} const key = "S001|admin" for i := 0; i < 3; i++ { if locked, _ := fl.locked(key); locked { t.Fatalf("第 %d 次失败前不应锁定", i+1) } if err := fl.recordFailure(key, 3, time.Minute); err != nil { t.Fatalf("recordFailure 报错:%v", err) } } if locked, _ := fl.locked(key); !locked { t.Fatal("达阈值应锁定") } // 锁定即清零计数:fail 键应已删除 failKey, _ := failKeys(key) if mr.Exists(failKey) { t.Fatal("锁定后 fail 计数键应被清零删除") } // TTL 到期自动解锁 mr.FastForward(61 * time.Second) if locked, _ := fl.locked(key); locked { t.Fatal("锁定到期应自动解锁") } // Reset 清空计数 _ = fl.recordFailure(key, 3, time.Minute) if err := fl.reset(key); err != nil { t.Fatalf("reset 报错:%v", err) } if mr.Exists(failKey) { t.Fatal("Reset 后 fail 键应删除") } } func TestRedisCounterIncrTTL(t *testing.T) { mr, s := newTestRedis(t) c := &redisCounter{rdb: s.rdb} n1, err := c.Incr("dq:test:x", time.Hour) if err != nil || n1 != 1 { t.Fatalf("首次应为 1,得到 %d err=%v", n1, err) } n2, _ := c.Incr("dq:test:x", time.Hour) if n2 != 2 { t.Fatalf("第二次应为 2,得到 %d", n2) } mr.FastForward(time.Hour + time.Second) n3, _ := c.Incr("dq:test:x", time.Hour) if n3 != 1 { t.Fatalf("TTL 过期后应重新从 1 计,得到 %d", n3) } }