Files
2026-08-19 22:05:49 +08:00

161 lines
3.5 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package middleware
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"fmt"
"net/http"
"server/internal/config"
"server/internal/pkg/errs"
"server/internal/pkg/httputil"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
)
type JWTMiddleware struct {
cfg *config.JWTConfig
}
type Claims struct {
UserID int32 `json:"user_id"`
TokenVersion int32 `json:"token_version"`
jwt.RegisteredClaims
}
func NewJWTMiddleware(cfg *config.Config) (*JWTMiddleware, error) {
jwtCfg := cfg.JWTConfig
// 不允许空密钥
if jwtCfg.Secret == "" {
return nil, fmt.Errorf("jwt.secret 未配置:请在 %s.yaml 中设置 jwt.secret", config.GetEnv())
}
m := &JWTMiddleware{cfg: &jwtCfg}
// 启动时校验签名算法配置
if _, err := m.signingMethod(); err != nil {
return nil, err
}
return m, nil
}
// signingMethod 返回配置的签名算法(仅支持 HMAC 家族)
func (m *JWTMiddleware) signingMethod() (jwt.SigningMethod, error) {
switch strings.ToUpper(m.cfg.SigningMethod) {
case "", "HS256":
return jwt.SigningMethodHS256, nil
case "HS384":
return jwt.SigningMethodHS384, nil
case "HS512":
return jwt.SigningMethodHS512, nil
default:
return nil, fmt.Errorf("不支持的 jwt.signing_method: %q仅支持 HS256/HS384/HS512", m.cfg.SigningMethod)
}
}
// ParseToken 解析accessToken
func (m *JWTMiddleware) ParseToken(tokenStr string) (*Claims, error) {
token, err := jwt.ParseWithClaims(
tokenStr,
&Claims{},
func(token *jwt.Token) (any, error) {
return []byte(m.cfg.Secret), nil
},
)
if err != nil || !token.Valid {
return nil, errs.ErrInvalidToken
}
claims, ok := token.Claims.(*Claims)
if !ok {
return nil, errs.ErrInvalidTokenClaims
}
if claims.UserID == 0 {
return nil, errs.ErrInvalidTokenClaims
}
return claims, nil
}
func (m *JWTMiddleware) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
auth := r.Header.Get("Authorization")
if auth == "" {
httputil.Fail(w, errs.ErrUnauthenticated)
return
}
parts := strings.SplitN(auth, " ", 2)
if len(parts) != 2 || parts[0] != "Bearer" {
httputil.Fail(w, errs.ErrUnauthenticated)
return
}
tokenStr := parts[1]
claims, err := m.ParseToken(tokenStr)
if err != nil {
httputil.Fail(w, errs.ErrInvalidToken)
return
}
userCtx := GetUserContext(r.Context())
userCtx.UserID = claims.UserID
userCtx.TokenVersion = claims.TokenVersion
next.ServeHTTP(w, r)
})
}
func (m *JWTMiddleware) GenerateAccessToken(userID int32, tokenVersion int32) (string, time.Time, error) {
method, err := m.signingMethod()
if err != nil {
return "", time.Time{}, err
}
now := time.Now()
expiresAt := now.Add(m.cfg.Expire)
accessClaims := Claims{
UserID: userID,
TokenVersion: tokenVersion,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(expiresAt),
IssuedAt: jwt.NewNumericDate(now),
},
}
// 使用配置的签名算法,与 ParseToken 的校验保持一致
accessToken := jwt.NewWithClaims(method, accessClaims)
token, err := accessToken.SignedString([]byte(m.cfg.Secret))
if err != nil {
return "", time.Time{}, err
}
return token, expiresAt, nil
}
func (m *JWTMiddleware) GenerateRefreshToken() (string, error) {
b := make([]byte, 32)
_, err := rand.Read(b)
if err != nil {
return "", err
}
return hex.EncodeToString(b), nil
}
// HashRefreshToken 哈希token 存储至redis
func (m *JWTMiddleware) HashRefreshToken(token string) string {
hash := sha256.Sum256([]byte(token))
return hex.EncodeToString(hash[:])
}