Browse Source

no message

panghu 7 months ago
parent
commit
2cc8595863
6 changed files with 64 additions and 34 deletions
  1. 2 2
      config/database.go
  2. 4 4
      config/redis.go
  3. 7 0
      main.go
  4. 33 11
      middleware/api.go
  5. 1 1
      router/api_ceshi.go
  6. 17 16
      server/redisServer.go

+ 2 - 2
config/database.go

@@ -12,8 +12,8 @@ import (
 var Mdb *gorm.DB
 var err error
 
-func init() {
-	dsn := fmt.Sprintf("%s:%s@tcp(%s:%s)/%s?charset=utf8mb4&parseTime=True&loc=Local", viper.GetString("database.username"), viper.GetString("database.password"), viper.GetString("database.host"), viper.GetInt("database.port"), viper.GetString("database.name"))
+func InitMysql() {
+	dsn := fmt.Sprintf("%s:%s@tcp(%s:%s)/%s?charset=utf8mb4&parseTime=True&loc=Local", viper.GetString("database.username"), viper.GetString("database.password"), viper.GetString("database.host"), viper.GetString("database.port"), viper.GetString("database.name"))
 
 	Mdb, err = gorm.Open(mysql.Open(dsn), &gorm.Config{
 		SkipDefaultTransaction:                   true, // 禁用默认事务(提高运行速度)

+ 4 - 4
config/redis.go

@@ -2,6 +2,7 @@ package config
 
 import (
 	"context"
+	"fmt"
 	"github.com/redis/go-redis/v9"
 	"github.com/spf13/viper"
 	"go.uber.org/zap"
@@ -10,15 +11,14 @@ import (
 var ctx = context.Background()
 var Rdb *redis.Client
 
-func init() {
+func InitRedis() {
 	Rdb = redis.NewClient(&redis.Options{
-		Addr:     viper.GetString("redis.host"),
+		Addr:     fmt.Sprintf("%s:%s", viper.GetString("redis.host"), viper.GetString("redis.port")),
 		Password: viper.GetString("redis.password"), // no password set
 		DB:       viper.GetInt("redis.db"),          // use default DB
 	})
 
-	_, err := Rdb.Ping(ctx).Result()
-	if err != nil {
+	if err = Rdb.Ping(ctx).Err(); err != nil {
 		zap.L().Error("Redis连接出错 " + err.Error())
 		panic("Redis连接出错,请检查参数: " + err.Error())
 	}

+ 7 - 0
main.go

@@ -6,6 +6,7 @@ import (
 	"github.com/gin-gonic/gin"
 	"github.com/spf13/viper"
 	"go.uber.org/zap"
+	"go_zh/config"
 	"go_zh/middleware"
 	"go_zh/pkg/logger"
 	"go_zh/router"
@@ -31,6 +32,10 @@ func main() {
 	}
 	// 确保程序退出前刷新日志缓冲区
 	defer logger.Sync()
+	//初始化mysql
+	config.InitMysql()
+	//初始化redis
+	config.InitRedis()
 	// 设置 Gin 模式
 	gin.SetMode(viper.GetString("app.mode"))
 	// 初始化Gin引擎
@@ -39,6 +44,8 @@ func main() {
 	engine.Use(middleware.GinLoggerMiddleware())   //日志中间件
 	engine.Use(middleware.GinRecoveryMiddleware()) //recover中间件
 
+	engine.Use(middleware.ApiMiddleware()) //api请求中间件
+
 	// 设置路由
 	router.SetupRoutes(engine)
 

+ 33 - 11
middleware/api.go

@@ -1,8 +1,10 @@
 package middleware
 
 import (
+	"encoding/json"
 	"github.com/gin-gonic/gin"
 	"go_zh/model"
+	"go_zh/server"
 	"net/http"
 )
 
@@ -25,23 +27,43 @@ func ApiMiddleware() gin.HandlerFunc {
 		}
 		if token == "" {
 			c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
-				"code":    401,
+				"code":    404,
 				"message": "缺少认证信息",
 				"data":    nil,
 			})
 			return
 		}
-		//根据token获取用户信息
-		memberModel := model.XMember{}
-		member, err := memberModel.GetMemberByToken(token)
-		if err != nil {
-			c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
-				"code":    401,
-				"message": "认证信息已过期",
-				"data":    nil,
-			})
-		}
+		// 1. 尝试从Redis缓存获取用户信息
+		var member *model.XMember
+		var err error
+
+		redisServer := server.GetRedisServerInstance()
+		data := redisServer.GetStr(token)
+		if len(data) == 0 { //没用缓存信息
+			//根据token获取用户信息
+			memberModel := model.XMember{}
+			member, err = memberModel.GetMemberByToken(token)
+			if err != nil {
+				c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
+					"code":    401,
+					"message": "认证信息已过期",
+					"data":    nil,
+				})
+			}
+			if member.ID <= 0 {
+				c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
+					"code":    401,
+					"message": "token失效",
+					"data":    nil,
+				})
+			}
+			//存入缓存
+			jsonData, _ := json.Marshal(member)
+			redisServer.SetStr(token, string(jsonData), 120)
+		} else {
+			_ = json.Unmarshal([]byte(data), &member)
 
+		}
 		c.Set("mId", member.ID)
 		c.Next()
 	}

+ 1 - 1
router/api_ceshi.go

@@ -7,7 +7,7 @@ import (
 
 func init() {
 	RegisterRoute(func(engine *gin.Engine) {
-		v1 := engine.Group("/v1")
+		v1 := engine.Group("/api")
 		// 用户路由组
 		ceshiGroup := v1.Group("/ceshi")
 		{

+ 17 - 16
server/redisServer.go

@@ -4,30 +4,31 @@ import (
 	"context"
 	"github.com/redis/go-redis/v9"
 	"go_zh/config"
+	"sync"
 	"time"
 )
 
-type RedisServer struct {
+type redisServer struct {
 	rdb *redis.Client
 }
 
-// 定义单例实例,但不在包级别初始化
-var redisServerInstance *RedisServer
+var (
+	redisInstance *redisServer
+	redisOnce     sync.Once
+)
 var ctx = context.Background()
 
-// 使用init函数初始化单例实例
-func init() {
-	redisServerInstance = &RedisServer{
-		rdb: config.Rdb,
-	}
+// GetRedisServerInstance 获取单例实例
+func GetRedisServerInstance() *redisServer {
+	redisOnce.Do(func() {
+		redisInstance = &redisServer{
+			rdb: config.Rdb,
+		}
+	})
+	return redisInstance
 }
 
-// 获取单例实例的方法
-func GetRedisServerInstance() *RedisServer {
-	return redisServerInstance
-}
-
-func (r *RedisServer) SetStr(key string, val string, timeout int) bool {
+func (r *redisServer) SetStr(key string, val string, timeout int) bool {
 	err := r.rdb.Set(ctx, key, val, time.Duration(timeout)*time.Second).Err()
 	if err != nil {
 		return false
@@ -35,11 +36,11 @@ func (r *RedisServer) SetStr(key string, val string, timeout int) bool {
 	return true
 }
 
-func (r *RedisServer) GetStr(key string) string {
+func (r *redisServer) GetStr(key string) string {
+
 	if key == "" {
 		return ""
 	}
-	ctx := context.Background()
 	val, err := r.rdb.Get(ctx, key).Result()
 	if err != nil {
 		return ""