package auth import ( "context" "fmt" "sync" "time" "github.com/redis/go-redis/v9" "github.com/gochat/gochat/internal/config" ) // 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 := time.Duration(s.cfg.RefreshExpiryHours) * time.Hour 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() } // 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) } 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) }