auth.go 4.6 KB

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