Files
nl-im-websocket-demo/main.txt
2025-07-03 08:50:38 +08:00

678 lines
22 KiB
Plaintext
Raw Permalink 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 main
import (
"context"
"encoding/json"
"fmt"
"log"
"net/http"
"os"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/go-redis/redis/v8"
"github.com/google/uuid"
"github.com/gorilla/websocket"
)
// WebSocket 升级器配置
var upgrader = websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: func(r *http.Request) bool {
return true // 允许跨域
},
}
// 统一消息结构(发送方使用)
type SendMessagePayload struct {
RequestType string `json:"request_type"` // 消息类型
TargetClientID string `json:"target_client_id"` // 目标客户端ID
SenderUserID string `json:"sender_user_id"` // 发送者用户ID
ReceiverUserID string `json:"receiver_user_id"` // 接收者用户ID
MessageType int `json:"message_type"` // 消息类型0-6
MessageContent string `json:"message_content"` // 消息内容
}
// 通过用户ID发送消息的结构
type SendToUserPayload struct {
SenderUserID string `json:"sender_user_id"` // 发送者用户ID
ReceiverUserID string `json:"receiver_user_id"` // 接收者用户ID
MessageType int `json:"message_type"` // 消息类型0-6
MessageContent string `json:"message_content"` // 消息内容
}
// 客户端接收消息结构(最终格式)
type ClientReceivedMessage struct {
SenderID string `json:"sender_id"` // 发送者连接ID
ReceiverID string `json:"receiver_id"` // 接收者连接ID
SenderUserID string `json:"sender_user_id"` // 发送者用户ID
ReceiverUserID string `json:"receiver_user_id"` // 接收者用户ID
MessageType int `json:"message_type"` // 消息类型0-6
Content string `json:"content"` // 消息内容
SendTime string `json:"send_time"` // 发送时间
}
// Redis 消息结构
type RedisMessage struct {
SenderNodeID string `json:"sender_node_id"` // 发送节点ID
ClientID string `json:"client_id"` // 目标客户端ID
Message string `json:"message"` // 消息内容
}
// 用户绑定请求结构
type BindRequest struct {
UserID string `json:"user_id"` // 用户ID
ClientID string `json:"client_id"` // 客户端ID
}
// 全局变量
var (
clients = make(map[string]*websocket.Conn) // 客户端连接映射
clientsMux sync.RWMutex // 连接映射的读写锁
redisCtx = context.Background() // Redis 上下文
nodeID string // 当前节点ID
redisCli *redis.Client // Redis 客户端
port string // 服务端口
)
// Redis键常量
const (
ClientUserKey = "client_user_mapping" // clientid -> userid映射
UserClientKey = "user_client_mapping" // userid -> clientid列表
)
func main() {
// 配置系统
configureSystem()
// 初始化Redis客户端
initRedisClient()
// 创建Gin路由器
router := gin.Default()
router.Use(gin.Recovery()) // 添加恢复中间件
// 注册路由
router.GET("/ws", handleWebSocket) // WebSocket连接端点
router.POST("/send", sendMessageHandler) // 消息发送端点通过客户端ID
router.POST("/bind", bindHandler) // 用户绑定端点
router.POST("/send-to-user", sendToUserHandler) // 新增通过用户ID发送消息
router.GET("/health", healthHandler) // 健康检查端点
// 输出启动信息
printStartupInfo()
// 启动Redis消息订阅
go subscribeToRedis()
// 启动HTTP服务
runServer(router)
}
// 配置系统环境
func configureSystem() {
nodeID = getEnv("NODE_ID", "local") // 获取节点ID默认"local"
port = getEnv("PORT", "12080") // 获取端口默认12080
log.SetPrefix(fmt.Sprintf("[Node:%s] ", nodeID)) // 设置日志前缀
log.SetFlags(log.LstdFlags | log.Lmicroseconds) // 设置日志格式
}
// 初始化Redis客户端
func initRedisClient() {
redisAddr := getEnv("REDIS_ADDR", "localhost:6379")
redisPassword := getEnv("REDIS_PASSWORD", "") // Redis密码
redisCli = redis.NewClient(&redis.Options{
Addr: redisAddr, // Redis地址
Password: redisPassword, // Redis密码
DB: 0, // 默认DB
})
// 测试Redis连接
if err := checkRedisConnection(); err != nil {
log.Fatalf("❌ Redis连接失败: %v", err)
}
}
// 检查Redis连接
func checkRedisConnection() error {
_, err := redisCli.Ping(redisCtx).Result()
return err
}
// 输出启动信息
func printStartupInfo() {
hostname, _ := os.Hostname()
log.Printf("🚀 WebSocket服务启动: NodeID=%s", nodeID)
log.Printf("🌐 监听端口: %s", port)
log.Printf("📡 Redis地址: %s", getEnv("REDIS_ADDR", "localhost:6379"))
log.Printf("💻 主机: %s", hostname)
log.Printf("🕒 启动时间: %s", time.Now().Format("2006-01-02 15:04:05"))
log.Printf("🔔 支持消息类型: 0:文本 1:图片 2:音频 3:视频 4:处方 5:病例 6:视频通话")
log.Printf("🔑 用户绑定键: %s, %s", ClientUserKey, UserClientKey)
log.Printf("📬 新增接口: POST /send-to-user (直接通过用户ID发送消息)")
log.Println("🔗 等待客户端连接...")
}
// 启动服务
func runServer(router *gin.Engine) {
addr := "127.0.0.1:" + port
log.Printf("🔌 服务地址: http://%s", addr)
log.Printf("🔌 WebSocket连接地址: ws://%s/ws", addr)
if err := router.Run(addr); err != nil {
log.Fatalf("❌ 服务启动失败: %v", err)
}
}
// 健康检查处理
func healthHandler(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"node": nodeID,
"time": time.Now().Format(time.RFC3339),
})
}
// WebSocket处理函数
func handleWebSocket(c *gin.Context) {
start := time.Now()
clientIP := c.ClientIP()
log.Printf("👤 客户端连接中: IP=%s", clientIP)
// 升级为WebSocket连接
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
log.Printf("⚠️ WebSocket升级失败: %v | ClientIP=%s", err, clientIP)
c.JSON(http.StatusBadRequest, gin.H{"error": "无法升级为WebSocket连接"})
return
}
defer conn.Close()
// 生成客户端ID
clientID := generateClientID()
log.Printf("✅ 客户端已连接: ClientID=%s | IP=%s | Duration=%s",
clientID, clientIP, time.Since(start))
// 将连接添加到客户端映射
addClient(clientID, conn)
defer removeClient(clientID)
// 发送客户端ID
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 handleClientMessage(clientID, message)
}
}
}
// 处理客户端消息
func handleClientMessage(senderID string, message []byte) {
log.Printf("📥 收到客户端消息: SenderID=%s | Size=%d bytes", senderID, len(message))
// 尝试解析为发送消息请求
var payload 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("🔗 处理绑定请求: ClientID=%s -> UserID=%s", senderID, payload.SenderUserID)
if err := bindClientToUser(senderID, payload.SenderUserID); err != nil {
log.Printf("❌ 绑定失败: %v", err)
} else {
log.Printf("✅ 绑定成功: ClientID=%s -> UserID=%s", senderID, payload.SenderUserID)
}
return
}
// 处理发送消息请求
if payload.RequestType == "send_message" {
// 如果指定了目标clientid则直接发送
if payload.TargetClientID != "" {
log.Printf("📨 处理客户端发起的发送请求: TargetID=%s | MsgType=%d",
payload.TargetClientID, payload.MessageType)
// 构建客户端消息结构
clientMsg := ClientReceivedMessage{
SenderID: senderID,
ReceiverID: payload.TargetClientID,
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
// 发送消息
sendMessageToClient(payload.TargetClientID, clientMsg)
return
}
// 如果指定了接收用户ID则通过用户ID发送
if payload.ReceiverUserID != "" {
log.Printf("📨 处理客户端发起的用户发送请求: ReceiverUserID=%s | MsgType=%d",
payload.ReceiverUserID, payload.MessageType)
// 构建客户端消息结构
clientMsg := ClientReceivedMessage{
SenderID: senderID,
ReceiverID: "", // 通过userID发送留空
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
// 通过用户ID发送消息
sendMessageToUser(payload.ReceiverUserID, clientMsg)
return
}
}
}
// 尝试解析为已封装的消息结构
var clientMsg ClientReceivedMessage
if err := json.Unmarshal(message, &clientMsg); err == nil && clientMsg.ReceiverID != "" {
log.Printf("📨 处理客户端封装消息: TargetID=%s | MsgType=%d",
clientMsg.ReceiverID, clientMsg.MessageType)
// 设置发送者ID避免客户端冒充
clientMsg.SenderID = senderID
// 发送消息
sendMessageToClient(clientMsg.ReceiverID, clientMsg)
return
}
// 无法识别的消息格式
log.Printf("⚠️ 无法识别的消息格式: Size=%d bytes | Message=%s", len(message), string(message))
}
// 消息发送处理HTTP API- 通过客户端ID发送
func sendMessageHandler(c *gin.Context) {
start := time.Now()
var payload SendMessagePayload
if err := c.ShouldBindJSON(&payload); err != nil {
log.Printf("⚠️ 无效的JSON请求格式: %v", err)
c.JSON(http.StatusBadRequest, gin.H{"error": "无效的JSON格式"})
return
}
if payload.RequestType == "" || (payload.TargetClientID == "" && payload.ReceiverUserID == "") ||
payload.SenderUserID == "" || payload.ReceiverUserID == "" {
log.Printf("⚠️ 缺少必要参数: request_type=%s target_client_id=%s receiver_user_id=%s sender_user_id=%s receiver_user_id=%s",
payload.RequestType, payload.TargetClientID, payload.ReceiverUserID, payload.SenderUserID, payload.ReceiverUserID)
c.JSON(http.StatusBadRequest, gin.H{"error": "缺少必要参数"})
return
}
// 只处理发送消息请求
if payload.RequestType == "send_message" {
log.Printf("📤 处理API发送请求: SenderUser=%s | MsgType=%d",
payload.SenderUserID, payload.MessageType)
// 构建客户端消息结构
clientMsg := ClientReceivedMessage{
SenderID: "system", // API消息标记为系统发送
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 = sendMessageToClient(payload.TargetClientID, clientMsg)
result = fmt.Sprintf("消息已发送到ClientID: %s", payload.TargetClientID)
} else if payload.ReceiverUserID != "" {
log.Printf("🎯 目标类型: UserID | ReceiverUser=%s", payload.ReceiverUserID)
sendErr = sendMessageToUser(payload.ReceiverUserID, clientMsg)
result = fmt.Sprintf("消息已发送到UserID: %s", payload.ReceiverUserID)
} else {
c.JSON(http.StatusBadRequest, gin.H{"error": "必须指定目标客户端ID或用户ID"})
return
}
if sendErr != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": sendErr.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"status": "success", "message": result})
log.Printf("✅ API请求完成: Duration=%s", time.Since(start))
return
}
c.JSON(http.StatusBadRequest, gin.H{"error": "不支持的request_type"})
}
// 用户绑定处理
func bindHandler(c *gin.Context) {
start := time.Now()
var req BindRequest
if err := c.ShouldBindJSON(&req); err != nil {
log.Printf("⚠️ 无效的绑定请求格式: %v", err)
c.JSON(http.StatusBadRequest, gin.H{"error": "无效的JSON格式"})
return
}
if req.UserID == "" || req.ClientID == "" {
log.Printf("⚠️ 缺少必要参数: user_id=%s client_id=%s", req.UserID, req.ClientID)
c.JSON(http.StatusBadRequest, gin.H{"error": "user_id和client_id不能为空"})
return
}
log.Printf("🔗 处理用户绑定请求: UserID=%s ClientID=%s", req.UserID, req.ClientID)
// 执行绑定
if err := bindClientToUser(req.ClientID, req.UserID); err != nil {
log.Printf("❌ 绑定失败: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{
"status": "success",
"message": fmt.Sprintf("绑定成功: ClientID %s -> UserID %s", req.ClientID, req.UserID),
})
log.Printf("✅ 绑定请求完成: UserID=%s ClientID=%s | Duration=%s",
req.UserID, req.ClientID, time.Since(start))
}
// 通过用户ID发送消息处理新增独立接口
func sendToUserHandler(c *gin.Context) {
start := time.Now()
var payload SendToUserPayload
if err := c.ShouldBindJSON(&payload); err != nil {
log.Printf("⚠️ 无效的JSON请求格式: %v", err)
c.JSON(http.StatusBadRequest, gin.H{"error": "无效的JSON格式"})
return
}
if payload.ReceiverUserID == "" || payload.SenderUserID == "" {
log.Printf("⚠️ 缺少必要参数: receiver_user_id=%s sender_user_id=%s",
payload.ReceiverUserID, payload.SenderUserID)
c.JSON(http.StatusBadRequest, gin.H{"error": "缺少接收方或发送方用户ID"})
return
}
log.Printf("📤 处理API用户发送请求: SenderUser=%s → ReceiverUser=%s | MsgType=%d",
payload.SenderUserID, payload.ReceiverUserID, payload.MessageType)
// 构建客户端消息结构
clientMsg := ClientReceivedMessage{
SenderID: "system", // 标记为系统发送
ReceiverID: "", // 留空通过用户ID发送
SenderUserID: payload.SenderUserID,
ReceiverUserID: payload.ReceiverUserID,
MessageType: payload.MessageType,
Content: payload.MessageContent,
SendTime: time.Now().Format(time.RFC3339),
}
if err := sendMessageToUser(payload.ReceiverUserID, clientMsg); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{
"code": 0,
"status": "success",
"message": fmt.Sprintf("消息已发送至用户 %s", payload.ReceiverUserID),
})
log.Printf("✅ API请求完成: Duration=%s", time.Since(start))
}
// 绑定客户端与用户
func bindClientToUser(clientID, userID string) error {
// 存储 clientid -> userid 映射
if err := redisCli.HSet(redisCtx, ClientUserKey, clientID, userID).Err(); err != nil {
return fmt.Errorf("保存client-user映射失败: %w", err)
}
// 存储 userid -> clientid 映射使用集合存储用户所有在线的clientID
userClientKey := fmt.Sprintf("%s:%s", UserClientKey, userID)
if err := redisCli.SAdd(redisCtx, userClientKey, clientID).Err(); err != nil {
return fmt.Errorf("保存user-client映射失败: %w", err)
}
// 设置过期时间(防止无效数据堆积)
expiration := 24 * time.Hour
if err := redisCli.Expire(redisCtx, userClientKey, expiration).Err(); err != nil {
log.Printf("⚠️ 设置过期时间失败: %v", err)
}
log.Printf("🔗 用户绑定成功: ClientID=%s → UserID=%s", clientID, userID)
return nil
}
// 通过用户ID发送消息
func sendMessageToUser(userID string, msg ClientReceivedMessage) error {
start := time.Now()
log.Printf("👤 通过用户ID发送消息: UserID=%s", userID)
// 获取用户关联的所有clientID
userClientKey := fmt.Sprintf("%s:%s", UserClientKey, userID)
clientIDs, err := redisCli.SMembers(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 {
// 创建消息副本并设置目标clientID
targetMsg := msg
targetMsg.ReceiverID = clientID
if err := sendMessageToClient(clientID, targetMsg); err != nil {
log.Printf("❌ 发送消息到客户端失败: ClientID=%s | %v", clientID, err)
failCount++
lastError = err
} else {
successCount++
}
}
log.Printf("📬 消息发送完成: 成功 %d, 失败 %d | UserID=%s | Duration=%s",
successCount, failCount, userID, time.Since(start))
if failCount > 0 {
return fmt.Errorf("部分消息发送失败,最后错误: %w", lastError)
}
return nil
}
// 发送消息给客户端(核心函数)
func sendMessageToClient(clientID string, msg ClientReceivedMessage) error {
start := time.Now()
// 序列化客户端消息
msgJSON, err := json.Marshal(msg)
if err != nil {
log.Printf("❌ 消息序列化失败: %v", err)
return fmt.Errorf("内部错误")
}
// 检查目标客户端是否在当前节点
if conn := getClient(clientID); conn != nil {
log.Printf("📤 向本地客户端发送消息: ClientID=%s | MsgType=%d | Size=%d bytes",
clientID, msg.MessageType, len(msgJSON))
if err := conn.WriteMessage(websocket.TextMessage, msgJSON); err != nil {
log.Printf("❌ 发送消息失败: %v | ClientID=%s", err, clientID)
removeClient(clientID) // 移除失效连接
return fmt.Errorf("发送消息失败")
}
log.Printf("✅ 消息成功发送到本地客户端: ClientID=%s | Duration=%s",
clientID, time.Since(start))
return nil
}
// 如果目标客户端不在当前节点则通过Redis转发
log.Printf("🌐 向远程节点转发消息: ClientID=%s", clientID)
redisMsg := RedisMessage{
SenderNodeID: nodeID,
ClientID: clientID,
Message: string(msgJSON), // 原始JSON字符串
}
redisMsgJSON, err := json.Marshal(redisMsg)
if err != nil {
log.Printf("❌ Redis消息序列化失败: %v", err)
return fmt.Errorf("内部错误")
}
// 发布到Redis的ws_messages频道
if err := redisCli.Publish(redisCtx, "ws_messages", redisMsgJSON).Err(); err != nil {
log.Printf("❌ 发布消息到Redis失败: %v", err)
return fmt.Errorf("无法转发消息")
}
log.Printf("📡 消息已转发到Redis: ClientID=%s | Size=%d bytes | Duration=%s",
clientID, len(redisMsgJSON), time.Since(start))
return nil
}
// 订阅Redis消息
func subscribeToRedis() {
pubsub := redisCli.Subscribe(redisCtx, "ws_messages")
defer pubsub.Close()
ch := pubsub.Channel()
log.Println("🔔 开始监听Redis消息...")
for msg := range ch {
start := time.Now()
var redisMsg RedisMessage
// 解析Redis消息
if err := json.Unmarshal([]byte(msg.Payload), &redisMsg); err != nil {
log.Printf("❌ 解析Redis消息失败: %v", err)
continue
}
log.Printf("📥 收到Redis消息: Sender=%s | ClientID=%s | Size=%d bytes",
redisMsg.SenderNodeID, redisMsg.ClientID, len(msg.Payload))
// 如果是本地节点发送的消息,忽略
if redisMsg.SenderNodeID == nodeID {
log.Println(" 忽略本节点转发的消息")
continue
}
// 发送给本地客户端
if conn := getClient(redisMsg.ClientID); conn != nil {
// 直接使用Redis中的原始消息字符串
if err := conn.WriteMessage(websocket.TextMessage, []byte(redisMsg.Message)); err != nil {
log.Printf("❌ 处理Redis消息失败: %v | ClientID=%s", err, redisMsg.ClientID)
removeClient(redisMsg.ClientID)
continue
}
log.Printf("✅ Redis消息已处理: ClientID=%s | Duration=%s",
redisMsg.ClientID, time.Since(start))
} else {
log.Printf("⚠️ 目标客户端不在本节点: ClientID=%s", redisMsg.ClientID)
}
}
}
// 生成客户端ID (格式: <节点ID>-<UUID>)
func generateClientID() string {
return fmt.Sprintf("%s-%s", nodeID, uuid.New().String())
}
// 添加客户端到映射
func addClient(clientID string, conn *websocket.Conn) {
clientsMux.Lock()
defer clientsMux.Unlock()
clients[clientID] = conn
log.Printf("📌 添加客户端到连接池: ClientID=%s | 当前连接数=%d",
clientID, len(clients))
}
// 从映射中移除客户端
func removeClient(clientID string) {
clientsMux.Lock()
defer clientsMux.Unlock()
if _, exists := clients[clientID]; !exists {
return
}
delete(clients, clientID)
log.Printf("🗑️ 从连接池移除客户端: ClientID=%s | 当前连接数=%d",
clientID, len(clients))
// 从Redis中移除clientid的绑定信息
// 注意这里只移除clientid->userid映射userid->clientid映射在过期时会自动清理
if err := redisCli.HDel(redisCtx, ClientUserKey, clientID).Err(); err != nil {
log.Printf("⚠️ 删除client_user_mapping失败: %v", err)
}
}
// 获取客户端连接
func getClient(clientID string) *websocket.Conn {
clientsMux.RLock()
defer clientsMux.RUnlock()
return clients[clientID]
}
// 获取环境变量值,如果不存在则使用默认值
func getEnv(key, defaultValue string) string {
if value := os.Getenv(key); value != "" {
return value
}
return defaultValue
}