199 lines
5.1 KiB
Go
199 lines
5.1 KiB
Go
package middleware
|
||
|
||
import (
|
||
"log"
|
||
"net/http"
|
||
"strings"
|
||
"sync"
|
||
|
||
"server/pkg/jwtutil"
|
||
|
||
beego "github.com/beego/beego/v2/server/web"
|
||
"github.com/beego/beego/v2/server/web/context"
|
||
)
|
||
|
||
// 鉴权开关(app.conf: auth_enforce)
|
||
// - off :完全不启用(等同改造前的行为)
|
||
// - warn :观察模式,只记录「无 token / token 无效」的接口,不拦截(默认,零风险上线)
|
||
// - on :真正拦截并返回 401
|
||
//
|
||
// 上线建议:先 warn 跑一段时间,用日志核对白名单是否遗漏,确认无误后再切 on。
|
||
const (
|
||
AuthModeOff = "off"
|
||
AuthModeWarn = "warn"
|
||
AuthModeOn = "on"
|
||
)
|
||
|
||
// protectedPrefixes 需要鉴权的业务前缀。
|
||
// 仅保护业务 API:官网(/ /site/ /index/…)、开放 API(/api/…)、
|
||
// 静态资源与 ACME 验证路径均不在其中,避免误伤公开接口。
|
||
var protectedPrefixes = []string{"/backend/", "/platform/", "/app/"}
|
||
|
||
// authWhitelist 受保护前缀内无需鉴权的路径(前缀匹配)
|
||
var authWhitelist = []string{
|
||
// 统一认证中心(P1 上线)
|
||
"/auth/",
|
||
|
||
// backend 登录/注册/找回密码
|
||
"/backend/login",
|
||
"/backend/sendLoginCode",
|
||
"/backend/register",
|
||
"/backend/sendRegisterCode",
|
||
"/backend/resetPassword",
|
||
"/backend/sendResetCode",
|
||
"/backend/verifyAccount",
|
||
"/backend/logout",
|
||
|
||
// platform 登录/找回密码
|
||
"/platform/login",
|
||
"/platform/sendLoginCode",
|
||
"/platform/resetPassword",
|
||
"/platform/logout",
|
||
// 平台开放接口(客户端更新通知,无需登录)
|
||
"/platform/api/",
|
||
|
||
// app 登录/注册/找回密码
|
||
"/app/login",
|
||
"/app/sendLoginCode",
|
||
"/app/register",
|
||
"/app/sendRegisterCode",
|
||
"/app/resetPassword",
|
||
"/app/sendResetCode",
|
||
"/app/verifyAccount",
|
||
"/app/logout",
|
||
}
|
||
|
||
var (
|
||
modeOnce sync.Once
|
||
modeVal = AuthModeWarn
|
||
warnSeen sync.Map // 观察模式下按 path 去重,避免日志刷屏
|
||
)
|
||
|
||
func authMode() string {
|
||
modeOnce.Do(func() {
|
||
if v, err := beego.AppConfig.String("auth_enforce"); err == nil {
|
||
switch strings.ToLower(strings.TrimSpace(v)) {
|
||
case AuthModeOff, AuthModeWarn, AuthModeOn:
|
||
modeVal = strings.ToLower(strings.TrimSpace(v))
|
||
}
|
||
}
|
||
// 额外白名单:auth_whitelist = /backend/xxx,/app/yyy
|
||
if extra, err := beego.AppConfig.String("auth_whitelist"); err == nil && strings.TrimSpace(extra) != "" {
|
||
for _, p := range strings.Split(extra, ",") {
|
||
if p = strings.TrimSpace(p); p != "" {
|
||
authWhitelist = append(authWhitelist, p)
|
||
}
|
||
}
|
||
}
|
||
log.Printf("[auth] JWT 全局鉴权模式=%s,受保护前缀=%v", modeVal, protectedPrefixes)
|
||
})
|
||
return modeVal
|
||
}
|
||
|
||
// cleanPath 去掉 query / fragment 后返回纯路径
|
||
func cleanPath(uri string) string {
|
||
if i := strings.IndexAny(uri, "?#"); i >= 0 {
|
||
return uri[:i]
|
||
}
|
||
return uri
|
||
}
|
||
|
||
func isProtected(path string) bool {
|
||
for _, p := range protectedPrefixes {
|
||
if strings.HasPrefix(path, p) {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
func isWhitelisted(path string) bool {
|
||
for _, p := range authWhitelist {
|
||
if strings.HasPrefix(path, p) {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
// warnLog 观察模式下记录一次「本会被拦截」的请求(按路径去重)
|
||
func warnLog(path, reason string) {
|
||
if _, loaded := warnSeen.LoadOrStore(path, struct{}{}); loaded {
|
||
return
|
||
}
|
||
log.Printf("[auth][warn] 该接口在 enforce=on 时会被拦截: %s (%s)", path, reason)
|
||
}
|
||
|
||
// JWTAuthMiddleware JWT认证中间件
|
||
func JWTAuthMiddleware() beego.FilterFunc {
|
||
return func(ctx *context.Context) {
|
||
mode := authMode()
|
||
if mode == AuthModeOff {
|
||
return
|
||
}
|
||
// 预检请求直接放行
|
||
if ctx.Input.Method() == http.MethodOptions {
|
||
return
|
||
}
|
||
|
||
path := cleanPath(ctx.Request.RequestURI)
|
||
if !isProtected(path) || isWhitelisted(path) {
|
||
return
|
||
}
|
||
|
||
// 失败处理:观察模式只记日志,强制模式才真正拦截
|
||
reject := func(status int, message string) {
|
||
if mode == AuthModeWarn {
|
||
warnLog(path, message)
|
||
return
|
||
}
|
||
ctx.Output.SetStatus(status)
|
||
ctx.Output.JSON(map[string]interface{}{
|
||
"success": false,
|
||
"message": message,
|
||
}, false, false)
|
||
}
|
||
|
||
// 从请求头中获取Authorization
|
||
authHeader := ctx.Request.Header.Get("Authorization")
|
||
if authHeader == "" {
|
||
reject(401, "未提供认证信息")
|
||
return
|
||
}
|
||
|
||
// 按空格分割
|
||
authParts := strings.SplitN(authHeader, " ", 2)
|
||
if !(len(authParts) == 2 && authParts[0] == "Bearer") {
|
||
reject(401, "认证信息格式错误")
|
||
return
|
||
}
|
||
|
||
// 解析token
|
||
claims, err := jwtutil.ParseToken(authParts[1])
|
||
if err != nil {
|
||
message := "无效的token"
|
||
if strings.Contains(err.Error(), "expired") {
|
||
message = "token已过期"
|
||
}
|
||
reject(401, message)
|
||
return
|
||
}
|
||
|
||
// 将用户信息存储在上下文
|
||
ctx.Input.SetData("userId", claims.UserID)
|
||
ctx.Input.SetData("username", claims.Username)
|
||
ctx.Input.SetData("tenantId", claims.TenantId)
|
||
|
||
// 从token中获取用户类型(如果token中没有,则默认为"user")
|
||
userType := claims.UserType
|
||
if userType == "" {
|
||
userType = "user"
|
||
}
|
||
ctx.Input.SetData("userType", userType)
|
||
|
||
// 认证中心扩展字段(旧 token 为空值,不影响现有逻辑)
|
||
ctx.Input.SetData("clientId", claims.ClientID)
|
||
ctx.Input.SetData("sid", claims.Sid)
|
||
}
|
||
}
|