|
@@ -2,6 +2,8 @@ package services
|
|
|
|
|
|
|
|
import (
|
|
import (
|
|
|
"context"
|
|
"context"
|
|
|
|
|
+ "crypto/rand"
|
|
|
|
|
+ "encoding/base64"
|
|
|
"errors"
|
|
"errors"
|
|
|
"fmt"
|
|
"fmt"
|
|
|
"time"
|
|
"time"
|
|
@@ -17,26 +19,34 @@ var (
|
|
|
ErrInvalidCreds = errors.New("invalid email or password")
|
|
ErrInvalidCreds = errors.New("invalid email or password")
|
|
|
ErrUserBanned = errors.New("account is banned")
|
|
ErrUserBanned = errors.New("account is banned")
|
|
|
ErrInvalidToken = errors.New("invalid or expired token")
|
|
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 {
|
|
type AuthService struct {
|
|
|
- userRepo *repository.UserRepo
|
|
|
|
|
- jwtSecret []byte
|
|
|
|
|
- refreshSecret []byte
|
|
|
|
|
|
|
+ userRepo *repository.UserRepo
|
|
|
|
|
+ refreshTokenRepo *repository.RefreshTokenRepo
|
|
|
|
|
+ jwtSecret []byte
|
|
|
|
|
+ refreshSecret []byte
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
-func NewAuthService(userRepo *repository.UserRepo, jwtSecret, refreshSecret string) *AuthService {
|
|
|
|
|
|
|
+func NewAuthService(userRepo *repository.UserRepo, refreshTokenRepo *repository.RefreshTokenRepo, jwtSecret, refreshSecret string) *AuthService {
|
|
|
return &AuthService{
|
|
return &AuthService{
|
|
|
- userRepo: userRepo,
|
|
|
|
|
- jwtSecret: []byte(jwtSecret),
|
|
|
|
|
- refreshSecret: []byte(refreshSecret),
|
|
|
|
|
|
|
+ userRepo: userRepo,
|
|
|
|
|
+ refreshTokenRepo: refreshTokenRepo,
|
|
|
|
|
+ jwtSecret: []byte(jwtSecret),
|
|
|
|
|
+ refreshSecret: []byte(refreshSecret),
|
|
|
}
|
|
}
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
type AuthResult struct {
|
|
type AuthResult struct {
|
|
|
- User *models.User `json:"user"`
|
|
|
|
|
- AccessToken string `json:"access_token"`
|
|
|
|
|
- RefreshToken string `json:"refresh_token"`
|
|
|
|
|
|
|
+ User *models.User `json:"user"`
|
|
|
|
|
+ AccessToken string `json:"access_token"`
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
type RegisterInput struct {
|
|
type RegisterInput struct {
|
|
@@ -46,15 +56,15 @@ type RegisterInput struct {
|
|
|
Name string
|
|
Name string
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
-func (s *AuthService) Register(ctx context.Context, input RegisterInput) (*AuthResult, error) {
|
|
|
|
|
|
|
+func (s *AuthService) Register(ctx context.Context, input RegisterInput) (*AuthResult, string, error) {
|
|
|
existing, _ := s.userRepo.GetByEmail(ctx, input.Email)
|
|
existing, _ := s.userRepo.GetByEmail(ctx, input.Email)
|
|
|
if existing != nil {
|
|
if existing != nil {
|
|
|
- return nil, ErrEmailExists
|
|
|
|
|
|
|
+ return nil, "", ErrEmailExists
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost)
|
|
hash, err := bcrypt.GenerateFromPassword([]byte(input.Password), bcrypt.DefaultCost)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
- return nil, fmt.Errorf("hash password: %w", err)
|
|
|
|
|
|
|
+ return nil, "", fmt.Errorf("hash password: %w", err)
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
user := &models.User{
|
|
user := &models.User{
|
|
@@ -65,7 +75,7 @@ func (s *AuthService) Register(ctx context.Context, input RegisterInput) (*AuthR
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
if err := s.userRepo.Create(ctx, user); err != nil {
|
|
if err := s.userRepo.Create(ctx, user); err != nil {
|
|
|
- return nil, fmt.Errorf("create user: %w", err)
|
|
|
|
|
|
|
+ return nil, "", fmt.Errorf("create user: %w", err)
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
return s.generateTokens(ctx, user)
|
|
return s.generateTokens(ctx, user)
|
|
@@ -76,21 +86,21 @@ type LoginInput struct {
|
|
|
Password string
|
|
Password string
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
-func (s *AuthService) Login(ctx context.Context, input LoginInput) (*AuthResult, error) {
|
|
|
|
|
|
|
+func (s *AuthService) Login(ctx context.Context, input LoginInput) (*AuthResult, string, error) {
|
|
|
user, err := s.userRepo.GetByEmail(ctx, input.Email)
|
|
user, err := s.userRepo.GetByEmail(ctx, input.Email)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
- return nil, fmt.Errorf("get user: %w", err)
|
|
|
|
|
|
|
+ return nil, "", fmt.Errorf("get user: %w", err)
|
|
|
}
|
|
}
|
|
|
if user == nil {
|
|
if user == nil {
|
|
|
- return nil, ErrInvalidCreds
|
|
|
|
|
|
|
+ return nil, "", ErrInvalidCreds
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
if user.Status == "banned" {
|
|
if user.Status == "banned" {
|
|
|
- return nil, ErrUserBanned
|
|
|
|
|
|
|
+ return nil, "", ErrUserBanned
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(input.Password)); err != nil {
|
|
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(input.Password)); err != nil {
|
|
|
- return nil, ErrInvalidCreds
|
|
|
|
|
|
|
+ return nil, "", ErrInvalidCreds
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
return s.generateTokens(ctx, user)
|
|
return s.generateTokens(ctx, user)
|
|
@@ -102,37 +112,36 @@ type TokenClaims struct {
|
|
|
jwt.RegisteredClaims
|
|
jwt.RegisteredClaims
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
-func (s *AuthService) generateTokens(ctx context.Context, user *models.User) (*AuthResult, error) {
|
|
|
|
|
|
|
+func (s *AuthService) generateTokens(ctx context.Context, user *models.User) (*AuthResult, string, error) {
|
|
|
now := time.Now()
|
|
now := time.Now()
|
|
|
|
|
|
|
|
accessClaims := TokenClaims{
|
|
accessClaims := TokenClaims{
|
|
|
UserID: user.ID,
|
|
UserID: user.ID,
|
|
|
Role: user.Role,
|
|
Role: user.Role,
|
|
|
RegisteredClaims: jwt.RegisteredClaims{
|
|
RegisteredClaims: jwt.RegisteredClaims{
|
|
|
- ExpiresAt: jwt.NewNumericDate(now.Add(15 * time.Minute)),
|
|
|
|
|
|
|
+ ExpiresAt: jwt.NewNumericDate(now.Add(AccessTokenExpiry)),
|
|
|
IssuedAt: jwt.NewNumericDate(now),
|
|
IssuedAt: jwt.NewNumericDate(now),
|
|
|
},
|
|
},
|
|
|
}
|
|
}
|
|
|
accessToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims).SignedString(s.jwtSecret)
|
|
accessToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims).SignedString(s.jwtSecret)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
- return nil, fmt.Errorf("sign access token: %w", err)
|
|
|
|
|
|
|
+ return nil, "", fmt.Errorf("sign access token: %w", err)
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
- refreshClaims := jwt.RegisteredClaims{
|
|
|
|
|
- ExpiresAt: jwt.NewNumericDate(now.Add(30 * 24 * time.Hour)),
|
|
|
|
|
- IssuedAt: jwt.NewNumericDate(now),
|
|
|
|
|
- Subject: user.ID,
|
|
|
|
|
- }
|
|
|
|
|
- refreshToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims).SignedString(s.refreshSecret)
|
|
|
|
|
|
|
+ refreshToken, err := generateSecureToken(RefreshTokenBytes)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
- return nil, fmt.Errorf("sign refresh token: %w", err)
|
|
|
|
|
|
|
+ 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{
|
|
return &AuthResult{
|
|
|
- User: user,
|
|
|
|
|
- AccessToken: accessToken,
|
|
|
|
|
- RefreshToken: refreshToken,
|
|
|
|
|
- }, nil
|
|
|
|
|
|
|
+ User: user,
|
|
|
|
|
+ AccessToken: accessToken,
|
|
|
|
|
+ }, refreshToken, nil
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
func (s *AuthService) ValidateAccessToken(tokenString string) (*TokenClaims, error) {
|
|
func (s *AuthService) ValidateAccessToken(tokenString string) (*TokenClaims, error) {
|
|
@@ -149,43 +158,63 @@ func (s *AuthService) ValidateAccessToken(tokenString string) (*TokenClaims, err
|
|
|
return claims, nil
|
|
return claims, nil
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
-func (s *AuthService) ValidateRefreshToken(tokenString string) (string, error) {
|
|
|
|
|
- token, err := jwt.ParseWithClaims(tokenString, &jwt.RegisteredClaims{}, func(t *jwt.Token) (interface{}, error) {
|
|
|
|
|
- return s.refreshSecret, nil
|
|
|
|
|
- })
|
|
|
|
|
|
|
+func (s *AuthService) RefreshSession(ctx context.Context, plainRefreshToken string) (*AuthResult, string, error) {
|
|
|
|
|
+ storedToken, err := s.refreshTokenRepo.GetValid(ctx, plainRefreshToken)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
- return "", ErrInvalidToken
|
|
|
|
|
- }
|
|
|
|
|
- claims, ok := token.Claims.(*jwt.RegisteredClaims)
|
|
|
|
|
- if !ok || !token.Valid {
|
|
|
|
|
- return "", ErrInvalidToken
|
|
|
|
|
|
|
+ return nil, "", fmt.Errorf("lookup refresh token: %w", err)
|
|
|
}
|
|
}
|
|
|
- return claims.Subject, nil
|
|
|
|
|
-}
|
|
|
|
|
-
|
|
|
|
|
-func (s *AuthService) RefreshSession(ctx context.Context, refreshToken string) (*AuthResult, error) {
|
|
|
|
|
- userID, err := s.ValidateRefreshToken(refreshToken)
|
|
|
|
|
- if err != nil {
|
|
|
|
|
- return nil, err
|
|
|
|
|
|
|
+ if storedToken == nil {
|
|
|
|
|
+ return nil, "", ErrInvalidToken
|
|
|
}
|
|
}
|
|
|
|
|
|
|
|
- user, err := s.userRepo.GetByID(ctx, userID)
|
|
|
|
|
|
|
+ user, err := s.userRepo.GetByID(ctx, storedToken.UserID)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
- return nil, fmt.Errorf("get user: %w", err)
|
|
|
|
|
|
|
+ return nil, "", fmt.Errorf("get user: %w", err)
|
|
|
}
|
|
}
|
|
|
if user == nil {
|
|
if user == nil {
|
|
|
- return nil, ErrInvalidToken
|
|
|
|
|
|
|
+ return nil, "", ErrInvalidToken
|
|
|
}
|
|
}
|
|
|
if user.Status == "banned" {
|
|
if user.Status == "banned" {
|
|
|
- return nil, ErrUserBanned
|
|
|
|
|
|
|
+ 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)
|
|
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 {
|
|
func strPtr(s string) *string {
|
|
|
if s == "" {
|
|
if s == "" {
|
|
|
return nil
|
|
return nil
|
|
|
}
|
|
}
|
|
|
return &s
|
|
return &s
|
|
|
-}
|
|
|
|
|
|
|
+}
|