Files
gochat/backend/internal/auth/refresh_store.go
T

129 lines
3.9 KiB
Go

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)
}