auth.go 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220
  1. package services
  2. import (
  3. "context"
  4. "crypto/rand"
  5. "encoding/base64"
  6. "errors"
  7. "fmt"
  8. "time"
  9. "github.com/golang-jwt/jwt/v5"
  10. "github.com/photoplaces/backend/internal/models"
  11. "github.com/photoplaces/backend/internal/repository"
  12. "golang.org/x/crypto/bcrypt"
  13. )
  14. var (
  15. ErrEmailExists = errors.New("email already exists")
  16. ErrInvalidCreds = errors.New("invalid email or password")
  17. ErrUserBanned = errors.New("account is banned")
  18. ErrInvalidToken = errors.New("invalid or expired token")
  19. ErrTokenReused = errors.New("token reused, possible theft")
  20. )
  21. const (
  22. AccessTokenExpiry = 15 * time.Minute
  23. RefreshTokenExpiry = 30 * 24 * time.Hour
  24. RefreshTokenBytes = 32
  25. )
  26. type AuthService struct {
  27. userRepo *repository.UserRepo
  28. refreshTokenRepo *repository.RefreshTokenRepo
  29. jwtSecret []byte
  30. refreshSecret []byte
  31. }
  32. func NewAuthService(userRepo *repository.UserRepo, refreshTokenRepo *repository.RefreshTokenRepo, jwtSecret, refreshSecret string) *AuthService {
  33. return &AuthService{
  34. userRepo: userRepo,
  35. refreshTokenRepo: refreshTokenRepo,
  36. jwtSecret: []byte(jwtSecret),
  37. refreshSecret: []byte(refreshSecret),
  38. }
  39. }
  40. type AuthResult struct {
  41. User *models.User `json:"user"`
  42. AccessToken string `json:"access_token"`
  43. }
  44. type RegisterInput struct {
  45. Email string
  46. Password string
  47. Role string
  48. Name string
  49. }
  50. func (s *AuthService) Register(ctx context.Context, input RegisterInput) (*AuthResult, string, error) {
  51. existing, _ := s.userRepo.GetByEmail(ctx, input.Email)
  52. if existing != nil {
  53. return nil, "", ErrEmailExists
  54. }
  55. hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost)
  56. if err != nil {
  57. return nil, "", fmt.Errorf("hash password: %w", err)
  58. }
  59. user := &models.User{
  60. Email: input.Email,
  61. PasswordHash: string(hash),
  62. Role: input.Role,
  63. Name: strPtr(input.Name),
  64. }
  65. if err := s.userRepo.Create(ctx, user); err != nil {
  66. return nil, "", fmt.Errorf("create user: %w", err)
  67. }
  68. return s.generateTokens(ctx, user)
  69. }
  70. type LoginInput struct {
  71. Email string
  72. Password string
  73. }
  74. func (s *AuthService) Login(ctx context.Context, input LoginInput) (*AuthResult, string, error) {
  75. user, err := s.userRepo.GetByEmail(ctx, input.Email)
  76. if err != nil {
  77. return nil, "", fmt.Errorf("get user: %w", err)
  78. }
  79. if user == nil {
  80. return nil, "", ErrInvalidCreds
  81. }
  82. if user.Status == "banned" {
  83. return nil, "", ErrUserBanned
  84. }
  85. if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(input.Password)); err != nil {
  86. return nil, "", ErrInvalidCreds
  87. }
  88. return s.generateTokens(ctx, user)
  89. }
  90. type TokenClaims struct {
  91. UserID string `json:"user_id"`
  92. Role string `json:"role"`
  93. jwt.RegisteredClaims
  94. }
  95. func (s *AuthService) generateTokens(ctx context.Context, user *models.User) (*AuthResult, string, error) {
  96. now := time.Now()
  97. accessClaims := TokenClaims{
  98. UserID: user.ID,
  99. Role: user.Role,
  100. RegisteredClaims: jwt.RegisteredClaims{
  101. ExpiresAt: jwt.NewNumericDate(now.Add(AccessTokenExpiry)),
  102. IssuedAt: jwt.NewNumericDate(now),
  103. },
  104. }
  105. accessToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims).SignedString(s.jwtSecret)
  106. if err != nil {
  107. return nil, "", fmt.Errorf("sign access token: %w", err)
  108. }
  109. refreshToken, err := generateSecureToken(RefreshTokenBytes)
  110. if err != nil {
  111. return nil, "", fmt.Errorf("generate refresh token: %w", err)
  112. }
  113. expiresAt := now.Add(RefreshTokenExpiry)
  114. if err := s.refreshTokenRepo.Create(ctx, user.ID, refreshToken, expiresAt); err != nil {
  115. return nil, "", fmt.Errorf("store refresh token: %w", err)
  116. }
  117. return &AuthResult{
  118. User: user,
  119. AccessToken: accessToken,
  120. }, refreshToken, nil
  121. }
  122. func (s *AuthService) ValidateAccessToken(tokenString string) (*TokenClaims, error) {
  123. token, err := jwt.ParseWithClaims(tokenString, &TokenClaims{}, func(t *jwt.Token) (interface{}, error) {
  124. return s.jwtSecret, nil
  125. })
  126. if err != nil {
  127. return nil, ErrInvalidToken
  128. }
  129. claims, ok := token.Claims.(*TokenClaims)
  130. if !ok || !token.Valid {
  131. return nil, ErrInvalidToken
  132. }
  133. return claims, nil
  134. }
  135. func (s *AuthService) RefreshSession(ctx context.Context, plainRefreshToken string) (*AuthResult, string, error) {
  136. storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
  137. if err != nil {
  138. return nil, "", fmt.Errorf("lookup refresh token: %w", err)
  139. }
  140. if storedToken == nil {
  141. return nil, "", ErrInvalidToken
  142. }
  143. user, err := s.userRepo.GetByID(ctx, storedToken.UserID)
  144. if err != nil {
  145. return nil, "", fmt.Errorf("get user: %w", err)
  146. }
  147. if user == nil {
  148. return nil, "", ErrInvalidToken
  149. }
  150. if user.Status == "banned" {
  151. return nil, "", ErrUserBanned
  152. }
  153. if err := s.refreshTokenRepo.Delete(ctx, storedToken.TokenHash); err != nil {
  154. return nil, "", fmt.Errorf("delete old refresh token: %w", err)
  155. }
  156. return s.generateTokens(ctx, user)
  157. }
  158. func (s *AuthService) RevokeSession(ctx context.Context, plainRefreshToken string) error {
  159. storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
  160. if err != nil {
  161. return fmt.Errorf("lookup refresh token: %w", err)
  162. }
  163. if storedToken == nil {
  164. return nil
  165. }
  166. return s.refreshTokenRepo.Revoke(ctx, storedToken.TokenHash)
  167. }
  168. func (s *AuthService) RevokeAllSessions(ctx context.Context, userID string) error {
  169. return s.refreshTokenRepo.RevokeAllForUser(ctx, userID)
  170. }
  171. func (s *AuthService) CleanupExpiredTokens(ctx context.Context) error {
  172. return s.refreshTokenRepo.CleanupExpired(ctx)
  173. }
  174. func generateSecureToken(n int) (string, error) {
  175. b := make([]byte, n)
  176. if _, err := rand.Read(b); err != nil {
  177. return "", err
  178. }
  179. return base64.URLEncoding.EncodeToString(b), nil
  180. }
  181. func strPtr(s string) *string {
  182. if s == "" {
  183. return nil
  184. }
  185. return &s
  186. }