auth.go 7.7 KB

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