154 lines
4.0 KiB
Go
154 lines
4.0 KiB
Go
package redispkg
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"encoding/hex"
|
|
"log"
|
|
"time"
|
|
|
|
"github.com/redis/go-redis/v9"
|
|
)
|
|
|
|
type R struct{ c *redis.Client }
|
|
|
|
func New(url string) *R {
|
|
if url == "" {
|
|
return &R{}
|
|
}
|
|
opt, err := redis.ParseURL(url)
|
|
if err != nil {
|
|
log.Printf("bad REDIS_URL (%v): redis disabled", err)
|
|
return &R{}
|
|
}
|
|
return &R{c: redis.NewClient(opt)}
|
|
}
|
|
|
|
func (r *R) Get(ctx context.Context, key string) (string, bool) {
|
|
if r.c == nil {
|
|
return "", false
|
|
}
|
|
v, err := r.c.Get(ctx, key).Result()
|
|
if err != nil {
|
|
return "", false // 故障=miss
|
|
}
|
|
return v, true
|
|
}
|
|
|
|
func (r *R) Set(ctx context.Context, key, val string, ttl time.Duration) {
|
|
if r.c == nil {
|
|
return
|
|
}
|
|
if err := r.c.Set(ctx, key, val, ttl).Err(); err != nil {
|
|
log.Printf("redis set %s: %v", key, err)
|
|
}
|
|
}
|
|
|
|
// incrWindowScript atomically increments and sets TTL on first value,
|
|
// preventing the INCR+EXPIRE race that could leave keys without TTL (B1).
|
|
var incrWindowScript = redis.NewScript(`
|
|
local n = redis.call('INCR', KEYS[1])
|
|
if n == 1 then
|
|
redis.call('EXPIRE', KEYS[1], ARGV[1])
|
|
end
|
|
return n
|
|
`)
|
|
|
|
func (r *R) IncrWindow(ctx context.Context, key string, ttl time.Duration) int {
|
|
if r.c == nil {
|
|
return 1
|
|
}
|
|
n, err := incrWindowScript.Run(ctx, r.c, []string{key}, int(ttl.Seconds())).Int()
|
|
if err != nil {
|
|
return 1 // fail-open
|
|
}
|
|
return n
|
|
}
|
|
|
|
func (r *R) Lock(ctx context.Context, key string, ttl time.Duration) (func(), bool) {
|
|
noop := func() {}
|
|
if r.c == nil {
|
|
return noop, true
|
|
}
|
|
b := make([]byte, 8)
|
|
if _, err := rand.Read(b); err != nil {
|
|
// B2: rand failure → degrade to no-lock instead of using a zero token.
|
|
log.Printf("rand.Read failed: %v (proceeding without lock)", err)
|
|
return noop, true
|
|
}
|
|
tok := hex.EncodeToString(b)
|
|
ok, err := r.c.SetNX(ctx, key, tok, ttl).Result()
|
|
if err != nil { // spec §9: Redis 故障降级放行,锁只做尽力去重
|
|
log.Printf("redis lock %s: %v (proceeding without lock)", key, err)
|
|
return noop, true
|
|
}
|
|
if !ok {
|
|
return noop, false // 锁被持有,别的副本在扫
|
|
}
|
|
return func() {
|
|
// B3: use WithoutCancel so unlock survives caller cancellation.
|
|
if err := r.c.Eval(context.WithoutCancel(ctx),
|
|
"if redis.call('get',KEYS[1])==ARGV[1] then return redis.call('del',KEYS[1]) else return 0 end",
|
|
[]string{key}, tok).Err(); err != nil {
|
|
log.Printf("redis unlock %s: %v", key, err)
|
|
}
|
|
}, true
|
|
}
|
|
|
|
// ScanLock acquires a distributed lock with automatic renewal.
|
|
// The lock is renewed every ttl/2 until unlock is called.
|
|
// Returns (unlock, true) on success, (noop, true) on redis failure (degrade),
|
|
// or (noop, false) if the lock is already held.
|
|
func (r *R) ScanLock(ctx context.Context, key string, ttl time.Duration) (func(), bool) {
|
|
noop := func() {}
|
|
if r.c == nil {
|
|
return noop, true
|
|
}
|
|
b := make([]byte, 8)
|
|
if _, err := rand.Read(b); err != nil {
|
|
log.Printf("rand.Read failed: %v (proceeding without lock)", err)
|
|
return noop, true
|
|
}
|
|
tok := hex.EncodeToString(b)
|
|
ok, err := r.c.SetNX(ctx, key, tok, ttl).Result()
|
|
if err != nil {
|
|
log.Printf("redis scanlock %s: %v (proceeding without lock)", key, err)
|
|
return noop, true
|
|
}
|
|
if !ok {
|
|
return noop, false
|
|
}
|
|
|
|
// Start renewal goroutine.
|
|
done := make(chan struct{})
|
|
go func() {
|
|
ticker := time.NewTicker(ttl / 2)
|
|
defer ticker.Stop()
|
|
for {
|
|
select {
|
|
case <-done:
|
|
return
|
|
case <-ticker.C:
|
|
// Renew only if we still own the lock.
|
|
if err := r.c.Eval(context.Background(),
|
|
`if redis.call('get',KEYS[1])==ARGV[1] then
|
|
return redis.call('expire',KEYS[1],ARGV[2])
|
|
else return 0 end`,
|
|
[]string{key}, tok, int(ttl.Seconds())).Err(); err != nil {
|
|
log.Printf("redis scanlock renew %s: %v", key, err)
|
|
}
|
|
}
|
|
}
|
|
}()
|
|
|
|
unlock := func() {
|
|
close(done) // stop renewal
|
|
if err := r.c.Eval(context.WithoutCancel(ctx),
|
|
"if redis.call('get',KEYS[1])==ARGV[1] then return redis.call('del',KEYS[1]) else return 0 end",
|
|
[]string{key}, tok).Err(); err != nil {
|
|
log.Printf("redis scanlock unlock %s: %v", key, err)
|
|
}
|
|
}
|
|
return unlock, true
|
|
}
|