package repository import ( "context" "crypto/sha256" "encoding/hex" "fmt" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) type RefreshTokenRepo struct { pool *pgxpool.Pool } func NewRefreshTokenRepo(pool *pgxpool.Pool) *RefreshTokenRepo { return &RefreshTokenRepo{pool: pool} } type RefreshToken struct { ID string UserID string TokenHash string ExpiresAt time.Time CreatedAt time.Time RevokedAt *time.Time } func hashToken(token string) string { sum := sha256.Sum256([]byte(token)) return hex.EncodeToString(sum[:]) } func (r *RefreshTokenRepo) Create(ctx context.Context, userID, plainToken string, expiresAt time.Time) error { tokenHash := hashToken(plainToken) _, err := r.pool.Exec(ctx, `INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`, userID, tokenHash, expiresAt, ) if err != nil { return fmt.Errorf("create refresh token: %w", err) } return nil } func (r *RefreshTokenRepo) GetValid(ctx context.Context, plainToken string) (*RefreshToken, error) { tokenHash := hashToken(plainToken) row := r.pool.QueryRow(ctx, `SELECT id, user_id, token_hash, expires_at, created_at, revoked_at FROM refresh_tokens WHERE token_hash = $1 AND expires_at > now() AND revoked_at IS NULL`, tokenHash, ) var rt RefreshToken err := row.Scan(&rt.ID, &rt.UserID, &rt.TokenHash, &rt.ExpiresAt, &rt.CreatedAt, &rt.RevokedAt) if err != nil { if err == pgx.ErrNoRows { return nil, nil } return nil, fmt.Errorf("get refresh token: %w", err) } return &rt, nil } func (r *RefreshTokenRepo) Revoke(ctx context.Context, tokenHash string) error { _, err := r.pool.Exec(ctx, `UPDATE refresh_tokens SET revoked_at = now() WHERE token_hash = $1 AND revoked_at IS NULL`, tokenHash, ) return err } func (r *RefreshTokenRepo) RevokeAllForUser(ctx context.Context, userID string) error { _, err := r.pool.Exec(ctx, `UPDATE refresh_tokens SET revoked_at = now() WHERE user_id = $1 AND revoked_at IS NULL`, userID, ) return err } func (r *RefreshTokenRepo) Delete(ctx context.Context, tokenHash string) error { _, err := r.pool.Exec(ctx, `DELETE FROM refresh_tokens WHERE token_hash = $1`, tokenHash, ) return err } func (r *RefreshTokenRepo) CleanupExpired(ctx context.Context) error { _, err := r.pool.Exec(ctx, `DELETE FROM refresh_tokens WHERE expires_at < now() - interval '1 day'`, ) return err }