auth.go 8.0 KB

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