Go: add context to redis client (#17689)

Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
Jin Hai
2026-08-02 22:50:54 +08:00
committed by GitHub
parent b2521ebf51
commit 266837eb33
31 changed files with 302 additions and 308 deletions

View File

@@ -106,7 +106,7 @@ const (
)
// Init InitRedis initializes Redis client
func Init() error {
func Init(ctx context.Context) error {
var initErr error
once.Do(func() {
globalConfig := server.GetConfig()
@@ -124,10 +124,10 @@ func Init() error {
})
// Test connection
ctx, cancel := context.WithTimeout(context.Background(), server.DefaultConnectTimeout)
redisCtx, cancel := context.WithTimeout(ctx, server.DefaultConnectTimeout)
defer cancel()
if err := client.Ping(ctx).Err(); err != nil {
if err := client.Ping(redisCtx).Err(); err != nil {
initErr = fmt.Errorf("failed to connect to Redis: %w", err)
return
}
@@ -167,11 +167,10 @@ func IsEnabled() bool {
}
// Health checks if Redis is healthy
func (r *Client) Health() bool {
func (r *Client) Health(ctx context.Context) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.Ping(ctx).Err(); err != nil {
return false
}
@@ -190,11 +189,10 @@ func (r *Client) Health() bool {
}
// Info returns Redis server information
func (r *Client) Info() map[string]interface{} {
func (r *Client) Info(ctx context.Context) map[string]interface{} {
if r.client == nil {
return nil
}
ctx := context.Background()
infoStr, err := r.client.Info(ctx).Result()
if err != nil {
common.Warn("Failed to get Redis info", zap.Error(err))
@@ -278,11 +276,10 @@ func (r *Client) IsAlive() bool {
}
// Exist checks if key exists
func (r *Client) Exist(key string) (bool, error) {
func (r *Client) Exist(ctx context.Context, key string) (bool, error) {
if r.client == nil {
return false, nil
}
ctx := context.Background()
exists, err := r.client.Exists(ctx, key).Result()
if err != nil {
common.Warn("Redis Exist error", zap.String("key", key), zap.Error(err))
@@ -292,11 +289,10 @@ func (r *Client) Exist(key string) (bool, error) {
}
// Get gets value by key
func (r *Client) Get(key string) (string, error) {
func (r *Client) Get(ctx context.Context, key string) (string, error) {
if r.client == nil {
return "", nil
}
ctx := context.Background()
val, err := r.client.Get(ctx, key).Result()
if err == redis.Nil {
return "", nil
@@ -309,11 +305,10 @@ func (r *Client) Get(key string) (string, error) {
}
// SetObj sets object with JSON serialization
func (r *Client) SetObj(key string, obj interface{}, exp time.Duration) bool {
func (r *Client) SetObj(ctx context.Context, key string, obj interface{}, exp time.Duration) bool {
if r.client == nil {
return false
}
ctx := context.Background()
data, err := json.Marshal(obj)
if err != nil {
common.Warn("Redis SetObj marshal error", zap.String("key", key), zap.Error(err))
@@ -326,12 +321,11 @@ func (r *Client) SetObj(key string, obj interface{}, exp time.Duration) bool {
return true
}
// GetObj gets and unmarshals object from Redis
func (r *Client) GetObj(key string, dest interface{}) bool {
// GetObj gets and unmarshal object from Redis
func (r *Client) GetObj(ctx context.Context, key string, dest interface{}) bool {
if r.client == nil {
return false
}
ctx := context.Background()
data, err := r.client.Get(ctx, key).Result()
if err == redis.Nil {
return false
@@ -348,11 +342,10 @@ func (r *Client) GetObj(key string, dest interface{}) bool {
}
// Set sets value with expiration
func (r *Client) Set(key string, value string, exp time.Duration) bool {
func (r *Client) Set(ctx context.Context, key string, value string, exp time.Duration) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.Set(ctx, key, value, exp).Err(); err != nil {
common.Warn("Redis Set error", zap.String("key", key), zap.Error(err))
return false
@@ -361,11 +354,10 @@ func (r *Client) Set(key string, value string, exp time.Duration) bool {
}
// SetNX sets value only if key does not exist
func (r *Client) SetNX(key string, value string, exp time.Duration) bool {
func (r *Client) SetNX(ctx context.Context, key string, value string, exp time.Duration) bool {
if r.client == nil {
return false
}
ctx := context.Background()
ok, err := r.client.SetNX(ctx, key, value, exp).Result()
if err != nil {
common.Warn("Redis SetNX error", zap.String("key", key), zap.Error(err))
@@ -376,11 +368,10 @@ func (r *Client) SetNX(key string, value string, exp time.Duration) bool {
// GetOrCreateKey atomically retrieves an existing key or creates a new one
// Uses Redis SETNX command to ensure atomicity across multiple goroutines/processes
func (r *Client) GetOrCreateKey(key string, value string) (string, error) {
func (r *Client) GetOrCreateKey(ctx context.Context, key string, value string) (string, error) {
if r.client == nil {
return "", nil
}
ctx := context.Background()
// First, try to get the existing key
existingKey, err := r.client.Get(ctx, key).Result()
if err == nil {
@@ -412,11 +403,10 @@ func (r *Client) GetOrCreateKey(key string, value string) (string, error) {
}
// SAdd adds member to set
func (r *Client) SAdd(key string, member string) bool {
func (r *Client) SAdd(ctx context.Context, key string, member string) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.SAdd(ctx, key, member).Err(); err != nil {
common.Warn("Redis SAdd error", zap.String("key", key), zap.Error(err))
return false
@@ -425,11 +415,10 @@ func (r *Client) SAdd(key string, member string) bool {
}
// SRem removes member from set
func (r *Client) SRem(key string, member string) bool {
func (r *Client) SRem(ctx context.Context, key string, member string) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.SRem(ctx, key, member).Err(); err != nil {
common.Warn("Redis SRem error", zap.String("key", key), zap.Error(err))
return false
@@ -438,11 +427,10 @@ func (r *Client) SRem(key string, member string) bool {
}
// SMembers returns all members of a set
func (r *Client) SMembers(key string) ([]string, error) {
func (r *Client) SMembers(ctx context.Context, key string) ([]string, error) {
if r.client == nil {
return nil, nil
}
ctx := context.Background()
members, err := r.client.SMembers(ctx, key).Result()
if err != nil {
common.Warn("Redis SMembers error", zap.String("key", key), zap.Error(err))
@@ -452,11 +440,10 @@ func (r *Client) SMembers(key string) ([]string, error) {
}
// SIsMember checks if member exists in set
func (r *Client) SIsMember(key string, member string) bool {
func (r *Client) SIsMember(ctx context.Context, key string, member string) bool {
if r.client == nil {
return false
}
ctx := context.Background()
ok, err := r.client.SIsMember(ctx, key, member).Result()
if err != nil {
common.Warn("Redis SIsMember error", zap.String("key", key), zap.Error(err))
@@ -466,11 +453,10 @@ func (r *Client) SIsMember(key string, member string) bool {
}
// ZAdd adds member with score to sorted set
func (r *Client) ZAdd(key string, member string, score float64) bool {
func (r *Client) ZAdd(ctx context.Context, key string, member string, score float64) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.ZAdd(ctx, key, redis.Z{Score: score, Member: member}).Err(); err != nil {
common.Warn("Redis ZAdd error", zap.String("key", key), zap.Error(err))
return false
@@ -479,11 +465,10 @@ func (r *Client) ZAdd(key string, member string, score float64) bool {
}
// ZCount returns count of members with score in range
func (r *Client) ZCount(key string, min, max float64) int64 {
func (r *Client) ZCount(ctx context.Context, key string, min, max float64) int64 {
if r.client == nil {
return 0
}
ctx := context.Background()
count, err := r.client.ZCount(ctx, key, fmt.Sprintf("%f", min), fmt.Sprintf("%f", max)).Result()
if err != nil {
common.Warn("Redis ZCount error", zap.String("key", key), zap.Error(err))
@@ -493,11 +478,10 @@ func (r *Client) ZCount(key string, min, max float64) int64 {
}
// ZPopMin pops minimum score members from sorted set
func (r *Client) ZPopMin(key string, count int) ([]redis.Z, error) {
func (r *Client) ZPopMin(ctx context.Context, key string, count int) ([]redis.Z, error) {
if r.client == nil {
return nil, nil
}
ctx := context.Background()
members, err := r.client.ZPopMin(ctx, key, int64(count)).Result()
if err != nil {
common.Warn("Redis ZPopMin error", zap.String("key", key), zap.Error(err))
@@ -507,11 +491,10 @@ func (r *Client) ZPopMin(key string, count int) ([]redis.Z, error) {
}
// ZRangeByScore returns members with score in range
func (r *Client) ZRangeByScore(key string, min, max float64) ([]string, error) {
func (r *Client) ZRangeByScore(ctx context.Context, key string, min, max float64) ([]string, error) {
if r.client == nil {
return nil, nil
}
ctx := context.Background()
members, err := r.client.ZRangeByScore(ctx, key, &redis.ZRangeBy{
Min: fmt.Sprintf("%f", min),
Max: fmt.Sprintf("%f", max),
@@ -524,11 +507,10 @@ func (r *Client) ZRangeByScore(key string, min, max float64) ([]string, error) {
}
// ZRemRangeByScore removes members with score in range
func (r *Client) ZRemRangeByScore(key string, min, max float64) int64 {
func (r *Client) ZRemRangeByScore(ctx context.Context, key string, min, max float64) int64 {
if r.client == nil {
return 0
}
ctx := context.Background()
count, err := r.client.ZRemRangeByScore(ctx, key, fmt.Sprintf("%f", min), fmt.Sprintf("%f", max)).Result()
if err != nil {
common.Warn("Redis ZRemRangeByScore error", zap.String("key", key), zap.Error(err))
@@ -538,11 +520,10 @@ func (r *Client) ZRemRangeByScore(key string, min, max float64) int64 {
}
// IncrBy increments key by increment
func (r *Client) IncrBy(key string, increment int64) (int64, error) {
func (r *Client) IncrBy(ctx context.Context, key string, increment int64) (int64, error) {
if r.client == nil {
return 0, nil
}
ctx := context.Background()
val, err := r.client.IncrBy(ctx, key, increment).Result()
if err != nil {
common.Warn("Redis IncrBy error", zap.String("key", key), zap.Error(err))
@@ -552,11 +533,10 @@ func (r *Client) IncrBy(key string, increment int64) (int64, error) {
}
// DecrBy decrements key by decrement
func (r *Client) DecrBy(key string, decrement int64) (int64, error) {
func (r *Client) DecrBy(ctx context.Context, key string, decrement int64) (int64, error) {
if r.client == nil {
return 0, nil
}
ctx := context.Background()
val, err := r.client.DecrBy(ctx, key, decrement).Result()
if err != nil {
common.Warn("Redis DecrBy error", zap.String("key", key), zap.Error(err))
@@ -566,7 +546,7 @@ func (r *Client) DecrBy(key string, decrement int64) (int64, error) {
}
// GenerateAutoIncrementID generates auto-increment ID
func (r *Client) GenerateAutoIncrementID(keyPrefix string, namespace string, increment int64, ensureMinimum *int64) int64 {
func (r *Client) GenerateAutoIncrementID(ctx context.Context, keyPrefix string, namespace string, increment int64, ensureMinimum *int64) int64 {
if r.client == nil {
return -1
}
@@ -581,7 +561,6 @@ func (r *Client) GenerateAutoIncrementID(keyPrefix string, namespace string, inc
}
redisKey := fmt.Sprintf("%s:%s", keyPrefix, namespace)
ctx := context.Background()
// Check if key exists
exists, err := r.client.Exists(ctx, redisKey).Result()
@@ -616,11 +595,10 @@ func (r *Client) GenerateAutoIncrementID(keyPrefix string, namespace string, inc
}
// Transaction sets key with NX flag (transaction-like behavior)
func (r *Client) Transaction(key string, value string, exp time.Duration) bool {
func (r *Client) Transaction(ctx context.Context, key string, value string, exp time.Duration) bool {
if r.client == nil {
return false
}
ctx := context.Background()
pipe := r.client.Pipeline()
pipe.SetNX(ctx, key, value, exp)
_, err := pipe.Exec(ctx)
@@ -632,11 +610,10 @@ func (r *Client) Transaction(key string, value string, exp time.Duration) bool {
}
// QueueProduct produces a message to Redis Stream
func (r *Client) QueueProduct(queue string, message interface{}) bool {
func (r *Client) QueueProduct(ctx context.Context, queue string, message interface{}) bool {
if r.client == nil {
return false
}
ctx := context.Background()
for i := 0; i < 3; i++ {
data, err := json.Marshal(message)
@@ -659,11 +636,10 @@ func (r *Client) QueueProduct(queue string, message interface{}) bool {
}
// QueueConsumer consumes a message from Redis Stream
func (r *Client) QueueConsumer(queueName, groupName, consumerName string, msgID string) (*Message, error) {
func (r *Client) QueueConsumer(ctx context.Context, queueName, groupName, consumerName string, msgID string) (*Message, error) {
if r.client == nil {
return nil, nil
}
ctx := context.Background()
for i := 0; i < 3; i++ {
// Create consumer group if not exists
@@ -730,11 +706,10 @@ func (r *Client) QueueConsumer(queueName, groupName, consumerName string, msgID
}
// Ack acknowledges the message
func (m *Message) Ack() bool {
func (m *Message) Ack(ctx context.Context) bool {
if m.consumer == nil {
return false
}
ctx := context.Background()
err := m.consumer.XAck(ctx, m.queueName, m.groupName, m.msgID).Err()
if err != nil {
common.Warn("Message Ack error", zap.Error(err))
@@ -754,11 +729,11 @@ func (m *Message) GetMsgID() string {
}
// GetPendingMsg gets pending messages
func (r *Client) GetPendingMsg(queue, groupName string) ([]redis.XPendingExt, error) {
func (r *Client) GetPendingMsg(ctx context.Context, queue, groupName string) ([]redis.XPendingExt, error) {
if r.client == nil {
return nil, nil
}
ctx := context.Background()
messages, err := r.client.XPendingExt(ctx, &redis.XPendingExtArgs{
Stream: queue,
Group: groupName,
@@ -776,11 +751,10 @@ func (r *Client) GetPendingMsg(queue, groupName string) ([]redis.XPendingExt, er
}
// RequeueMsg re-enqueues a message
func (r *Client) RequeueMsg(queue, groupName, msgID string) {
func (r *Client) RequeueMsg(ctx context.Context, queue, groupName, msgID string) {
if r.client == nil {
return
}
ctx := context.Background()
for i := 0; i < 3; i++ {
msgs, err := r.client.XRange(ctx, queue, msgID, msgID).Result()
@@ -803,11 +777,10 @@ func (r *Client) RequeueMsg(queue, groupName, msgID string) {
}
// QueueInfo returns queue group info
func (r *Client) QueueInfo(queue, groupName string) (map[string]interface{}, error) {
func (r *Client) QueueInfo(ctx context.Context, queue, groupName string) (map[string]interface{}, error) {
if r.client == nil {
return nil, nil
}
ctx := context.Background()
for i := 0; i < 3; i++ {
groups, err := r.client.XInfoGroups(ctx, queue).Result()
@@ -833,11 +806,10 @@ func (r *Client) QueueInfo(queue, groupName string) (map[string]interface{}, err
}
// DeleteIfEqual deletes key if its value equals expected value (atomic)
func (r *Client) DeleteIfEqual(key, expectedValue string) bool {
func (r *Client) DeleteIfEqual(ctx context.Context, key, expectedValue string) bool {
if r.client == nil {
return false
}
ctx := context.Background()
result, err := r.luaDeleteIfEqual.Run(ctx, r.client, []string{key}, expectedValue).Result()
if err != nil {
common.Warn("Redis DeleteIfEqual error", zap.Error(err))
@@ -847,11 +819,10 @@ func (r *Client) DeleteIfEqual(key, expectedValue string) bool {
}
// Delete deletes a key
func (r *Client) Delete(key string) bool {
func (r *Client) Delete(ctx context.Context, key string) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.Del(ctx, key).Err(); err != nil {
common.Warn("Redis Delete error", zap.String("key", key), zap.Error(err))
return false
@@ -860,11 +831,10 @@ func (r *Client) Delete(key string) bool {
}
// Expire sets expiration on a key
func (r *Client) Expire(key string, exp time.Duration) bool {
func (r *Client) Expire(ctx context.Context, key string, exp time.Duration) bool {
if r.client == nil {
return false
}
ctx := context.Background()
if err := r.client.Expire(ctx, key, exp).Err(); err != nil {
common.Warn("Redis Expire error", zap.String("key", key), zap.Error(err))
return false
@@ -873,11 +843,10 @@ func (r *Client) Expire(key string, exp time.Duration) bool {
}
// TTL gets remaining time to live of a key
func (r *Client) TTL(key string) time.Duration {
func (r *Client) TTL(ctx context.Context, key string) time.Duration {
if r.client == nil {
return -2
}
ctx := context.Background()
ttl, err := r.client.TTL(ctx, key).Result()
if err != nil {
common.Warn("Redis TTL error", zap.String("key", key), zap.Error(err))
@@ -913,13 +882,13 @@ func NewDistributedLock(lockKey string, lockValue string, timeout time.Duration,
}
// Acquire acquires the lock
func (l *DistributedLock) Acquire() bool {
func (l *DistributedLock) Acquire(ctx context.Context) bool {
if l.client == nil {
return false
}
// Delete if stale
l.client.DeleteIfEqual(l.lockKey, l.lockValue)
return l.client.SetNX(l.lockKey, l.lockValue, l.timeout)
l.client.DeleteIfEqual(ctx, l.lockKey, l.lockValue)
return l.client.SetNX(ctx, l.lockKey, l.lockValue, l.timeout)
}
// SpinAcquire keeps trying to acquire the lock
@@ -929,8 +898,8 @@ func (l *DistributedLock) SpinAcquire(ctx context.Context) error {
case <-ctx.Done():
return ctx.Err()
default:
l.client.DeleteIfEqual(l.lockKey, l.lockValue)
if l.client.SetNX(l.lockKey, l.lockValue, l.timeout) {
l.client.DeleteIfEqual(ctx, l.lockKey, l.lockValue)
if l.client.SetNX(ctx, l.lockKey, l.lockValue, l.timeout) {
return nil
}
time.Sleep(10 * time.Second)
@@ -939,11 +908,11 @@ func (l *DistributedLock) SpinAcquire(ctx context.Context) error {
}
// Release releases the lock
func (l *DistributedLock) Release() bool {
func (l *DistributedLock) Release(ctx context.Context) bool {
if l.client == nil {
return false
}
return l.client.DeleteIfEqual(l.lockKey, l.lockValue)
return l.client.DeleteIfEqual(ctx, l.lockKey, l.lockValue)
}
// TokenBucket token bucket rate limiter
@@ -968,11 +937,10 @@ func NewTokenBucket(key string, capacity, rate float64) *TokenBucket {
}
// Allow checks if request is allowed
func (tb *TokenBucket) Allow(cost float64) (bool, float64) {
func (tb *TokenBucket) Allow(ctx context.Context, cost float64) (bool, float64) {
if tb.client == nil || tb.client.client == nil {
return true, 0
}
ctx := context.Background()
now := float64(time.Now().Unix())
result, err := tb.client.luaTokenBucket.Run(ctx, tb.client.client, []string{tb.key},

View File

@@ -52,7 +52,7 @@ func newStrictTestClient(t *testing.T) (*Client, *miniredis.Miniredis) {
// third should be denied. This is the happy-path security gate.
func TestEvalTokenBucketStrict_AllowedThenDenied(t *testing.T) {
r, _ := newStrictTestClient(t)
ctx := context.Background()
ctx := t.Context()
for i := 1; i <= 2; i++ {
ok, err := r.EvalTokenBucketStrict(ctx, "tb:webhook", 2, 0.1)
@@ -80,7 +80,7 @@ func TestEvalTokenBucketStrict_RedisDownFailsClosed(t *testing.T) {
r, mr := newStrictTestClient(t)
mr.Close() // break the connection
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
ctx, cancel := context.WithTimeout(t.Context(), 500*time.Millisecond)
defer cancel()
ok, err := r.EvalTokenBucketStrict(ctx, "tb:webhook", 5, 1)
@@ -97,7 +97,7 @@ func TestEvalTokenBucketStrict_RedisDownFailsClosed(t *testing.T) {
// (false, error) so the webhook handler can surface 102.
func TestEvalTokenBucketStrict_NilClient(t *testing.T) {
var r *Client
ok, err := r.EvalTokenBucketStrict(context.Background(), "tb:webhook", 1, 1)
ok, err := r.EvalTokenBucketStrict(t.Context(), "tb:webhook", 1, 1)
if err == nil {
t.Fatalf("expected error on nil client, got nil")
}