package middleware import ( "cn/oneao/base-go/utils/jwtx" "cn/oneao/base-go/utils/pathx" "cn/oneao/base-go/utils/response" "context" "strconv" "github.com/cloudwego/hertz/pkg/app" ) // 私有 key,防止冲突 type ctxKeyUser struct{} // ContextUser 保存用户信息,包括 UserID 和 SpaceID type ContextUser struct { UserID int64 SpaceID int64 } // Auth 中间件,直接在中间件里过滤白名单路径 func Auth() app.HandlerFunc { whitelistPaths := []string{ "/auth/login", "/auth/register", "/auth/refresh/token", } return func(ctx context.Context, c *app.RequestContext) { path := string(c.Request.URI().Path()) if pathx.MatchPath(path, whitelistPaths) { c.Next(ctx) return } token := string(c.GetHeader("token")) spaceId := string(c.GetHeader("x-space-id")) if token == "" { response.Fail(c).Code(response.HttpCode.Unauthorized).Send() c.Abort() return } verifyToken := jwtx.VerifyToken(token) if verifyToken.IsValid { user := &ContextUser{ UserID: verifyToken.Claims.UserID, SpaceID: 0, } if spaceId != "" { i, err := strconv.ParseInt(spaceId, 10, 64) if err != nil { response.Fail(c).Message("Invalid x-space-id").Send() c.Abort() return } user.SpaceID = i } ctx = setContextUser(ctx, user) c.Next(ctx) return } response.Fail(c).Code(response.HttpCode.RefreshToken).Send() c.Abort() } } // GetContextUser 从 context 获取用户信息 func GetContextUser(ctx context.Context) *ContextUser { user, ok := ctx.Value(ctxKeyUser{}).(*ContextUser) if !ok || user == nil { return &ContextUser{} } return user } // GetUserID 从 context 获取 user_id func GetUserID(ctx context.Context) int64 { return GetContextUser(ctx).UserID } // GetSpaceID 从 context 获取 space_id,并返回是否有效 func GetSpaceID(ctx context.Context) (int64, bool) { spaceID := GetContextUser(ctx).SpaceID return spaceID, spaceID != 0 } // 内部中间件设置 ContextUser func setContextUser(ctx context.Context, user *ContextUser) context.Context { return context.WithValue(ctx, ctxKeyUser{}, user) }