| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287 |
- package services
- import (
- "context"
- "crypto/rand"
- "encoding/base64"
- "errors"
- "fmt"
- "log/slog"
- "time"
- "github.com/golang-jwt/jwt/v5"
- "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/log"
- "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/models"
- "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/pointer"
- "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/repository"
- "golang.org/x/crypto/bcrypt"
- )
- var (
- ErrEmailExists = errors.New("email already exists")
- ErrInvalidCreds = errors.New("invalid email or password")
- ErrUserBanned = errors.New("account is banned")
- ErrInvalidToken = errors.New("invalid or expired token")
- ErrTokenReused = errors.New("token reused, possible theft")
- )
- const (
- AccessTokenExpiry = 15 * time.Minute
- RefreshTokenExpiry = 30 * 24 * time.Hour
- RefreshTokenBytes = 32
- )
- type UserRepo interface {
- Create(ctx context.Context, user *models.User) error
- GetByID(ctx context.Context, id string) (*models.User, error)
- GetByEmail(ctx context.Context, email string) (*models.User, error)
- Update(ctx context.Context, user *models.User) error
- UpdateRole(ctx context.Context, id, role string) error
- UpdateStatus(ctx context.Context, id, status string) error
- List(ctx context.Context, filter models.UserFilter) ([]*models.User, error)
- }
- type RefreshTokenRepo interface {
- Create(ctx context.Context, userID, plainToken string, expiresAt time.Time) error
- GetValid(ctx context.Context, plainToken string) (*models.RefreshToken, error)
- GetRevoked(ctx context.Context, plainToken string) (*models.RefreshToken, error)
- Revoke(ctx context.Context, tokenHash string) error
- RevokeAllForUser(ctx context.Context, userID string) error
- Delete(ctx context.Context, tokenHash string) error
- ReplaceIfExists(ctx context.Context, tokenHash string) (bool, error)
- CleanupExpired(ctx context.Context) error
- }
- type AuthService struct {
- userRepo UserRepo
- refreshTokenRepo RefreshTokenRepo
- jwtSecret []byte
- refreshSecret []byte
- }
- func NewAuthService(userRepo UserRepo, refreshTokenRepo RefreshTokenRepo, jwtSecret, refreshSecret string) *AuthService {
- return &AuthService{
- userRepo: userRepo,
- refreshTokenRepo: refreshTokenRepo,
- jwtSecret: []byte(jwtSecret),
- refreshSecret: []byte(refreshSecret),
- }
- }
- type AuthResult struct {
- User *models.User `json:"user"`
- AccessToken string `json:"access_token"`
- }
- type RegisterInput struct {
- Email string
- Password string
- Role string
- Name string
- }
- func (s *AuthService) Register(ctx context.Context, input RegisterInput) (*AuthResult, string, error) {
- logger := log.FromContext(ctx)
- hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost)
- if err != nil {
- logger.ErrorContext(ctx, "password hashing failed", log.WithError(err))
- return nil, "", fmt.Errorf("hash password: %w", err)
- }
- user := &models.User{
- Email: input.Email,
- PasswordHash: string(hash),
- Role: input.Role,
- Name: pointer.Str(input.Name),
- }
- if err := s.userRepo.Create(ctx, user); err != nil {
- if errors.Is(err, repository.ErrEmailExists) {
- logger.WarnContext(ctx, "registration attempt with existing email", slog.String("email", input.Email))
- return nil, "", ErrEmailExists
- }
- logger.ErrorContext(ctx, "user creation failed", log.WithError(err), slog.String("email", input.Email))
- return nil, "", fmt.Errorf("create user: %w", err)
- }
- logger.InfoContext(ctx, "user registered", slog.String("user_id", user.ID), slog.String("email", user.Email), slog.String("role", user.Role))
- return s.generateTokens(ctx, user)
- }
- type LoginInput struct {
- Email string
- Password string
- }
- func (s *AuthService) Login(ctx context.Context, input LoginInput) (*AuthResult, string, error) {
- logger := log.FromContext(ctx)
- user, err := s.userRepo.GetByEmail(ctx, input.Email)
- if err != nil {
- logger.ErrorContext(ctx, "user lookup failed", log.WithError(err), slog.String("email", input.Email))
- return nil, "", fmt.Errorf("get user: %w", err)
- }
- if user == nil {
- logger.WarnContext(ctx, "login attempt for non-existent user", slog.String("email", input.Email))
- return nil, "", ErrInvalidCreds
- }
- if user.Status == "banned" {
- logger.WarnContext(ctx, "login attempt for banned user", slog.String("user_id", user.ID))
- return nil, "", ErrUserBanned
- }
- if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(input.Password)); err != nil {
- logger.WarnContext(ctx, "invalid password attempt", slog.String("user_id", user.ID))
- return nil, "", ErrInvalidCreds
- }
- logger.InfoContext(ctx, "user logged in", slog.String("user_id", user.ID), slog.String("role", user.Role))
- return s.generateTokens(ctx, user)
- }
- type TokenClaims struct {
- UserID string `json:"user_id"`
- Role string `json:"role"`
- jwt.RegisteredClaims
- }
- func (s *AuthService) generateTokens(ctx context.Context, user *models.User) (*AuthResult, string, error) {
- now := time.Now()
- accessClaims := TokenClaims{
- UserID: user.ID,
- Role: user.Role,
- RegisteredClaims: jwt.RegisteredClaims{
- ExpiresAt: jwt.NewNumericDate(now.Add(AccessTokenExpiry)),
- IssuedAt: jwt.NewNumericDate(now),
- },
- }
- accessToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims).SignedString(s.jwtSecret)
- if err != nil {
- return nil, "", fmt.Errorf("sign access token: %w", err)
- }
- refreshToken, err := generateSecureToken(RefreshTokenBytes)
- if err != nil {
- return nil, "", fmt.Errorf("generate refresh token: %w", err)
- }
- expiresAt := now.Add(RefreshTokenExpiry)
- if err := s.refreshTokenRepo.Create(ctx, user.ID, refreshToken, expiresAt); err != nil {
- return nil, "", fmt.Errorf("store refresh token: %w", err)
- }
- return &AuthResult{
- User: user,
- AccessToken: accessToken,
- }, refreshToken, nil
- }
- func (s *AuthService) ValidateAccessToken(tokenString string) (*TokenClaims, error) {
- token, err := jwt.ParseWithClaims(tokenString, &TokenClaims{}, func(t *jwt.Token) (interface{}, error) {
- return s.jwtSecret, nil
- })
- if err != nil {
- return nil, ErrInvalidToken
- }
- claims, ok := token.Claims.(*TokenClaims)
- if !ok || !token.Valid {
- return nil, ErrInvalidToken
- }
- return claims, nil
- }
- func (s *AuthService) RefreshSession(ctx context.Context, plainRefreshToken string) (*AuthResult, string, error) {
- logger := log.FromContext(ctx)
- storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
- if err != nil {
- logger.ErrorContext(ctx, "refresh token lookup failed", log.WithError(err))
- return nil, "", fmt.Errorf("lookup refresh token: %w", err)
- }
- if storedToken == nil {
- revokedToken, err := s.refreshTokenRepo.GetRevoked(ctx, plainRefreshToken)
- if err != nil {
- logger.ErrorContext(ctx, "revoked token lookup failed", log.WithError(err))
- return nil, "", fmt.Errorf("lookup revoked token: %w", err)
- }
- if revokedToken != nil {
- logger.WarnContext(ctx, "token reuse detected, revoking all user tokens", slog.String("user_id", revokedToken.UserID))
- if revokeErr := s.refreshTokenRepo.RevokeAllForUser(ctx, revokedToken.UserID); revokeErr != nil {
- logger.ErrorContext(ctx, "failed to revoke all user tokens after reuse", log.WithError(revokeErr))
- }
- return nil, "", ErrTokenReused
- }
- logger.WarnContext(ctx, "refresh attempt with invalid/expired token")
- return nil, "", ErrInvalidToken
- }
- user, err := s.userRepo.GetByID(ctx, storedToken.UserID)
- if err != nil {
- logger.ErrorContext(ctx, "user lookup failed during refresh", log.WithError(err), slog.String("user_id", storedToken.UserID))
- return nil, "", fmt.Errorf("get user: %w", err)
- }
- if user == nil {
- logger.WarnContext(ctx, "refresh attempt for deleted user", slog.String("user_id", storedToken.UserID))
- return nil, "", ErrInvalidToken
- }
- if user.Status == "banned" {
- logger.WarnContext(ctx, "refresh attempt for banned user", slog.String("user_id", user.ID))
- return nil, "", ErrUserBanned
- }
- replaced, err := s.refreshTokenRepo.ReplaceIfExists(ctx, storedToken.TokenHash)
- if err != nil {
- logger.ErrorContext(ctx, "failed to replace old refresh token", log.WithError(err))
- return nil, "", fmt.Errorf("replace old refresh token: %w", err)
- }
- if !replaced {
- // Token was already replaced (grace period after lost cookie / fast F5).
- // Not theft — browser didn't receive the new cookie. Revoke this specific
- // token to prevent further reuse, then issue fresh tokens.
- if storedToken.ReplacedAt != nil {
- logger.InfoContext(ctx, "grace period token reuse, issuing new tokens",
- slog.String("user_id", user.ID),
- slog.Time("replaced_at", *storedToken.ReplacedAt))
- if revokeErr := s.refreshTokenRepo.Revoke(ctx, storedToken.TokenHash); revokeErr != nil {
- logger.ErrorContext(ctx, "failed to revoke grace-period token", log.WithError(revokeErr))
- }
- return s.generateTokens(ctx, user)
- }
- logger.WarnContext(ctx, "concurrent token rotation detected", slog.String("user_id", user.ID))
- return nil, "", ErrTokenReused
- }
- logger.InfoContext(ctx, "token refreshed", slog.String("user_id", user.ID))
- return s.generateTokens(ctx, user)
- }
- func (s *AuthService) GetUser(ctx context.Context, userID string) (*models.User, error) {
- return s.userRepo.GetByID(ctx, userID)
- }
- func (s *AuthService) RevokeSession(ctx context.Context, plainRefreshToken string) error {
- logger := log.FromContext(ctx)
- storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
- if err != nil {
- logger.ErrorContext(ctx, "revoke: refresh token lookup failed", log.WithError(err))
- return fmt.Errorf("lookup refresh token: %w", err)
- }
- if storedToken == nil {
- return nil
- }
- logger.InfoContext(ctx, "session revoked", slog.String("user_id", storedToken.UserID))
- return s.refreshTokenRepo.Revoke(ctx, storedToken.TokenHash)
- }
- func generateSecureToken(n int) (string, error) {
- b := make([]byte, n)
- if _, err := rand.Read(b); err != nil {
- return "", err
- }
- return base64.URLEncoding.EncodeToString(b), nil
- }
|