206 lines
6.6 KiB
Go
206 lines
6.6 KiB
Go
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 == RoleSuperAdmin || role == RoleSuperAdmin || typeValue == "superadmin" {
|
|
userType = RoleSuperAdmin
|
|
}
|
|
|
|
// 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)
|
|
}
|