package auth import ( "crypto/rand" "encoding/hex" "errors" "fmt" "strings" "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 UserType string `json:"user_type,omitempty"` // user/super_admin platform identity Provider string `json:"provider"` // email/google/oidc CustomRoleID uint `json:"custom_role_id,omitempty"` // enterprise custom role ClientID string `json:"client_id,omitempty"` // DeviseTokenAuth-compatible session id 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} } // InsecureHeaderAuthAllowed reports whether trusted dev/test callers may use // X-User-ID instead of a signed token. Release config validation rejects this. func (s *JWTService) InsecureHeaderAuthAllowed() bool { return s != nil && s.cfg != nil && s.cfg.AllowInsecureHeaderAuth } // 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) { return s.GenerateTokenPairForClient(user, accountID, role, "") } func (s *JWTService) GenerateTokenPairForClient(user *model.User, accountID uint, role, clientID string) (*TokenPair, error) { userType := "user" typeValue := strings.ToLower(strings.ReplaceAll(strings.TrimSpace(user.Type), "_", "")) if user.Role == "super_admin" || role == "super_admin" || typeValue == "superadmin" { userType = "super_admin" } // Access Token now := time.Now() accessExpiry := now.Add(s.cfg.ExpiryDuration()) issuer := s.cfg.Issuer if issuer == "" { issuer = "gochat" } accessClaims := &Claims{ UserID: user.ID, AccountID: accountID, Role: role, UserType: userType, Provider: user.Provider, CustomRoleID: func() uint { if user.CustomRoleID != nil { return *user.CustomRoleID } return 0 }(), ClientID: clientID, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(accessExpiry), IssuedAt: jwt.NewNumericDate(now), Subject: fmt.Sprintf("user_%d", user.ID), Issuer: issuer, Audience: jwt.ClaimStrings{s.cfg.Audience}, }, } 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 := now.Add(s.cfg.RefreshExpiryDuration()) refreshID := make([]byte, 16) if _, err := rand.Read(refreshID); err != nil { return nil, fmt.Errorf("failed to generate refresh token id: %w", err) } refreshClaims := &Claims{ UserID: user.ID, Provider: user.Provider, ClientID: clientID, RegisteredClaims: jwt.RegisteredClaims{ ID: hex.EncodeToString(refreshID), ExpiresAt: jwt.NewNumericDate(refreshExpiry), IssuedAt: jwt.NewNumericDate(now), Subject: fmt.Sprintf("refresh_%d", user.ID), Issuer: issuer, Audience: jwt.ClaimStrings{s.cfg.Audience}, }, } 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) { return s.validateToken(tokenString, "user", "access") } // ValidateRefreshToken validates a refresh JWT token and returns claims. func (s *JWTService) ValidateRefreshToken(tokenString string) (*Claims, error) { return s.validateToken(tokenString, "refresh", "refresh") } func (s *JWTService) validateToken(tokenString, subjectPrefix, kind string) (*Claims, error) { if s == nil || s.cfg == nil { return nil, errors.New("JWT service is not configured") } secrets := append([]string{s.cfg.Secret}, s.cfg.PreviousSecrets...) var lastErr error for index, secret := range secrets { claims := &Claims{} token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) { return []byte(secret), nil }, s.validationOptions(index == 0)...) if err != nil { lastErr = err continue } if !token.Valid { lastErr = errors.New("invalid token claims") continue } if claims.Subject != fmt.Sprintf("%s_%d", subjectPrefix, claims.UserID) { return nil, fmt.Errorf("invalid token type: expected %s token", kind) } return claims, nil } return nil, fmt.Errorf("failed to parse %s token: %w", kind, lastErr) } func (s *JWTService) validationOptions(validateClaims bool) []jwt.ParserOption { options := []jwt.ParserOption{jwt.WithValidMethods([]string{jwt.SigningMethodHS256.Alg()})} if !validateClaims { return options } if s.cfg.Issuer != "" { options = append(options, jwt.WithIssuer(s.cfg.Issuer)) } if s.cfg.Audience != "" { options = append(options, jwt.WithAudience(s.cfg.Audience)) } return options } // 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.GenerateTokenPairForClient(user, accountID, role, claims.ClientID) }