Files
Meshray-Manager/internal/api/middleware/auth.go
T
2026-06-30 15:14:37 +08:00

328 lines
7.7 KiB
Go
Raw 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 (
"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("内部错误")