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 }