fix: make rate limit increment atomic

This commit is contained in:
0x37
2026-04-21 11:05:06 +01:00
parent 33d6816ac6
commit 72ce8e8cc9
4 changed files with 80 additions and 12 deletions

View File

@@ -70,6 +70,14 @@ else
end
`)
var incrementWithExpireScript = redis.NewScript(`
local count = redis.call("INCR", KEYS[1])
if count == 1 then
redis.call("PEXPIRE", KEYS[1], ARGV[1])
end
return count
`)
func (c *Client) Unlock(ctx context.Context, key string, token string) error {
if c == nil || c.rdb == nil {
return nil
@@ -82,15 +90,10 @@ func (c *Client) IncrementWithExpire(ctx context.Context, key string, expire tim
if c == nil || c.rdb == nil {
return 0, nil
}
count, err := c.rdb.Incr(ctx, key).Result()
if err != nil {
return 0, err
}
if count == 1 {
err = c.rdb.Expire(ctx, key, expire).Err()
if err != nil {
return 0, err
}
}
return count, nil
}
return incrementWithExpireScript.Run(
ctx,
c.rdb,
[]string{key},
expire.Milliseconds(),
).Int64()
}

View File

@@ -0,0 +1,56 @@
package redis
import (
"context"
"testing"
"time"
miniredis "github.com/alicebob/miniredis/v2"
goredis "github.com/redis/go-redis/v9"
)
func TestIncrementWithExpireSetsTTLWithoutExtendingWindow(t *testing.T) {
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("start miniredis: %v", err)
}
defer mr.Close()
client := &Client{
rdb: goredis.NewClient(&goredis.Options{Addr: mr.Addr()}),
}
defer client.Close()
ctx := context.Background()
key := "feedsystem:ratelimit:test"
expire := 30 * time.Second
count, err := client.IncrementWithExpire(ctx, key, expire)
if err != nil {
t.Fatalf("first increment: %v", err)
}
if count != 1 {
t.Fatalf("expected count 1, got %d", count)
}
firstTTL := mr.TTL(key)
if firstTTL <= 0 || firstTTL > expire {
t.Fatalf("expected ttl in (0, %s], got %s", expire, firstTTL)
}
mr.FastForward(5 * time.Second)
ttlBeforeSecond := mr.TTL(key)
count, err = client.IncrementWithExpire(ctx, key, expire)
if err != nil {
t.Fatalf("second increment: %v", err)
}
if count != 2 {
t.Fatalf("expected count 2, got %d", count)
}
ttlAfterSecond := mr.TTL(key)
if ttlAfterSecond != ttlBeforeSecond {
t.Fatalf("expected ttl to stay at %s, got %s", ttlBeforeSecond, ttlAfterSecond)
}
}