auth.go 9.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287
  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. "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/repository"
  15. "golang.org/x/crypto/bcrypt"
  16. )
  17. var (
  18. ErrEmailExists = errors.New("email already exists")
  19. ErrInvalidCreds = errors.New("invalid email or password")
  20. ErrUserBanned = errors.New("account is banned")
  21. ErrInvalidToken = errors.New("invalid or expired token")
  22. ErrTokenReused = errors.New("token reused, possible theft")
  23. )
  24. const (
  25. AccessTokenExpiry = 15 * time.Minute
  26. RefreshTokenExpiry = 30 * 24 * time.Hour
  27. RefreshTokenBytes = 32
  28. )
  29. type UserRepo interface {
  30. Create(ctx context.Context, user *models.User) error
  31. GetByID(ctx context.Context, id string) (*models.User, error)
  32. GetByEmail(ctx context.Context, email string) (*models.User, error)
  33. Update(ctx context.Context, user *models.User) error
  34. UpdateRole(ctx context.Context, id, role string) error
  35. UpdateStatus(ctx context.Context, id, status string) error
  36. List(ctx context.Context, filter models.UserFilter) ([]*models.User, error)
  37. }
  38. type RefreshTokenRepo interface {
  39. Create(ctx context.Context, userID, plainToken string, expiresAt time.Time) error
  40. GetValid(ctx context.Context, plainToken string) (*models.RefreshToken, error)
  41. GetRevoked(ctx context.Context, plainToken string) (*models.RefreshToken, error)
  42. Revoke(ctx context.Context, tokenHash string) error
  43. RevokeAllForUser(ctx context.Context, userID string) error
  44. Delete(ctx context.Context, tokenHash string) error
  45. ReplaceIfExists(ctx context.Context, tokenHash string) (bool, error)
  46. CleanupExpired(ctx context.Context) error
  47. }
  48. type AuthService struct {
  49. userRepo UserRepo
  50. refreshTokenRepo RefreshTokenRepo
  51. jwtSecret []byte
  52. refreshSecret []byte
  53. }
  54. func NewAuthService(userRepo UserRepo, refreshTokenRepo RefreshTokenRepo, jwtSecret, refreshSecret string) *AuthService {
  55. return &AuthService{
  56. userRepo: userRepo,
  57. refreshTokenRepo: refreshTokenRepo,
  58. jwtSecret: []byte(jwtSecret),
  59. refreshSecret: []byte(refreshSecret),
  60. }
  61. }
  62. type AuthResult struct {
  63. User *models.User `json:"user"`
  64. AccessToken string `json:"access_token"`
  65. }
  66. type RegisterInput struct {
  67. Email string
  68. Password string
  69. Role string
  70. Name string
  71. }
  72. func (s *AuthService) Register(ctx context.Context, input RegisterInput) (*AuthResult, string, error) {
  73. logger := log.FromContext(ctx)
  74. hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost)
  75. if err != nil {
  76. logger.ErrorContext(ctx, "password hashing failed", log.WithError(err))
  77. return nil, "", fmt.Errorf("hash password: %w", err)
  78. }
  79. user := &models.User{
  80. Email: input.Email,
  81. PasswordHash: string(hash),
  82. Role: input.Role,
  83. Name: pointer.Str(input.Name),
  84. }
  85. if err := s.userRepo.Create(ctx, user); err != nil {
  86. if errors.Is(err, repository.ErrEmailExists) {
  87. logger.WarnContext(ctx, "registration attempt with existing email", slog.String("email", input.Email))
  88. return nil, "", ErrEmailExists
  89. }
  90. logger.ErrorContext(ctx, "user creation failed", log.WithError(err), slog.String("email", input.Email))
  91. return nil, "", fmt.Errorf("create user: %w", err)
  92. }
  93. logger.InfoContext(ctx, "user registered", slog.String("user_id", user.ID), slog.String("email", user.Email), slog.String("role", user.Role))
  94. return s.generateTokens(ctx, user)
  95. }
  96. type LoginInput struct {
  97. Email string
  98. Password string
  99. }
  100. func (s *AuthService) Login(ctx context.Context, input LoginInput) (*AuthResult, string, error) {
  101. logger := log.FromContext(ctx)
  102. user, err := s.userRepo.GetByEmail(ctx, input.Email)
  103. if err != nil {
  104. logger.ErrorContext(ctx, "user lookup failed", log.WithError(err), slog.String("email", input.Email))
  105. return nil, "", fmt.Errorf("get user: %w", err)
  106. }
  107. if user == nil {
  108. logger.WarnContext(ctx, "login attempt for non-existent user", slog.String("email", input.Email))
  109. return nil, "", ErrInvalidCreds
  110. }
  111. if user.Status == "banned" {
  112. logger.WarnContext(ctx, "login attempt for banned user", slog.String("user_id", user.ID))
  113. return nil, "", ErrUserBanned
  114. }
  115. if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(input.Password)); err != nil {
  116. logger.WarnContext(ctx, "invalid password attempt", slog.String("user_id", user.ID))
  117. return nil, "", ErrInvalidCreds
  118. }
  119. logger.InfoContext(ctx, "user logged in", slog.String("user_id", user.ID), slog.String("role", user.Role))
  120. return s.generateTokens(ctx, user)
  121. }
  122. type TokenClaims struct {
  123. UserID string `json:"user_id"`
  124. Role string `json:"role"`
  125. jwt.RegisteredClaims
  126. }
  127. func (s *AuthService) generateTokens(ctx context.Context, user *models.User) (*AuthResult, string, error) {
  128. now := time.Now()
  129. accessClaims := TokenClaims{
  130. UserID: user.ID,
  131. Role: user.Role,
  132. RegisteredClaims: jwt.RegisteredClaims{
  133. ExpiresAt: jwt.NewNumericDate(now.Add(AccessTokenExpiry)),
  134. IssuedAt: jwt.NewNumericDate(now),
  135. },
  136. }
  137. accessToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims).SignedString(s.jwtSecret)
  138. if err != nil {
  139. return nil, "", fmt.Errorf("sign access token: %w", err)
  140. }
  141. refreshToken, err := generateSecureToken(RefreshTokenBytes)
  142. if err != nil {
  143. return nil, "", fmt.Errorf("generate refresh token: %w", err)
  144. }
  145. expiresAt := now.Add(RefreshTokenExpiry)
  146. if err := s.refreshTokenRepo.Create(ctx, user.ID, refreshToken, expiresAt); err != nil {
  147. return nil, "", fmt.Errorf("store refresh token: %w", err)
  148. }
  149. return &AuthResult{
  150. User: user,
  151. AccessToken: accessToken,
  152. }, refreshToken, nil
  153. }
  154. func (s *AuthService) ValidateAccessToken(tokenString string) (*TokenClaims, error) {
  155. token, err := jwt.ParseWithClaims(tokenString, &TokenClaims{}, func(t *jwt.Token) (interface{}, error) {
  156. return s.jwtSecret, nil
  157. })
  158. if err != nil {
  159. return nil, ErrInvalidToken
  160. }
  161. claims, ok := token.Claims.(*TokenClaims)
  162. if !ok || !token.Valid {
  163. return nil, ErrInvalidToken
  164. }
  165. return claims, nil
  166. }
  167. func (s *AuthService) RefreshSession(ctx context.Context, plainRefreshToken string) (*AuthResult, string, error) {
  168. logger := log.FromContext(ctx)
  169. storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
  170. if err != nil {
  171. logger.ErrorContext(ctx, "refresh token lookup failed", log.WithError(err))
  172. return nil, "", fmt.Errorf("lookup refresh token: %w", err)
  173. }
  174. if storedToken == nil {
  175. revokedToken, err := s.refreshTokenRepo.GetRevoked(ctx, plainRefreshToken)
  176. if err != nil {
  177. logger.ErrorContext(ctx, "revoked token lookup failed", log.WithError(err))
  178. return nil, "", fmt.Errorf("lookup revoked token: %w", err)
  179. }
  180. if revokedToken != nil {
  181. logger.WarnContext(ctx, "token reuse detected, revoking all user tokens", slog.String("user_id", revokedToken.UserID))
  182. if revokeErr := s.refreshTokenRepo.RevokeAllForUser(ctx, revokedToken.UserID); revokeErr != nil {
  183. logger.ErrorContext(ctx, "failed to revoke all user tokens after reuse", log.WithError(revokeErr))
  184. }
  185. return nil, "", ErrTokenReused
  186. }
  187. logger.WarnContext(ctx, "refresh attempt with invalid/expired token")
  188. return nil, "", ErrInvalidToken
  189. }
  190. user, err := s.userRepo.GetByID(ctx, storedToken.UserID)
  191. if err != nil {
  192. logger.ErrorContext(ctx, "user lookup failed during refresh", log.WithError(err), slog.String("user_id", storedToken.UserID))
  193. return nil, "", fmt.Errorf("get user: %w", err)
  194. }
  195. if user == nil {
  196. logger.WarnContext(ctx, "refresh attempt for deleted user", slog.String("user_id", storedToken.UserID))
  197. return nil, "", ErrInvalidToken
  198. }
  199. if user.Status == "banned" {
  200. logger.WarnContext(ctx, "refresh attempt for banned user", slog.String("user_id", user.ID))
  201. return nil, "", ErrUserBanned
  202. }
  203. replaced, err := s.refreshTokenRepo.ReplaceIfExists(ctx, storedToken.TokenHash)
  204. if err != nil {
  205. logger.ErrorContext(ctx, "failed to replace old refresh token", log.WithError(err))
  206. return nil, "", fmt.Errorf("replace old refresh token: %w", err)
  207. }
  208. if !replaced {
  209. // Token was already replaced (grace period after lost cookie / fast F5).
  210. // Not theft — browser didn't receive the new cookie. Revoke this specific
  211. // token to prevent further reuse, then issue fresh tokens.
  212. if storedToken.ReplacedAt != nil {
  213. logger.InfoContext(ctx, "grace period token reuse, issuing new tokens",
  214. slog.String("user_id", user.ID),
  215. slog.Time("replaced_at", *storedToken.ReplacedAt))
  216. if revokeErr := s.refreshTokenRepo.Revoke(ctx, storedToken.TokenHash); revokeErr != nil {
  217. logger.ErrorContext(ctx, "failed to revoke grace-period token", log.WithError(revokeErr))
  218. }
  219. return s.generateTokens(ctx, user)
  220. }
  221. logger.WarnContext(ctx, "concurrent token rotation detected", slog.String("user_id", user.ID))
  222. return nil, "", ErrTokenReused
  223. }
  224. logger.InfoContext(ctx, "token refreshed", slog.String("user_id", user.ID))
  225. return s.generateTokens(ctx, user)
  226. }
  227. func (s *AuthService) GetUser(ctx context.Context, userID string) (*models.User, error) {
  228. return s.userRepo.GetByID(ctx, userID)
  229. }
  230. func (s *AuthService) RevokeSession(ctx context.Context, plainRefreshToken string) error {
  231. logger := log.FromContext(ctx)
  232. storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
  233. if err != nil {
  234. logger.ErrorContext(ctx, "revoke: refresh token lookup failed", log.WithError(err))
  235. return fmt.Errorf("lookup refresh token: %w", err)
  236. }
  237. if storedToken == nil {
  238. return nil
  239. }
  240. logger.InfoContext(ctx, "session revoked", slog.String("user_id", storedToken.UserID))
  241. return s.refreshTokenRepo.Revoke(ctx, storedToken.TokenHash)
  242. }
  243. func generateSecureToken(n int) (string, error) {
  244. b := make([]byte, n)
  245. if _, err := rand.Read(b); err != nil {
  246. return "", err
  247. }
  248. return base64.URLEncoding.EncodeToString(b), nil
  249. }