refresh_tokens.go 2.4 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798
  1. package repository
  2. import (
  3. "context"
  4. "crypto/sha256"
  5. "encoding/hex"
  6. "fmt"
  7. "time"
  8. "github.com/jackc/pgx/v5"
  9. "github.com/jackc/pgx/v5/pgxpool"
  10. )
  11. type RefreshTokenRepo struct {
  12. pool *pgxpool.Pool
  13. }
  14. func NewRefreshTokenRepo(pool *pgxpool.Pool) *RefreshTokenRepo {
  15. return &RefreshTokenRepo{pool: pool}
  16. }
  17. type RefreshToken struct {
  18. ID string
  19. UserID string
  20. TokenHash string
  21. ExpiresAt time.Time
  22. CreatedAt time.Time
  23. RevokedAt *time.Time
  24. }
  25. func hashToken(token string) string {
  26. sum := sha256.Sum256([]byte(token))
  27. return hex.EncodeToString(sum[:])
  28. }
  29. func (r *RefreshTokenRepo) Create(ctx context.Context, userID, plainToken string, expiresAt time.Time) error {
  30. tokenHash := hashToken(plainToken)
  31. _, err := r.pool.Exec(ctx,
  32. `INSERT INTO refresh_tokens (user_id, token_hash, expires_at)
  33. VALUES ($1, $2, $3)`,
  34. userID, tokenHash, expiresAt,
  35. )
  36. if err != nil {
  37. return fmt.Errorf("create refresh token: %w", err)
  38. }
  39. return nil
  40. }
  41. func (r *RefreshTokenRepo) GetValid(ctx context.Context, plainToken string) (*RefreshToken, error) {
  42. tokenHash := hashToken(plainToken)
  43. row := r.pool.QueryRow(ctx,
  44. `SELECT id, user_id, token_hash, expires_at, created_at, revoked_at
  45. FROM refresh_tokens
  46. WHERE token_hash = $1 AND expires_at > now() AND revoked_at IS NULL`,
  47. tokenHash,
  48. )
  49. var rt RefreshToken
  50. err := row.Scan(&rt.ID, &rt.UserID, &rt.TokenHash, &rt.ExpiresAt, &rt.CreatedAt, &rt.RevokedAt)
  51. if err != nil {
  52. if err == pgx.ErrNoRows {
  53. return nil, nil
  54. }
  55. return nil, fmt.Errorf("get refresh token: %w", err)
  56. }
  57. return &rt, nil
  58. }
  59. func (r *RefreshTokenRepo) Revoke(ctx context.Context, tokenHash string) error {
  60. _, err := r.pool.Exec(ctx,
  61. `UPDATE refresh_tokens SET revoked_at = now() WHERE token_hash = $1 AND revoked_at IS NULL`,
  62. tokenHash,
  63. )
  64. return err
  65. }
  66. func (r *RefreshTokenRepo) RevokeAllForUser(ctx context.Context, userID string) error {
  67. _, err := r.pool.Exec(ctx,
  68. `UPDATE refresh_tokens SET revoked_at = now() WHERE user_id = $1 AND revoked_at IS NULL`,
  69. userID,
  70. )
  71. return err
  72. }
  73. func (r *RefreshTokenRepo) Delete(ctx context.Context, tokenHash string) error {
  74. _, err := r.pool.Exec(ctx,
  75. `DELETE FROM refresh_tokens WHERE token_hash = $1`,
  76. tokenHash,
  77. )
  78. return err
  79. }
  80. func (r *RefreshTokenRepo) CleanupExpired(ctx context.Context) error {
  81. _, err := r.pool.Exec(ctx,
  82. `DELETE FROM refresh_tokens WHERE expires_at < now() - interval '1 day'`,
  83. )
  84. return err
  85. }