修改websockt身份验证处理
This commit is contained in:
@@ -1,9 +1,14 @@
|
|||||||
package api
|
package api
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
grpcClient "go-desk-service/grpc-client"
|
||||||
|
"go-desk-service/libs"
|
||||||
|
userpb "go-desk-service/proto/gen"
|
||||||
"log"
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
@@ -12,8 +17,12 @@ import (
|
|||||||
type Websocks struct {
|
type Websocks struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
var clients = make(map[*websocket.Conn]bool) // 已连接的客户端
|
type Message struct {
|
||||||
var clientsMu sync.Mutex // 保护 clients 映射的互斥锁
|
}
|
||||||
|
|
||||||
|
var Clients = make(map[int64]map[string]*websocket.Conn) // 已连接的客户端
|
||||||
|
|
||||||
|
var ClientsMu sync.Mutex // 保护 clients 映射的互斥锁
|
||||||
|
|
||||||
var upgrader = websocket.Upgrader{
|
var upgrader = websocket.Upgrader{
|
||||||
CheckOrigin: func(r *http.Request) bool {
|
CheckOrigin: func(r *http.Request) bool {
|
||||||
@@ -22,17 +31,43 @@ var upgrader = websocket.Upgrader{
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (*Websocks) Init(ctx *gin.Context) {
|
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)
|
conn, err := upgrader.Upgrade(ctx.Writer, ctx.Request, nil)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Printf("WebSocket 升级失败: %v", err)
|
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
|
return
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
// 将保存连接状态
|
// 将保存连接状态
|
||||||
clientsMu.Lock()
|
ClientsMu.Lock()
|
||||||
clients[conn] = true
|
Clients[TokenStatus.UserId] = conn
|
||||||
clientsMu.Unlock()
|
ClientsMu.Unlock()
|
||||||
log.Printf("客户端已连接: %s", conn.RemoteAddr())
|
log.Printf("客户端已连接: %s", conn.RemoteAddr())
|
||||||
for {
|
for {
|
||||||
mt, message, err := conn.ReadMessage()
|
mt, message, err := conn.ReadMessage()
|
||||||
|
|||||||
@@ -20,4 +20,7 @@ var ErrorCode = map[string]*ErrorInfo{
|
|||||||
"BrowserOpenFail": {Code: 0, Data: "", Msg: "浏览器开启失败"},
|
"BrowserOpenFail": {Code: 0, Data: "", Msg: "浏览器开启失败"},
|
||||||
"AccountFailed": {Code: 0, Data: "", Msg: "账号不存在"},
|
"AccountFailed": {Code: 0, Data: "", Msg: "账号不存在"},
|
||||||
"TokenExpired": {Code: 0, Data: "", Msg: "Token已过期"},
|
"TokenExpired": {Code: 0, Data: "", Msg: "Token已过期"},
|
||||||
|
"LoginError": {Code: 0, Data: "", Msg: "登录失败"},
|
||||||
|
"RegistrationFailed": {Code: 0, Data: "", Msg: "注册失败"},
|
||||||
|
"WebsocktFailed": {Code: 0, Data: "", Msg: "WebSocket 连接失败"},
|
||||||
}
|
}
|
||||||
|
|||||||
20
models/profile.go
Normal file
20
models/profile.go
Normal file
@@ -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"
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@ package router
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"go-desk-service/api"
|
"go-desk-service/api"
|
||||||
|
"go-desk-service/middleware"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
@@ -11,6 +12,6 @@ type Test struct{}
|
|||||||
var apiTest api.Test
|
var apiTest api.Test
|
||||||
|
|
||||||
func (*Test) Init(app *gin.Engine) {
|
func (*Test) Init(app *gin.Engine) {
|
||||||
test := app.Group("/test")
|
test := app.Group("/test", middleware.TokenAuth())
|
||||||
test.GET("/status", apiTest.Status)
|
test.GET("/status", apiTest.Status)
|
||||||
}
|
}
|
||||||
|
|||||||
26
services/message_service.go
Normal file
26
services/message_service.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
32
services/profile_service.go
Normal file
32
services/profile_service.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user