Files

165 lines
4.6 KiB
Go

package auth
import (
"errors"
"time"
"github.com/golang-jwt/jwt/v5"
)
const (
tokenTypeAccess = "access"
tokenTypeRefresh = "refresh"
)
var ErrInvalidToken = errors.New("invalid token")
// TokenValidationError 仅携带可安全记录的失败类别,不包含原始令牌或签名信息。
type TokenValidationError struct {
Reason string
err error
}
func (e *TokenValidationError) Error() string {
return "token validation failed: " + e.Reason
}
func (e *TokenValidationError) Unwrap() error {
return e.err
}
// AuthFailureReason 让 HTTP 鉴权层可在不依赖具体认证包的情况下读取失败类别。
func (e *TokenValidationError) AuthFailureReason() string {
return e.Reason
}
// TokenFailureReason 返回适合写入日志的令牌校验失败类别。
func TokenFailureReason(err error) string {
var validationErr *TokenValidationError
if errors.As(err, &validationErr) && validationErr.Reason != "" {
return validationErr.Reason
}
return "invalid"
}
type JWTManager struct {
secret []byte
accessTTL time.Duration
refreshTTL time.Duration
signingMethod jwt.SigningMethod
}
type Claims struct {
UserID uint64 `json:"uid"`
Phone string `json:"phone"`
TokenType string `json:"typ"`
SubjectType string `json:"sub_type"`
TokenVersion int64 `json:"ver,omitempty"`
jwt.RegisteredClaims
}
type TokenPair struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
TokenType string `json:"token_type"`
ExpiresInSeconds int64 `json:"expires_in"`
}
func NewJWTManager(secret string) *JWTManager {
return &JWTManager{
secret: []byte(secret),
accessTTL: 2 * time.Hour,
refreshTTL: 14 * 24 * time.Hour,
signingMethod: jwt.SigningMethodHS256,
}
}
func (m *JWTManager) GeneratePair(userID uint64, phone string) (TokenPair, error) {
return m.GenerateSubjectPair(userID, phone, "user")
}
func (m *JWTManager) GenerateSubjectPair(userID uint64, subject string, subjectType string) (TokenPair, error) {
return m.GenerateSubjectPairWithVersion(userID, subject, subjectType, 0)
}
func (m *JWTManager) GenerateSubjectPairWithVersion(userID uint64, subject string, subjectType string, tokenVersion int64) (TokenPair, error) {
accessToken, err := m.generate(userID, subject, subjectType, tokenTypeAccess, tokenVersion, m.accessTTL)
if err != nil {
return TokenPair{}, err
}
refreshToken, err := m.generate(userID, subject, subjectType, tokenTypeRefresh, tokenVersion, m.refreshTTL)
if err != nil {
return TokenPair{}, err
}
return TokenPair{
AccessToken: accessToken,
RefreshToken: refreshToken,
TokenType: "Bearer",
ExpiresInSeconds: int64(m.accessTTL.Seconds()),
}, nil
}
func (m *JWTManager) Parse(tokenText, expectedType string) (*Claims, error) {
return m.ParseSubject(tokenText, expectedType, "")
}
func (m *JWTManager) ParseSubject(tokenText, expectedType string, expectedSubjectType string) (*Claims, error) {
claims := &Claims{}
token, err := jwt.ParseWithClaims(tokenText, claims, func(token *jwt.Token) (any, error) {
if token.Method != m.signingMethod {
return nil, newTokenValidationError("signing_method")
}
return m.secret, nil
})
if err != nil {
return nil, newTokenValidationError(jwtFailureReason(err))
}
if !token.Valid {
return nil, newTokenValidationError("invalid")
}
if claims.TokenType != expectedType {
return nil, newTokenValidationError("token_type")
}
if expectedSubjectType != "" && claims.SubjectType != expectedSubjectType {
return nil, newTokenValidationError("subject_type")
}
return claims, nil
}
func newTokenValidationError(reason string) error {
return &TokenValidationError{Reason: reason, err: ErrInvalidToken}
}
func jwtFailureReason(err error) string {
switch {
case errors.Is(err, jwt.ErrTokenExpired):
return "expired"
case errors.Is(err, jwt.ErrTokenSignatureInvalid):
return "signature_invalid"
case errors.Is(err, jwt.ErrTokenMalformed):
return "malformed"
case errors.Is(err, jwt.ErrTokenNotValidYet):
return "not_valid_yet"
default:
return "invalid"
}
}
func (m *JWTManager) generate(userID uint64, subject string, subjectType string, tokenType string, tokenVersion int64, ttl time.Duration) (string, error) {
now := time.Now()
claims := Claims{
UserID: userID,
Phone: subject,
TokenType: tokenType,
SubjectType: subjectType,
TokenVersion: tokenVersion,
RegisteredClaims: jwt.RegisteredClaims{
Subject: subject,
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(ttl)),
},
}
token := jwt.NewWithClaims(m.signingMethod, claims)
return token.SignedString(m.secret)
}