package services import ( "context" "crypto/rand" "encoding/base64" "errors" "fmt" "time" "github.com/golang-jwt/jwt/v5" "github.com/photoplaces/backend/internal/models" "github.com/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 AuthService struct { userRepo *repository.UserRepo refreshTokenRepo *repository.RefreshTokenRepo jwtSecret []byte refreshSecret []byte } func NewAuthService(userRepo *repository.UserRepo, refreshTokenRepo *repository.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) { existing, _ := s.userRepo.GetByEmail(ctx, input.Email) if existing != nil { return nil, "", ErrEmailExists } hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost) if err != nil { return nil, "", fmt.Errorf("hash password: %w", err) } user := &models.User{ Email: input.Email, PasswordHash: string(hash), Role: input.Role, Name: strPtr(input.Name), } if err := s.userRepo.Create(ctx, user); err != nil { return nil, "", fmt.Errorf("create user: %w", err) } return s.generateTokens(ctx, user) } type LoginInput struct { Email string Password string } func (s *AuthService) Login(ctx context.Context, input LoginInput) (*AuthResult, string, error) { user, err := s.userRepo.GetByEmail(ctx, input.Email) if err != nil { return nil, "", fmt.Errorf("get user: %w", err) } if user == nil { return nil, "", ErrInvalidCreds } if user.Status == "banned" { return nil, "", ErrUserBanned } if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(input.Password)); err != nil { return nil, "", ErrInvalidCreds } 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) { storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken) if err != nil { return nil, "", fmt.Errorf("lookup refresh token: %w", err) } if storedToken == nil { return nil, "", ErrInvalidToken } user, err := s.userRepo.GetByID(ctx, storedToken.UserID) if err != nil { return nil, "", fmt.Errorf("get user: %w", err) } if user == nil { return nil, "", ErrInvalidToken } if user.Status == "banned" { return nil, "", ErrUserBanned } if err := s.refreshTokenRepo.Delete(ctx, storedToken.TokenHash); err != nil { return nil, "", fmt.Errorf("delete old refresh token: %w", err) } return s.generateTokens(ctx, user) } func (s *AuthService) RevokeSession(ctx context.Context, plainRefreshToken string) error { storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken) if err != nil { return fmt.Errorf("lookup refresh token: %w", err) } if storedToken == nil { return nil } return s.refreshTokenRepo.Revoke(ctx, storedToken.TokenHash) } func (s *AuthService) RevokeAllSessions(ctx context.Context, userID string) error { return s.refreshTokenRepo.RevokeAllForUser(ctx, userID) } func (s *AuthService) CleanupExpiredTokens(ctx context.Context) error { return s.refreshTokenRepo.CleanupExpired(ctx) } 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 } func strPtr(s string) *string { if s == "" { return nil } return &s }