161 lines
3.5 KiB
Go
161 lines
3.5 KiB
Go
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[:])
|
||
}
|