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