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] }