package middleware import ( "errors" "net/http" "strings" "time" "git.zkcoi.com/zkcoi/meshray/internal/service" "git.zkcoi.com/zkcoi/meshray/internal/store/sqlite" "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" "go.uber.org/zap" ) // Claims JWT 声明 type Claims struct { UserID uint `json:"user_id"` Username string `json:"username"` TokenType string `json:"token_type"` // access | refresh jwt.RegisteredClaims } // Response 统一响应格式 type Response struct { Code int `json:"code"` Message string `json:"message"` Data interface{} `json:"data,omitempty"` } // Success 成功响应 func Success(c *gin.Context, data interface{}) { c.JSON(http.StatusOK, Response{ Code: 0, Message: "success", Data: data, }) } // Error 错误响应 func Error(c *gin.Context, code int, message string) { c.JSON(http.StatusOK, Response{ Code: code, Message: message, }) } // GenerateToken 生成 JWT Token func GenerateToken(secret string, userID uint, username string, tokenType string, duration time.Duration) (string, error) { claims := Claims{ UserID: userID, Username: username, TokenType: tokenType, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(time.Now().Add(duration)), IssuedAt: jwt.NewNumericDate(time.Now()), Issuer: "meshray", }, } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) return token.SignedString([]byte(secret)) } // JWTAuth JWT 鉴权中间件 func JWTAuth(secret string) gin.HandlerFunc { return func(c *gin.Context) { var tokenString string // 1. 优先从 Authorization Header 获取 authHeader := c.GetHeader("Authorization") if authHeader != "" { parts := strings.SplitN(authHeader, " ", 2) if len(parts) == 2 && parts[0] == "Bearer" { tokenString = parts[1] } } // 2. 如果 Header 中没有,尝试从 query 参数获取(WebSocket 场景) if tokenString == "" { tokenString = c.Query("token") } // 3. 都没有则拒绝 if tokenString == "" { Error(c, 401, "未提供认证令牌") c.Abort() return } // 解析和验证 Token token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) { return []byte(secret), nil }) if err != nil { Error(c, 401, "无效的认证令牌") c.Abort() return } if !token.Valid { Error(c, 401, "认证令牌已过期或无效") c.Abort() return } claims, ok := token.Claims.(*Claims) if !ok { Error(c, 401, "无法解析认证声明") c.Abort() return } // 检查 Token 类型 if claims.TokenType != "access" { Error(c, 403, "Token 类型错误") c.Abort() return } // 将用户信息存入上下文 c.Set("user_id", claims.UserID) c.Set("username", claims.Username) c.Next() } } // LoginRequest 登录请求 type LoginRequest struct { Username string `json:"username" binding:"required"` Password string `json:"password" binding:"required"` } // LoginResponse 登录响应 type LoginResponse struct { AccessToken string `json:"access_token"` RefreshToken string `json:"refresh_token"` ExpiresIn int64 `json:"expires_in"` } // LoginHandler 登录 Handler func LoginHandler(jwtSecret string, logger *zap.Logger, store *sqlite.Store) gin.HandlerFunc { return func(c *gin.Context) { var req LoginRequest if err := c.ShouldBindJSON(&req); err != nil { Error(c, 400, "请求参数错误") c.Abort() return } logger.Info("Login request received", zap.String("username", req.Username)) // 创建 UserService 并验证用户 userService := service.NewUserService(store) user, err := userService.Authenticate(req.Username, req.Password) if err != nil { logger.Warn("Authentication failed", zap.String("username", req.Username), zap.Error(err)) Error(c, 401, err.Error()) c.Abort() return } // 生成 Access Token token, err := GenerateToken( jwtSecret, user.ID, user.Username, "access", 2*time.Hour, ) if err != nil { logger.Error("Failed to generate token", zap.Error(err)) Error(c, 500, "无法生成访问令牌") c.Abort() return } // 生成 Refresh Token refreshToken, err := GenerateToken( jwtSecret, user.ID, user.Username, "refresh", 7*24*time.Hour, ) if err != nil { logger.Error("Failed to generate refresh token", zap.Error(err)) Error(c, 500, "无法生成刷新令牌") c.Abort() return } logger.Info("User login successful", zap.String("username", user.Username), zap.Uint("user_id", user.ID)) Success(c, LoginResponse{ AccessToken: token, RefreshToken: refreshToken, ExpiresIn: 7200, }) } } // RefreshTokenRequest 刷新 Token 请求 type RefreshTokenRequest struct { RefreshToken string `json:"refresh_token" binding:"required"` } // RefreshTokenHandler 刷新 Token Handler func RefreshTokenHandler(jwtSecret string, logger *zap.Logger) gin.HandlerFunc { return func(c *gin.Context) { var req RefreshTokenRequest if err := c.ShouldBindJSON(&req); err != nil { Error(c, 400, "请求参数错误") c.Abort() return } logger.Info("Refresh token request received") // 验证刷新 Token token, err := jwt.ParseWithClaims(req.RefreshToken, &Claims{}, func(token *jwt.Token) (interface{}, error) { return []byte(jwtSecret), nil }) if err != nil || !token.Valid { Error(c, 401, "无效的刷新令牌") c.Abort() return } claims, ok := token.Claims.(*Claims) if !ok || claims.TokenType != "refresh" { Error(c, 403, "Token 类型错误") c.Abort() return } // 生成新的访问令牌 newAccessToken, err := GenerateToken( jwtSecret, claims.UserID, claims.Username, "access", 2*time.Hour, ) if err != nil { logger.Error("Failed to generate new access token", zap.Error(err)) Error(c, 500, "无法生成访问令牌") c.Abort() return } Success(c, gin.H{ "access_token": newAccessToken, "expires_in": 7200, }) } } // RequestLogger 请求日志中间件 func RequestLogger(logger *zap.Logger) gin.HandlerFunc { return func(c *gin.Context) { start := time.Now() path := c.Request.URL.Path query := c.Request.URL.RawQuery c.Next() latency := time.Since(start) statusCode := c.Writer.Status() logger.Info("HTTP request", zap.Int("status", statusCode), zap.String("method", c.Request.Method), zap.String("path", path), zap.String("query", query), zap.String("ip", c.ClientIP()), zap.String("user_agent", c.Request.UserAgent()), zap.Duration("latency", latency), ) } } // CORS CORS 中间件 func CORS() gin.HandlerFunc { // 允许的 Origin 白名单 allowedOrigins := map[string]bool{ "http://localhost:9531": true, "http://127.0.0.1:9531": true, // 生产环境可以添加域名 // "https://meshray.example.com": true, } return func(c *gin.Context) { origin := c.Request.Header.Get("Origin") // ✅ 修复:检查 Origin 是否在白名单内 if !allowedOrigins[origin] { // 不在白名单,使用默认值(不设置 Access-Control-Allow-Origin) c.Next() return } // 在白名单内,设置 CORS 头 c.Writer.Header().Set("Access-Control-Allow-Origin", origin) c.Writer.Header().Set("Access-Control-Allow-Credentials", "true") c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With") c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE") if c.Request.Method == "OPTIONS" { c.AbortWithStatus(http.StatusNoContent) return } c.Next() } } // InternalError 内部错误 var InternalError = errors.New("内部错误")