Files
nl-im-websocket-demo/controller/websocket_controller.go
2025-07-03 09:16:16 +08:00

565 lines
18 KiB
Go
Raw Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package controller
import (
"context"
"encoding/json"
"fmt"
"log"
"net/http"
"os"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/go-redis/redis/v8"
"github.com/gorilla/websocket"
"xk-websocket/models" // 替换为实际路径
"xk-websocket/utils" // 导入新的utils包
)
var upgrader = websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: func(r *http.Request) bool {
return true
},
}
// WebSocketController WebSocket控制器结构体
type WebSocketController struct {
Clients map[string]*websocket.Conn
ClientsMux sync.RWMutex
RedisCli *redis.Client
NodeID string
Port string
RedisCtx context.Context
Logger *log.Logger
}
func NewWebSocketController() *WebSocketController {
return &WebSocketController{
Clients: make(map[string]*websocket.Conn),
RedisCtx: context.Background(),
}
}
// ConfigureSystem 配置系统参数
func (c *WebSocketController) ConfigureSystem() {
c.NodeID = utils.GetEnv("NODE_ID", "local")
c.Port = utils.GetEnv("PORT", "12080")
log.SetPrefix(fmt.Sprintf("[Node:%s] ", c.NodeID))
log.SetFlags(log.LstdFlags | log.Lmicroseconds)
c.configureLogger()
}
// 配置日志记录器
func (c *WebSocketController) configureLogger() {
logDir := utils.GetEnv("LOG_DIR", "./logs")
if err := os.MkdirAll(logDir, 0755); err != nil {
log.Fatalf("❌❌ 创建日志目录失败: %v", err)
}
logFreq := utils.GetEnv("LOG_ROTATE_FREQ", "daily")
if c.Logger == nil {
c.Logger = log.New(os.Stdout, "", log.LstdFlags|log.Lmicroseconds)
}
c.Logger.SetPrefix(fmt.Sprintf("[Node:%s] ", c.NodeID))
log.SetOutput(c.getLogWriter(logDir, logFreq))
}
// 获取日志文件写入器
func (c *WebSocketController) getLogWriter(logDir, freq string) *utils.DailyFileWriter {
return utils.NewDailyFileWriter(logDir, freq, c.Logger, c.NodeID)
}
// InitRedisClient 初始化Redis客户端
func (c *WebSocketController) InitRedisClient() {
redisAddr := utils.GetEnv("REDIS_ADDR", "localhost:6379")
redisPassword := utils.GetEnv("REDIS_PASSWORD", "")
c.RedisCli = redis.NewClient(&redis.Options{
Addr: redisAddr,
Password: redisPassword,
DB: 0,
})
if err := c.checkRedisConnection(); err != nil {
log.Fatalf("❌❌ Redis连接失败: %v", err)
}
}
// 检查Redis连接
func (c *WebSocketController) checkRedisConnection() error {
_, err := c.RedisCli.Ping(c.RedisCtx).Result()
return err
}
// PrintStartupInfo 打印启动信息
func (c *WebSocketController) PrintStartupInfo() {
hostname, _ := os.Hostname()
log.Printf("🚀🚀 WebSocket服务启动: NodeID=%s", c.NodeID)
log.Printf("🌐🌐 监听端口: %s", c.Port)
log.Printf("📡📡 Redis地址: %s", utils.GetEnv("REDIS_ADDR", "localhost:6379"))
log.Printf("💻💻 主机: %s", hostname)
log.Printf("🕒🕒🕒 启动时间: %s", time.Now().Format("2006-01-02 15:04:05"))
log.Printf("🔔🔔 支持消息类型: \n 0:文本\n 1:图片\n 2:音频\n 3:视频\n 4:处方\n 5:病例\n 6:视频通话")
log.Printf("🔑🔑 用户绑定键: \n %s\n %s", models.ClientUserKey, models.UserClientKey)
log.Printf("📬📬 新增接口: POST /send-to-user (直接通过用户ID发送消息)")
log.Println("🔗🔗 等待客户端连接...")
}
// HealthHandler 健康检查处理器
func (c *WebSocketController) HealthHandler(ctx *gin.Context) {
ctx.JSON(http.StatusOK, gin.H{
"status": "ok",
"node": c.NodeID,
"time": time.Now().Format(time.RFC3339),
})
}
// HandleWebSocket WebSocket连接处理器
func (c *WebSocketController) HandleWebSocket(ctx *gin.Context) {
start := time.Now()
clientIP := ctx.ClientIP()
log.Printf("👤👤 客户端连接中: IP=%s", clientIP)
conn, err := upgrader.Upgrade(ctx.Writer, ctx.Request, nil)
if err != nil {
log.Printf("⚠️ WebSocket升级失败: %v | ClientIP=%s", err, clientIP)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "无法升级为WebSocket连接"})
return
}
defer conn.Close()
clientID := utils.GenerateClientID(c.NodeID)
log.Printf("✅ 客户端已连接: \n ClientID=%s \n IP=%s \n Duration=%s",
clientID, clientIP, time.Since(start))
c.addClient(clientID, conn)
defer c.removeClient(clientID)
if err := conn.WriteJSON(gin.H{"clientId": clientID}); err != nil {
log.Printf("⚠️ 发送客户端ID失败: %v | ClientID=%s", err, clientID)
return
}
conn.SetPongHandler(func(string) error {
conn.SetReadDeadline(time.Now().Add(60 * time.Second))
return nil
})
for {
messageType, message, err := conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway) {
log.Printf("❌❌ 连接意外断开: %v | ClientID=%s", err, clientID)
} else {
log.Printf("⚠️ 连接正常关闭: ClientID=%s", clientID)
}
break
}
if messageType == websocket.TextMessage {
go c.handleClientMessage(clientID, message)
}
}
}
// 处理客户端消息
func (c *WebSocketController) handleClientMessage(senderID string, message []byte) {
log.Printf("📥📥 收到客户端消息: \n SenderID=%s \n Size=%d bytes", senderID, len(message))
var payload models.SendMessagePayload
if err := json.Unmarshal(message, &payload); err == nil && payload.RequestType != "" {
log.Printf("📦📦 解析JSON消息成功: Type=%s", payload.RequestType)
if payload.RequestType == "bind" && payload.SenderUserID != "" {
log.Printf("🔗🔗 处理绑定请求: \n ClientID=%s \n UserID=%s", senderID, payload.SenderUserID)
if err := c.bindClientToUser(senderID, payload.SenderUserID); err != nil {
log.Printf("❌❌ 绑定失败: %v", err)
} else {
log.Printf("✅ 绑定成功: \n ClientID=%s \n UserID=%s", senderID, payload.SenderUserID)
}
return
}
if payload.RequestType == "send_message" {
if payload.TargetClientID != "" {
log.Printf("📨📨 处理客户端发起的发送请求: \n TargetID=%s \n MsgType=%d",
payload.TargetClientID, payload.MessageType)
clientMsg := models.ClientReceivedMessage{
SenderID: senderID,
ReceiverID: payload.TargetClientID,
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
c.sendMessageToClient(payload.TargetClientID, clientMsg)
return
}
if payload.ReceiverUserID != "" {
log.Printf("📨📨 处理客户端发起的用户发送请求: \n ReceiverUserID=%s \n MsgType=%d",
payload.ReceiverUserID, payload.MessageType)
clientMsg := models.ClientReceivedMessage{
SenderID: senderID,
ReceiverID: "",
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
c.sendMessageToUser(payload.ReceiverUserID, clientMsg)
return
}
}
}
var clientMsg models.ClientReceivedMessage
if err := json.Unmarshal(message, &clientMsg); err == nil && clientMsg.ReceiverID != "" {
log.Printf("📨📨 处理客户端封装消息: \n TargetID=%s \n MsgType=%d",
clientMsg.ReceiverID, clientMsg.MessageType)
clientMsg.SenderID = senderID
c.sendMessageToClient(clientMsg.ReceiverID, clientMsg)
return
}
log.Printf("⚠️ 无法识别的消息格式: \n Size=%d bytes \n Message=%s", len(message), string(message))
}
// SendMessageHandler API消息发送处理器
func (c *WebSocketController) SendMessageHandler(ctx *gin.Context) {
start := time.Now()
var payload models.SendMessagePayload
if err := ctx.ShouldBindJSON(&payload); err != nil {
log.Printf("⚠️ 无效的JSON请求格式: %v", err)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "无效的JSON格式"})
return
}
if payload.RequestType == "" || (payload.TargetClientID == "" && payload.ReceiverUserID == "") ||
payload.SenderUserID == "" || payload.ReceiverUserID == "" {
log.Printf("⚠️ 缺少必要参数: \n request_type=%s \n target_client_id=%s \n receiver_user_id=%s \n sender_user_id=%s \n receiver_user_id=%s",
payload.RequestType, payload.TargetClientID, payload.ReceiverUserID, payload.SenderUserID, payload.ReceiverUserID)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "缺少必要参数"})
return
}
if payload.RequestType == "send_message" {
log.Printf("📤📤 处理API发送请求: \n SenderUser=%s \n MsgType=%d",
payload.SenderUserID, payload.MessageType)
clientMsg := models.ClientReceivedMessage{
SenderID: "system",
ReceiverID: payload.TargetClientID,
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
var result string
var sendErr error
if payload.TargetClientID != "" {
log.Printf("🎯🎯 目标类型: ClientID | Target=%s", payload.TargetClientID)
sendErr = c.sendMessageToClient(payload.TargetClientID, clientMsg)
result = fmt.Sprintf("消息已发送到ClientID: %s", payload.TargetClientID)
} else if payload.ReceiverUserID != "" {
log.Printf("🎯🎯 目标类型: UserID | ReceiverUser=%s", payload.ReceiverUserID)
sendErr = c.sendMessageToUser(payload.ReceiverUserID, clientMsg)
result = fmt.Sprintf("消息已发送到UserID: %s", payload.ReceiverUserID)
} else {
ctx.JSON(http.StatusBadRequest, gin.H{"error": "必须指定目标客户端ID或用户ID"})
return
}
if sendErr != nil {
ctx.JSON(http.StatusInternalServerError, gin.H{"error": sendErr.Error()})
return
}
ctx.JSON(http.StatusOK, gin.H{"status": "success", "message": result})
log.Printf("✅ API请求完成: Duration=%s", time.Since(start))
return
}
ctx.JSON(http.StatusBadRequest, gin.H{"error": "不支持的request_type"})
}
// 用户绑定处理器
func (c *WebSocketController) BindHandler(ctx *gin.Context) {
start := time.Now()
var req models.BindRequest
if err := ctx.ShouldBindJSON(&req); err != nil {
log.Printf("⚠️ 无效的绑定请求格式: %v", err)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "无效的JSON格式"})
return
}
if req.UserID == "" || req.ClientID == "" {
log.Printf("⚠️ 缺少必要参数: \n user_id=%s \n client_id=%s", req.UserID, req.ClientID)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "user_id和client_id不能为空"})
return
}
log.Printf("🔗🔗 处理用户绑定请求: \n UserID=%s \n ClientID=%s", req.UserID, req.ClientID)
if err := c.bindClientToUser(req.ClientID, req.UserID); err != nil {
log.Printf("❌❌ 绑定失败: %v", err)
ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
ctx.JSON(http.StatusOK, gin.H{
"status": "success",
"message": fmt.Sprintf("绑定成功: ClientID %s -> UserID %s", req.ClientID, req.UserID),
})
log.Printf("✅ 绑定请求完成: \n UserID=%s \n ClientID=%s \n Duration=%s",
req.UserID, req.ClientID, time.Since(start))
}
// 通过用户ID发送消息处理器
func (c *WebSocketController) SendToUserHandler(ctx *gin.Context) {
start := time.Now()
var payload models.SendToUserPayload
if err := ctx.ShouldBindJSON(&payload); err != nil {
log.Printf("⚠️ 无效的JSON请求格式: %v", err)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "无效的JSON格式"})
return
}
if payload.ReceiverUserID == "" || payload.SenderUserID == "" {
log.Printf("⚠️ 缺少必要参数: \n receiver_user_id=%s \n sender_user_id=%s",
payload.ReceiverUserID, payload.SenderUserID)
ctx.JSON(http.StatusBadRequest, gin.H{"error": "缺少接收方或发送方用户ID"})
return
}
log.Printf("📤📤 处理API用户发送请求: \n SenderUser=%s → ReceiverUser=%s \n MsgType=%d",
payload.SenderUserID, payload.ReceiverUserID, payload.MessageType)
clientMsg := models.ClientReceivedMessage{
SenderID: "system",
ReceiverID: "",
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
if err := c.sendMessageToUser(payload.ReceiverUserID, clientMsg); err != nil {
ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
ctx.JSON(http.StatusOK, gin.H{
"code": 0,
"status": "success",
"message": fmt.Sprintf("消息已发送至用户 %s", payload.ReceiverUserID),
})
log.Printf("✅ API请求完成: Duration=%s", time.Since(start))
}
// 绑定客户端与用户
func (c *WebSocketController) bindClientToUser(clientID, userID string) error {
if err := c.RedisCli.HSet(c.RedisCtx, models.ClientUserKey, clientID, userID).Err(); err != nil {
return fmt.Errorf("保存client-user映射失败: %w", err)
}
userClientKey := fmt.Sprintf("%s:%s", models.UserClientKey, userID)
if err := c.RedisCli.SAdd(c.RedisCtx, userClientKey, clientID).Err(); err != nil {
return fmt.Errorf("保存user-client映射失败: %w", err)
}
expiration := 24 * time.Hour
if err := c.RedisCli.Expire(c.RedisCtx, userClientKey, expiration).Err(); err != nil {
log.Printf("⚠️ 设置过期时间失败: %v", err)
}
log.Printf("🔗🔗 用户绑定成功: \n ClientID=%s → UserID=%s", clientID, userID)
return nil
}
// 通过用户ID发送消息
func (c *WebSocketController) sendMessageToUser(userID string, msg models.ClientReceivedMessage) error {
start := time.Now()
log.Printf("👤👤 通过用户ID发送消息: UserID=%s", userID)
userClientKey := fmt.Sprintf("%s:%s", models.UserClientKey, userID)
clientIDs, err := c.RedisCli.SMembers(c.RedisCtx, userClientKey).Result()
if err != nil {
return fmt.Errorf("获取用户clientID失败: %w", err)
}
if len(clientIDs) == 0 {
return fmt.Errorf("用户未绑定任何客户端")
}
log.Printf("📡📡 找到 %d 个关联的客户端: UserID=%s", len(clientIDs), userID)
var successCount, failCount int
var lastError error
for _, clientID := range clientIDs {
targetMsg := msg
targetMsg.ReceiverID = clientID
if err := c.sendMessageToClient(clientID, targetMsg); err != nil {
log.Printf("❌❌ 发送消息到客户端失败: \n ClientID=%s \n %v", clientID, err)
failCount++
lastError = err
} else {
successCount++
}
}
log.Printf("📬📬 消息发送完成: \n 成功 %d \n 失败 %d \n UserID=%s \n Duration=%s",
successCount, failCount, userID, time.Since(start))
if failCount > 0 {
return fmt.Errorf("部分消息发送失败,最后错误: %w", lastError)
}
return nil
}
// 发送消息给客户端
func (c *WebSocketController) sendMessageToClient(clientID string, msg models.ClientReceivedMessage) error {
start := time.Now()
msgJSON, err := json.Marshal(msg)
if err != nil {
log.Printf("❌❌ 消息序列化失败: %v", err)
return fmt.Errorf("内部错误")
}
if conn := c.getClient(clientID); conn != nil {
log.Printf("📤📤 向本地客户端发送消息: \n ClientID=%s \n MsgType=%d \n Size=%d bytes",
clientID, msg.MessageType, len(msgJSON))
if err := conn.WriteMessage(websocket.TextMessage, msgJSON); err != nil {
log.Printf("❌❌ 发送消息失败: %v | ClientID=%s", err, clientID)
c.removeClient(clientID)
return fmt.Errorf("发送消息失败")
}
log.Printf("✅ 消息成功发送到本地客户端: \n ClientID=%s \n Duration=%s",
clientID, time.Since(start))
return nil
}
log.Printf("🌐🌐 向远程节点转发消息: ClientID=%s", clientID)
redisMsg := models.RedisMessage{
SenderNodeID: c.NodeID,
ClientID: clientID,
Message: string(msgJSON),
}
redisMsgJSON, err := json.Marshal(redisMsg)
if err != nil {
log.Printf("❌❌ Redis消息序列化失败: %v", err)
return fmt.Errorf("内部错误")
}
if err := c.RedisCli.Publish(c.RedisCtx, "ws_messages", redisMsgJSON).Err(); err != nil {
log.Printf("❌❌ 发布消息到Redis失败: %v", err)
return fmt.Errorf("无法转发消息")
}
log.Printf("📡📡 消息已转发到Redis: \n ClientID=%s \n Size=%d bytes \n Duration=%s",
clientID, len(redisMsgJSON), time.Since(start))
return nil
}
// SubscribeToRedis 订阅Redis消息
func (c *WebSocketController) SubscribeToRedis() {
pubsub := c.RedisCli.Subscribe(c.RedisCtx, "ws_messages")
defer pubsub.Close()
ch := pubsub.Channel()
log.Println("🔔🔔 开始监听Redis消息...")
for msg := range ch {
start := time.Now()
var redisMsg models.RedisMessage
if err := json.Unmarshal([]byte(msg.Payload), &redisMsg); err != nil {
log.Printf("❌❌ 解析Redis消息失败: %v", err)
continue
}
log.Printf("📥📥 收到Redis消息: \n Sender=%s \n ClientID=%s \n Size=%d bytes",
redisMsg.SenderNodeID, redisMsg.ClientID, len(msg.Payload))
if redisMsg.SenderNodeID == c.NodeID {
log.Println(" 忽略本节点转发的消息")
continue
}
if conn := c.getClient(redisMsg.ClientID); conn != nil {
if err := conn.WriteMessage(websocket.TextMessage, []byte(redisMsg.Message)); err != nil {
log.Printf("❌❌ 处理Redis消息失败: %v | ClientID=%s", err, redisMsg.ClientID)
c.removeClient(redisMsg.ClientID)
continue
}
log.Printf("✅ Redis消息已处理: \n ClientID=%s \n Duration=%s",
redisMsg.ClientID, time.Since(start))
} else {
log.Printf("⚠️ 目标客户端不在本节点: ClientID=%s", redisMsg.ClientID)
}
}
}
// ========================== 辅助方法 ========================== //
// 添加客户端到映射
func (c *WebSocketController) addClient(clientID string, conn *websocket.Conn) {
c.ClientsMux.Lock()
defer c.ClientsMux.Unlock()
c.Clients[clientID] = conn
log.Printf("📌📌 添加客户端到连接池: \n ClientID=%s \n 当前连接数=%d",
clientID, len(c.Clients))
}
// 从映射中移除客户端
func (c *WebSocketController) removeClient(clientID string) {
c.ClientsMux.Lock()
defer c.ClientsMux.Unlock()
if _, exists := c.Clients[clientID]; !exists {
return
}
delete(c.Clients, clientID)
log.Printf("🗑🗑️ 从连接池移除客户端: \n ClientID=%s \n 当前连接数=%d",
clientID, len(c.Clients))
if err := c.RedisCli.HDel(c.RedisCtx, models.ClientUserKey, clientID).Err(); err != nil {
log.Printf("⚠️ 删除client_user_mapping失败: %v", err)
}
}
// 获取客户端连接
func (c *WebSocketController) getClient(clientID string) *websocket.Conn {
c.ClientsMux.RLock()
defer c.ClientsMux.RUnlock()
return c.Clients[clientID]
}