From bdb17437bf90f8d5734c694e05ade206375b0b16 Mon Sep 17 00:00:00 2001 From: zyj <18107291228@163.com> Date: Thu, 11 Sep 2025 18:15:45 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E6=94=B9websockt=E8=BA=AB=E4=BB=BD?= =?UTF-8?q?=E9=AA=8C=E8=AF=81=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/websocks.go | 47 ++++++++++++++++++++++++++++++++----- libs/error_code.go | 3 +++ models/profile.go | 20 ++++++++++++++++ router/test.go | 3 ++- services/message_service.go | 26 ++++++++++++++++++++ services/profile_service.go | 32 +++++++++++++++++++++++++ 6 files changed, 124 insertions(+), 7 deletions(-) create mode 100644 models/profile.go create mode 100644 services/message_service.go create mode 100644 services/profile_service.go diff --git a/api/websocks.go b/api/websocks.go index d79806f..5f813a8 100644 --- a/api/websocks.go +++ b/api/websocks.go @@ -1,9 +1,14 @@ package api import ( + "context" + grpcClient "go-desk-service/grpc-client" + "go-desk-service/libs" + userpb "go-desk-service/proto/gen" "log" "net/http" "sync" + "time" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" @@ -12,8 +17,12 @@ import ( type Websocks struct { } -var clients = make(map[*websocket.Conn]bool) // 已连接的客户端 -var clientsMu sync.Mutex // 保护 clients 映射的互斥锁 +type Message struct { +} + +var Clients = make(map[int64]map[string]*websocket.Conn) // 已连接的客户端 + +var ClientsMu sync.Mutex // 保护 clients 映射的互斥锁 var upgrader = websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { @@ -22,17 +31,43 @@ var upgrader = websocket.Upgrader{ } func (*Websocks) Init(ctx *gin.Context) { + WebsocktFailed := libs.ErrorCode["WebsocktFailed"] + // 判断当前连接是否合法 + tokenStr := ctx.Request.Header.Get("Sec-WebSocket-Protocol") + if tokenStr == "" { + ctx.JSON(http.StatusInternalServerError, gin.H{"code": WebsocktFailed.Code, "data": WebsocktFailed.Data, "msg": WebsocktFailed.Msg}) + return + } + client := grpcClient.GetUserClient() + ctx1, cancel := context.WithTimeout(context.Background(), time.Second*5) + defer cancel() + TokenStatus, err1 := client.ValidateToken(ctx1, &userpb.ValidateTokenRequest{ + AccessToken: tokenStr, + }) + // 验证当前token + if err1 != nil { + ctx.JSON(http.StatusInternalServerError, gin.H{"code": WebsocktFailed.Code, "data": WebsocktFailed.Data, "msg": WebsocktFailed.Msg}) + return + } + if !TokenStatus.IsValid { + ctx.JSON(http.StatusInternalServerError, gin.H{"code": WebsocktFailed.Code, "data": WebsocktFailed.Data, "msg": WebsocktFailed.Msg}) + return + } + conn, err := upgrader.Upgrade(ctx.Writer, ctx.Request, nil) + if err != nil { log.Printf("WebSocket 升级失败: %v", err) - ctx.JSON(http.StatusInternalServerError, gin.H{"error": "无法建立 WebSocket 连接"}) + + ctx.JSON(http.StatusInternalServerError, gin.H{"code": WebsocktFailed.Code, "data": WebsocktFailed.Data, "msg": WebsocktFailed.Msg}) return } defer conn.Close() + // 将保存连接状态 - clientsMu.Lock() - clients[conn] = true - clientsMu.Unlock() + ClientsMu.Lock() + Clients[TokenStatus.UserId] = conn + ClientsMu.Unlock() log.Printf("客户端已连接: %s", conn.RemoteAddr()) for { mt, message, err := conn.ReadMessage() diff --git a/libs/error_code.go b/libs/error_code.go index 42bd307..62bb19a 100644 --- a/libs/error_code.go +++ b/libs/error_code.go @@ -20,4 +20,7 @@ var ErrorCode = map[string]*ErrorInfo{ "BrowserOpenFail": {Code: 0, Data: "", Msg: "浏览器开启失败"}, "AccountFailed": {Code: 0, Data: "", Msg: "账号不存在"}, "TokenExpired": {Code: 0, Data: "", Msg: "Token已过期"}, + "LoginError": {Code: 0, Data: "", Msg: "登录失败"}, + "RegistrationFailed": {Code: 0, Data: "", Msg: "注册失败"}, + "WebsocktFailed": {Code: 0, Data: "", Msg: "WebSocket 连接失败"}, } diff --git a/models/profile.go b/models/profile.go new file mode 100644 index 0000000..7f49e0d --- /dev/null +++ b/models/profile.go @@ -0,0 +1,20 @@ +package models + +import "time" + +type Profile struct { + ID int8 `gorm:"id;primary_key"` + UserId int8 `gorm:"user_id"` + Nickname string `gorm:"nickname"` + Email string `gorm:"email"` + Phone string `gorm:"phone"` + Status int `gorm:"status"` + LoginStatus int `gorm:"login_status"` + CreatedAt time.Time `gorm:"created_at;type:timestamptz"` + UpdatedAt time.Time `gorm:"updated_at;type:timestamptz"` +} + +// 实现 TableName 方法指定表名 +func (Profile) TableName() string { + return "gd_profile" +} diff --git a/router/test.go b/router/test.go index a189d3c..f6f3feb 100644 --- a/router/test.go +++ b/router/test.go @@ -2,6 +2,7 @@ package router import ( "go-desk-service/api" + "go-desk-service/middleware" "github.com/gin-gonic/gin" ) @@ -11,6 +12,6 @@ type Test struct{} var apiTest api.Test func (*Test) Init(app *gin.Engine) { - test := app.Group("/test") + test := app.Group("/test", middleware.TokenAuth()) test.GET("/status", apiTest.Status) } diff --git a/services/message_service.go b/services/message_service.go new file mode 100644 index 0000000..9bda591 --- /dev/null +++ b/services/message_service.go @@ -0,0 +1,26 @@ +package services + +import ( + "log" + + "github.com/gorilla/websocket" +) + +type MessageService struct{} + +// websockt消息处理 +func (*MessageService) WebSocketMessage(conn *websocket.Conn) { + for { + mt, message, err := conn.ReadMessage() + if err != nil { + log.Printf("读取错误: %v (客户端: %s)", err, conn.RemoteAddr()) + break + } + log.Printf("收到来自 %s 的消息: %s", conn.RemoteAddr(), string(message)) + err = conn.WriteMessage(mt, message) + if err != nil { + log.Println("write:", err) + break + } + } +} diff --git a/services/profile_service.go b/services/profile_service.go new file mode 100644 index 0000000..c52ac3c --- /dev/null +++ b/services/profile_service.go @@ -0,0 +1,32 @@ +package services + +import ( + "go-desk-service/libs" + "go-desk-service/models" + + "gorm.io/gorm" +) + +type ProfileService struct { + db *gorm.DB +} + +func InitProfileService() *ProfileService { + db := libs.GetDB() + return &ProfileService{ + db: db, + } +} + +// 获取个人信息 +func (p *ProfileService) GetInfo(params *models.Profile) models.Profile { + var user models.Profile + p.db.Model(&models.Profile{}).Where(params).First(&user) + return user +} + +// 更新账号信息 +func (p *ProfileService) UpdateAccount(userId int8, params *models.Profile) bool { + p.db.Model(&models.Profile{}).Where("user_id = ?", userId).Updates(params) + return true +}