56 lines
1.2 KiB
Go
56 lines
1.2 KiB
Go
package middleware
|
|
|
|
import (
|
|
errorcode "ai_scheduler/internal/data/error"
|
|
"strings"
|
|
|
|
"ai_scheduler/internal/pkg"
|
|
|
|
"github.com/gofiber/fiber/v2"
|
|
)
|
|
|
|
// AuthMiddleware 登录校验中间件
|
|
func AuthMiddleware(secret string, whiteList ...string) fiber.Handler {
|
|
whiteMap := make(map[string]struct{}, len(whiteList))
|
|
for _, p := range whiteList {
|
|
whiteMap[p] = struct{}{}
|
|
}
|
|
|
|
return func(c *fiber.Ctx) error {
|
|
if _, ok := whiteMap[c.Path()]; ok {
|
|
return c.Next()
|
|
}
|
|
|
|
tokenStr := extractToken(c)
|
|
if tokenStr == "" {
|
|
return pkg.HandleResponse(c, nil, errorcode.AuthNotFound)
|
|
}
|
|
|
|
claims, err := pkg.ParseToken(tokenStr, secret)
|
|
if err != nil {
|
|
return pkg.HandleResponse(c, nil, errorcode.AuthNotFound)
|
|
}
|
|
|
|
c.Locals("user_id", claims.UserID)
|
|
c.Locals("username", claims.Username)
|
|
c.Locals("claims", claims)
|
|
|
|
return c.Next()
|
|
}
|
|
}
|
|
|
|
func extractToken(c *fiber.Ctx) string {
|
|
auth := c.Get("Authorization")
|
|
if auth != "" {
|
|
parts := strings.SplitN(auth, " ", 2)
|
|
if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
|
|
return strings.TrimSpace(parts[1])
|
|
}
|
|
return strings.TrimSpace(auth)
|
|
}
|
|
if t := c.Query("token"); t != "" {
|
|
return t
|
|
}
|
|
return ""
|
|
}
|