package auth import ( "errors" "fmt" "time" "github.com/golang-jwt/jwt/v5" "github.com/gochat/gochat/internal/config" "github.com/gochat/gochat/internal/model" ) // Reference: P2E §1.2 — JWT Claims structure // Replaces Chatwoot's DeviseTokenAuth multi-token mechanism with single JWT + Refresh token. // Claims represents JWT token claims. type Claims struct { UserID uint `json:"user_id"` AccountID uint `json:"account_id"` // current active account Role string `json:"role"` // agent/administrator/custom_role Provider string `json:"provider"` // email/google/saml CustomRoleID uint `json:"custom_role_id,omitempty"` // enterprise custom role jwt.RegisteredClaims } // TokenPair holds an access token and refresh token. type TokenPair struct { AccessToken string `json:"access_token"` RefreshToken string `json:"refresh_token"` ExpiresAt time.Time `json:"expires_at"` } // JWTService provides JWT token generation and validation. type JWTService struct { cfg *config.JWTConfig } // NewJWTService creates a JWT service with configuration. func NewJWTService(cfg *config.JWTConfig) *JWTService { return &JWTService{cfg: cfg} } // GenerateTokenPair generates an access token + refresh token pair. // Access Token: 15min expiry with full Claims // Refresh Token: 7 days expiry, only UserID + Provider func (s *JWTService) GenerateTokenPair(user *model.User, accountID uint, role string) (*TokenPair, error) { // Access Token accessExpiry := time.Now().Add(time.Duration(s.cfg.ExpiryHours) * time.Hour) accessClaims := &Claims{ UserID: user.ID, AccountID: accountID, Role: role, Provider: user.Provider, CustomRoleID: func() uint { if user.CustomRoleID != nil { return *user.CustomRoleID }; return 0 }(), RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(accessExpiry), IssuedAt: jwt.NewNumericDate(time.Now()), Subject: fmt.Sprintf("user_%d", user.ID), Issuer: "gochat", }, } accessToken := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims) accessTokenString, err := accessToken.SignedString([]byte(s.cfg.Secret)) if err != nil { return nil, fmt.Errorf("failed to sign access token: %w", err) } // Refresh Token refreshExpiry := time.Now().Add(time.Duration(s.cfg.RefreshExpiryHours) * time.Hour) refreshClaims := &Claims{ UserID: user.ID, Provider: user.Provider, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(refreshExpiry), IssuedAt: jwt.NewNumericDate(time.Now()), Subject: fmt.Sprintf("refresh_%d", user.ID), Issuer: "gochat", }, } refreshToken := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims) refreshTokenString, err := refreshToken.SignedString([]byte(s.cfg.Secret)) if err != nil { return nil, fmt.Errorf("failed to sign refresh token: %w", err) } return &TokenPair{ AccessToken: accessTokenString, RefreshToken: refreshTokenString, ExpiresAt: accessExpiry, }, nil } // ValidateAccessToken validates an access JWT token and returns claims. func (s *JWTService) ValidateAccessToken(tokenString string) (*Claims, error) { token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) { if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) } return []byte(s.cfg.Secret), nil }) if err != nil { return nil, fmt.Errorf("failed to parse token: %w", err) } claims, ok := token.Claims.(*Claims) if !ok || !token.Valid { return nil, errors.New("invalid token claims") } // Verify it's an access token (not a refresh token) if claims.Subject != fmt.Sprintf("user_%d", claims.UserID) { return nil, errors.New("invalid token type: expected access token") } return claims, nil } // ValidateRefreshToken validates a refresh JWT token and returns claims. func (s *JWTService) ValidateRefreshToken(tokenString string) (*Claims, error) { token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) { if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) } return []byte(s.cfg.Secret), nil }) if err != nil { return nil, fmt.Errorf("failed to parse refresh token: %w", err) } claims, ok := token.Claims.(*Claims) if !ok || !token.Valid { return nil, errors.New("invalid refresh token claims") } // Verify it's a refresh token if claims.Subject != fmt.Sprintf("refresh_%d", claims.UserID) { return nil, errors.New("invalid token type: expected refresh token") } return claims, nil } // RefreshAccessToken generates a new access token from a valid refresh token. func (s *JWTService) RefreshAccessToken(refreshTokenString string, accountID uint, role string) (*TokenPair, error) { claims, err := s.ValidateRefreshToken(refreshTokenString) if err != nil { return nil, err } // Generate new token pair with the user ID from refresh claims // Note: we need the full user object for GenerateTokenPair, but we only have claims here // In production, we fetch user from DB using claims.UserID user := &model.User{ Base: model.Base{ID: claims.UserID}, Provider: claims.Provider, } return s.GenerateTokenPair(user, accountID, role) }