package auth import ( "context" "fmt" "strings" "sync" "time" "github.com/redis/go-redis/v9" "github.com/gochat/gochat/internal/config" ) var compareAndSwapRefreshToken = redis.NewScript(` if redis.call("GET", KEYS[1]) ~= ARGV[1] then return 0 end redis.call("SET", KEYS[1], ARGV[2], "PX", ARGV[3]) return 1 `) // Reference: P2E §1.4 — Refresh Token storage in Redis for rotation tracking // Refresh tokens are stored in Redis with TTL matching their JWT expiry. // This enables: token rotation, revocation, and audit trail. // RefreshTokenStore manages refresh token storage in Redis. type RefreshTokenStore struct { rdb *redis.Client cfg *config.JWTConfig mu sync.RWMutex mem map[string]refreshTokenEntry } type refreshTokenEntry struct { token string expiresAt time.Time } // NewRefreshTokenStore creates a refresh token store backed by Redis. func NewRefreshTokenStore(rdb *redis.Client, cfg *config.JWTConfig) *RefreshTokenStore { return &RefreshTokenStore{rdb: rdb, cfg: cfg, mem: map[string]refreshTokenEntry{}} } // Store saves a refresh token in Redis with TTL. // Key pattern: refresh_token:{user_id}:{token_hash} func (s *RefreshTokenStore) Store(ctx context.Context, userID uint, refreshToken string) error { return s.StoreForClient(ctx, userID, "", refreshToken) } func (s *RefreshTokenStore) StoreForClient(ctx context.Context, userID uint, clientID, refreshToken string) error { key := s.key(userID, clientID) ttl := s.cfg.RefreshExpiryDuration() if s.rdb == nil { s.mu.Lock() defer s.mu.Unlock() s.mem[key] = refreshTokenEntry{token: refreshToken, expiresAt: time.Now().Add(ttl)} return nil } return s.rdb.Set(ctx, key, refreshToken, ttl).Err() } // Validate checks if a refresh token exists and matches the stored value. func (s *RefreshTokenStore) Validate(ctx context.Context, userID uint, refreshToken string) (bool, error) { return s.ValidateForClient(ctx, userID, "", refreshToken) } func (s *RefreshTokenStore) ValidateForClient(ctx context.Context, userID uint, clientID, refreshToken string) (bool, error) { key := s.key(userID, clientID) if s.rdb == nil { s.mu.Lock() defer s.mu.Unlock() stored, ok := s.mem[key] if ok && time.Now().After(stored.expiresAt) { delete(s.mem, key) return false, nil } return ok && stored.token == refreshToken, nil } stored, err := s.rdb.Get(ctx, key).Result() if err == redis.Nil { return false, nil // token not found (expired or revoked) } if err != nil { return false, fmt.Errorf("redis error: %w", err) } return stored == refreshToken, nil } // Revoke removes a refresh token from Redis (logout). func (s *RefreshTokenStore) Revoke(ctx context.Context, userID uint) error { return s.RevokeClient(ctx, userID, "") } func (s *RefreshTokenStore) RevokeClient(ctx context.Context, userID uint, clientID string) error { key := s.key(userID, clientID) if s.rdb == nil { s.mu.Lock() defer s.mu.Unlock() delete(s.mem, key) return nil } return s.rdb.Del(ctx, key).Err() } // RevokeUser removes every legacy and client-scoped refresh token for a user. func (s *RefreshTokenStore) RevokeUser(ctx context.Context, userID uint) error { prefix := s.key(userID, "") if s.rdb == nil { s.mu.Lock() defer s.mu.Unlock() for key := range s.mem { if key == prefix || strings.HasPrefix(key, prefix+":") { delete(s.mem, key) } } return nil } if err := s.rdb.Del(ctx, prefix).Err(); err != nil { return err } var cursor uint64 for { keys, next, err := s.rdb.Scan(ctx, cursor, prefix+":*", 100).Result() if err != nil { return err } if len(keys) > 0 { if err := s.rdb.Del(ctx, keys...).Err(); err != nil { return err } } cursor = next if cursor == 0 { return nil } } } // Rotate replaces an old refresh token with a new one (refresh token rotation). // This ensures each refresh token can only be used once. func (s *RefreshTokenStore) Rotate(ctx context.Context, userID uint, newRefreshToken string) error { return s.Store(ctx, userID, newRefreshToken) } func (s *RefreshTokenStore) RotateForClient(ctx context.Context, userID uint, clientID, newRefreshToken string) error { return s.StoreForClient(ctx, userID, clientID, newRefreshToken) } // CompareAndSwapForClient atomically consumes oldRefreshToken and stores newRefreshToken. func (s *RefreshTokenStore) CompareAndSwapForClient(ctx context.Context, userID uint, clientID, oldRefreshToken, newRefreshToken string) (bool, error) { key := s.key(userID, clientID) ttl := time.Duration(s.cfg.RefreshExpiryHours) * time.Hour if s.rdb == nil { s.mu.Lock() defer s.mu.Unlock() stored, ok := s.mem[key] if !ok || time.Now().After(stored.expiresAt) || stored.token != oldRefreshToken { if ok && time.Now().After(stored.expiresAt) { delete(s.mem, key) } return false, nil } s.mem[key] = refreshTokenEntry{token: newRefreshToken, expiresAt: time.Now().Add(ttl)} return true, nil } swapped, err := compareAndSwapRefreshToken.Run(ctx, s.rdb, []string{key}, oldRefreshToken, newRefreshToken, ttl.Milliseconds()).Int() if err != nil { return false, fmt.Errorf("redis error: %w", err) } return swapped == 1, nil } func (s *RefreshTokenStore) HasClient(ctx context.Context, userID uint, clientID string) (bool, error) { key := s.key(userID, clientID) if s.rdb == nil { s.mu.Lock() defer s.mu.Unlock() entry, ok := s.mem[key] if ok && time.Now().After(entry.expiresAt) { delete(s.mem, key) return false, nil } return ok, nil } n, err := s.rdb.Exists(ctx, key).Result() return n > 0, err } func (s *RefreshTokenStore) key(userID uint, clientID string) string { if clientID == "" { return fmt.Sprintf("gochat:refresh_token:%d", userID) } return fmt.Sprintf("gochat:refresh_token:%d:%s", userID, clientID) }