Go语言Gin框架中间件链与JWT认证鉴权拦截器设计实战

Gin框架的中间件机制是其核心设计之一。中间件以洋葱模型执行,请求依次穿过每一层中间件到达业务处理函数,响应再按相反顺序返回。合理设计中间件链可以实现认证、鉴权、限流、日志等横切关注点的解耦。本文通过一个完整的JWT认证与RBAC鉴权中间件实现,展示Gin中间件的设计模式。

Gin中间件执行机制与洋葱模型原理

Gin中间件通过c.Next()控制执行流。调用Next()前的代码在请求阶段执行,Next()后的代码在响应阶段执行,形成洋葱式穿透:

func main() {
    r := gin.New()
    
    // 中间件执行顺序:A -> B -> handler -> B -> A
    r.Use(MiddlewareA(), MiddlewareB())
    
    r.GET("/api", func(c *gin.Context) {
        c.JSON(200, gin.H{"message": "ok"})
    })
}

func MiddlewareA() gin.HandlerFunc {
    return func(c *gin.Context) {
        start := time.Now()
        log.Printf("MiddlewareA: 请求开始")
        
        c.Next()  // 进入下一层中间件或handler
        
        latency := time.Since(start)
        log.Printf("MiddlewareA: 请求完成, 耗时: %v", latency)
    }
}

func MiddlewareB() gin.HandlerFunc {
    return func(c *gin.Context) {
        log.Printf("MiddlewareB: 请求开始")
        c.Next()
        log.Printf("MiddlewareB: 请求完成")
    }
}
// 输出顺序:
// MiddlewareA: 请求开始
// MiddlewareB: 请求开始
// (handler执行)
// MiddlewareB: 请求完成
// MiddlewareA: 请求完成, 耗时: xxx

调用c.Abort()会中断后续中间件和handler的执行,常用于认证失败时提前返回。结合AbortWithStatusJSON可以一步完成中断和响应。

JWT认证中间件实现与Token签发验证

package middleware

import (
    "net/http"
    "strings"
    "time"
    "errors"
    
    "github.com/gin-gonic/gin"
    "github.com/golang-jwt/jwt/v5"
)

var jwtSecret = []byte("your-256-bit-secret-key-change-in-production")

// Claims定义JWT的载荷结构
type Claims struct {
    UserID   int64    `json:"user_id"`
    Username string   `json:"username"`
    Role     string   `json:"role"`
    jwt.RegisteredClaims
}

// GenerateToken 签发JWT Token
func GenerateToken(userID int64, username, role string) (string, error) {
    claims := Claims{
        UserID:   userID,
        Username: username,
        Role:     role,
        RegisteredClaims: jwt.RegisteredClaims{
            ExpiresAt: jwt.NewNumericDate(time.Now().Add(24 * time.Hour)),
            IssuedAt:  jwt.NewNumericDate(time.Now()),
            Issuer:    "myapp",
            Subject:   username,
        },
    }
    
    token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
    return token.SignedString(jwtSecret)
}

// ParseToken 解析JWT Token
func ParseToken(tokenString string) (*Claims, error) {
    claims := &Claims{}
    token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) {
        return jwtSecret, nil
    })
    if err != nil {
        return nil, err
    }
    if !token.Valid {
        return nil, errors.New("invalid token")
    }
    return claims, nil
}

// JWTAuth JWT认证中间件
func JWTAuth() gin.HandlerFunc {
    return func(c *gin.Context) {
        authHeader := c.GetHeader("Authorization")
        if authHeader == "" {
            c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
                "code":    401,
                "message": "缺少认证信息",
            })
            return
        }
        
        parts := strings.SplitN(authHeader, " ", 2)
        if len(parts) != 2 || parts[0] != "Bearer" {
            c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
                "code":    401,
                "message": "认证格式错误",
            })
            return
        }
        
        tokenString := parts[1]
        claims, err := ParseToken(tokenString)
        if err != nil {
            if err.Error() == "token has invalid claims: token is expired" {
                c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
                    "code":    401001,
                    "message": "token已过期",
                })
            } else {
                c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
                    "code":    401002,
                    "message": "token无效",
                })
            }
            return
        }
        
        c.Set("user_id", claims.UserID)
        c.Set("username", claims.Username)
        c.Set("role", claims.Role)
        
        c.Next()
    }
}

RBAC权限控制中间件与角色路由分组

// RequireRole 角色鉴权中间件
func RequireRole(roles ...string) gin.HandlerFunc {
    return func(c *gin.Context) {
        userRole, exists := c.Get("role")
        if !exists {
            c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
                "code":    401,
                "message": "未认证",
            })
            return
        }
        
        roleStr, ok := userRole.(string)
        if !ok {
            c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{
                "code":    500,
                "message": "角色信息异常",
            })
            return
        }
        
        allowed := false
        for _, r := range roles {
            if r == roleStr {
                allowed = true
                break
            }
        }
        
        if !allowed {
            c.AbortWithStatusJSON(http.StatusForbidden, gin.H{
                "code":    403,
                "message": "权限不足",
            })
            return
        }
        
        c.Next()
    }
}

// RequirePermission 细粒度权限检查中间件
func RequirePermission(permission string) gin.HandlerFunc {
    return func(c *gin.Context) {
        userID, _ := c.Get("user_id")
        
        permissions, err := getUserPermissions(userID.(int64))
        if err != nil {
            c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{
                "code":    500,
                "message": "权限查询失败",
            })
            return
        }
        
        hasPerm := false
        for _, p := range permissions {
            if p == permission || p == "*" {
                hasPerm = true
                break
            }
        }
        
        if !hasPerm {
            c.AbortWithStatusJSON(http.StatusForbidden, gin.H{
                "code":    403,
                "message": "缺少权限: " + permission,
            })
            return
        }
        
        c.Next()
    }
}

func getUserPermissions(userID int64) ([]string, error) {
    permMap := map[int64][]string{
        1: {"user:read", "user:write", "user:delete", "*"},
        2: {"user:read", "user:write"},
        3: {"user:read"},
    }
    return permMap[userID], nil
}

中间件链组装与路由分组配置

func SetupRouter() *gin.Engine {
    r := gin.New()
    
    // 全局中间件
    r.Use(gin.Recovery())
    r.Use(Logger())
    r.Use(corsMiddleware())
    r.Use(rateLimitMiddleware(100, time.Minute))
    
    // 公开路由组
    public := r.Group("/api/v1")
    {
        public.POST("/login", handleLogin)
        public.POST("/register", handleRegister)
        public.GET("/health", handleHealthCheck)
    }
    
    // 认证路由组
    authed := r.Group("/api/v1")
    authed.Use(middleware.JWTAuth())
    {
        authed.GET("/profile", handleGetProfile)
        authed.PUT("/profile", handleUpdateProfile)
        
        editor := authed.Group("")
        editor.Use(middleware.RequireRole("editor", "admin"))
        {
            editor.GET("/posts", handleListPosts)
            editor.POST("/posts", handleCreatePost)
            editor.PUT("/posts/:id", handleUpdatePost)
        }
        
        admin := authed.Group("/admin")
        admin.Use(middleware.RequireRole("admin"))
        {
            admin.GET("/users", handleListUsers)
            admin.POST("/users", handleCreateUser)
            admin.DELETE("/users/:id", 
                middleware.RequirePermission("user:delete"), 
                handleDeleteUser)
        }
    }
    
    return r
}

// 限流中间件(基于令牌桶)
func rateLimitMiddleware(rate int, window time.Duration) gin.HandlerFunc {
    limiters := sync.Map{}
    
    return func(c *gin.Context) {
        ip := c.ClientIP()
        
        var limiter *rate.Limiter
        if val, ok := limiters.Load(ip); ok {
            limiter = val.(*rate.Limiter)
        } else {
            limiter = rate.NewLimiter(rate.Every(window/time.Duration(rate)), rate)
            limiters.Store(ip, limiter)
        }
        
        if !limiter.Allow() {
            c.Header("Retry-After", fmt.Sprintf("%d", int(window.Seconds())))
            c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{
                "code":    429,
                "message": "请求过于频繁,请稍后再试",
            })
            return
        }
        
        c.Next()
    }
}

结构化日志中间件与请求链路追踪

func Logger() gin.HandlerFunc {
    return func(c *gin.Context) {
        requestID := uuid.New().String()
        c.Set("request_id", requestID)
        c.Header("X-Request-ID", requestID)
        
        start := time.Now()
        path := c.Request.URL.Path
        method := c.Request.Method
        
        logFields := logrus.Fields{
            "request_id": requestID,
            "method":     method,
            "path":       path,
            "ip":         c.ClientIP(),
            "user_agent": c.Request.UserAgent(),
        }
        
        c.Next()
        
        latency := time.Since(start)
        status := c.Writer.Status()
        bodySize := c.Writer.Size()
        
        if userID, exists := c.Get("user_id"); exists {
            logFields["user_id"] = userID
        }
        
        logFields["status"] = status
        logFields["latency_ms"] = latency.Milliseconds()
        logFields["body_size"] = bodySize
        
        switch {
        case status >= 500:
            logrus.WithFields(logFields).Error("服务端错误")
        case status >= 400:
            logrus.WithFields(logFields).Warn("客户端错误")
        default:
            logrus.WithFields(logFields).Info("请求完成")
        }
    }
}

这套中间件链的执行顺序为:Recovery -> Logger -> CORS -> RateLimit -> JWTAuth -> RequireRole -> RequirePermission -> Handler。每一层可独立测试和替换,符合单一职责原则。Token签发使用HMAC-SHA256,生产环境中jwtSecret应从环境变量或密钥管理服务读取,不应硬编码在代码中。

原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/go-yu-yan-gin-kuang-jia-zhong-jian-jian-lian-yu-jwt-ren/

(0)
小编小编
上一篇 18小时前
下一篇 18小时前

相关推荐