fix: make rate limit increment atomic
This commit is contained in:
@@ -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()
|
||||
}
|
||||
|
||||
56
backend/internal/middleware/redis/redis_test.go
Normal file
56
backend/internal/middleware/redis/redis_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user