hub.go 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254
  1. package websocket
  2. import (
  3. "log"
  4. "sync"
  5. "github.com/2930134478/AI-CS/backend/models"
  6. )
  7. // OnClientConnectCallback 客户端连接时的回调函数。
  8. // conversationID: 对话ID
  9. // isVisitor: 是否是访客
  10. // visitorCount: 该对话当前的访客连接数
  11. // agentID: 客服ID(如果是客服连接)
  12. type OnClientConnectCallback func(conversationID uint, isVisitor bool, visitorCount int, agentID uint)
  13. // OnClientDisconnectCallback 客户端断开连接时的回调函数。
  14. // conversationID: 对话ID
  15. // isVisitor: 是否是访客
  16. // visitorCount: 该对话当前的访客连接数(断开后)
  17. type OnClientDisconnectCallback func(conversationID uint, isVisitor bool, visitorCount int)
  18. // Hub 管理所有 WebSocket 连接
  19. // 每个对话(conversation)可以有多个人连接(访客和客服)
  20. type Hub struct {
  21. // 每个对话ID对应的客户端连接列表
  22. // conversationID -> []*Client
  23. conversations map[uint]map[*Client]bool
  24. // 注册新客户端(当有人连接时)
  25. register chan *Client
  26. // 注销客户端(当有人断开连接时)
  27. unregister chan *Client
  28. // 广播消息(当有新消息时,推送给所有相关的客户端)
  29. broadcast chan *Message
  30. // 互斥锁(保护并发访问)
  31. mu sync.RWMutex
  32. // 回调函数
  33. onConnect OnClientConnectCallback
  34. onDisconnect OnClientDisconnectCallback
  35. }
  36. // Message 是要广播的消息
  37. type Message struct {
  38. ConversationID uint `json:"conversation_id"`
  39. Data interface{} `json:"data"` // 消息内容(可以是 Message 对象)
  40. Type string `json:"type"` // 消息类型:new_message, conversation_update 等
  41. }
  42. // NewHub 创建一个新的 Hub
  43. func NewHub(onConnect OnClientConnectCallback, onDisconnect OnClientDisconnectCallback) *Hub {
  44. return &Hub{
  45. conversations: make(map[uint]map[*Client]bool),
  46. register: make(chan *Client),
  47. unregister: make(chan *Client),
  48. broadcast: make(chan *Message, 256),
  49. onConnect: onConnect,
  50. onDisconnect: onDisconnect,
  51. }
  52. }
  53. // Run 启动 Hub,处理所有事件
  54. func (h *Hub) Run() {
  55. for {
  56. select {
  57. // 新客户端连接
  58. case client := <-h.register:
  59. h.mu.Lock()
  60. // 如果这个对话还没有客户端,创建一个新的 map
  61. if h.conversations[client.conversationID] == nil {
  62. h.conversations[client.conversationID] = make(map[*Client]bool)
  63. }
  64. // 把这个客户端加入到对话中
  65. h.conversations[client.conversationID][client] = true
  66. // 统计该对话的访客连接数
  67. visitorCount := 0
  68. for c := range h.conversations[client.conversationID] {
  69. if c.isVisitor {
  70. visitorCount++
  71. }
  72. }
  73. h.mu.Unlock()
  74. // 调用连接回调函数
  75. if h.onConnect != nil {
  76. h.onConnect(client.conversationID, client.isVisitor, visitorCount, client.agentID)
  77. }
  78. // 客户端断开连接
  79. case client := <-h.unregister:
  80. h.mu.Lock()
  81. // 从对话中移除这个客户端
  82. wasVisitor := client.isVisitor
  83. if clients, ok := h.conversations[client.conversationID]; ok {
  84. if _, ok := clients[client]; ok {
  85. delete(clients, client)
  86. // 关闭发送通道(避免重复关闭导致 panic)
  87. select {
  88. case _, ok := <-client.send:
  89. if !ok {
  90. // 通道已经关闭,不需要再次关闭
  91. }
  92. default:
  93. // 通道未关闭,关闭它
  94. close(client.send)
  95. }
  96. // 统计该对话的访客连接数(断开后)
  97. visitorCount := 0
  98. for c := range clients {
  99. if c.isVisitor {
  100. visitorCount++
  101. }
  102. }
  103. // 如果这个对话没有客户端了,删除对话
  104. if len(clients) == 0 {
  105. delete(h.conversations, client.conversationID)
  106. }
  107. h.mu.Unlock()
  108. // 调用断开回调函数
  109. if h.onDisconnect != nil {
  110. h.onDisconnect(client.conversationID, wasVisitor, visitorCount)
  111. }
  112. } else {
  113. h.mu.Unlock()
  114. log.Printf("⚠️ 客户端断开时未找到: 对话ID=%d", client.conversationID)
  115. }
  116. } else {
  117. h.mu.Unlock()
  118. log.Printf("⚠️ 客户端断开时对话不存在: 对话ID=%d", client.conversationID)
  119. }
  120. // 广播消息
  121. case message := <-h.broadcast:
  122. h.mu.RLock()
  123. // 找到这个对话的所有客户端
  124. clients, ok := h.conversations[message.ConversationID]
  125. if !ok {
  126. h.mu.RUnlock()
  127. log.Printf("⚠️ 广播消息失败: 对话ID=%d 没有客户端连接", message.ConversationID)
  128. continue
  129. }
  130. // 创建一个客户端列表的副本(避免在遍历时修改)
  131. clientList := make([]*Client, 0, len(clients))
  132. for client := range clients {
  133. clientList = append(clientList, client)
  134. }
  135. h.mu.RUnlock()
  136. // 给所有客户端发送消息
  137. for _, client := range clientList {
  138. select {
  139. case client.send <- message:
  140. default:
  141. // 如果发送失败(客户端可能已经断开),关闭连接
  142. log.Printf("⚠️ 发送消息失败: 对话ID=%d, 客户端断开", client.conversationID)
  143. close(client.send)
  144. h.mu.Lock()
  145. delete(h.conversations[client.conversationID], client)
  146. h.mu.Unlock()
  147. }
  148. }
  149. }
  150. }
  151. }
  152. // BroadcastMessage 广播消息到指定对话的所有客户端
  153. func (h *Hub) BroadcastMessage(conversationID uint, messageType string, data interface{}) {
  154. h.broadcast <- &Message{
  155. ConversationID: conversationID,
  156. Type: messageType,
  157. Data: data,
  158. }
  159. }
  160. // BroadcastToAllAgents 广播消息到所有客服客户端(不管连接到哪个对话)
  161. // 用于 visitor_status_update 等需要所有客服都收到的事件
  162. func (h *Hub) BroadcastToAllAgents(messageType string, data interface{}) {
  163. h.mu.RLock()
  164. // 收集所有客服客户端(isVisitor == false)
  165. allAgents := make([]*Client, 0)
  166. for _, clients := range h.conversations {
  167. for client := range clients {
  168. if !client.isVisitor {
  169. allAgents = append(allAgents, client)
  170. }
  171. }
  172. }
  173. h.mu.RUnlock()
  174. // 为每个客服客户端创建消息并发送
  175. for _, client := range allAgents {
  176. // 如果 data 是 Message 对象,使用消息的 conversation_id
  177. // 否则使用客户端连接的对话ID
  178. var conversationID uint
  179. if msg, ok := data.(*models.Message); ok {
  180. conversationID = msg.ConversationID
  181. } else if convID, ok := data.(map[string]interface{})["conversation_id"]; ok {
  182. if id, ok := convID.(uint); ok {
  183. conversationID = id
  184. } else if id, ok := convID.(float64); ok {
  185. conversationID = uint(id)
  186. } else {
  187. conversationID = client.conversationID
  188. }
  189. } else {
  190. conversationID = client.conversationID
  191. }
  192. message := &Message{
  193. ConversationID: conversationID,
  194. Type: messageType,
  195. Data: data,
  196. }
  197. select {
  198. case client.send <- message:
  199. default:
  200. // 如果发送失败(客户端可能已经断开),关闭连接
  201. log.Printf("⚠️ 发送消息到客服失败: 对话ID=%d, 客户端断开", client.conversationID)
  202. close(client.send)
  203. h.mu.Lock()
  204. if clients, ok := h.conversations[client.conversationID]; ok {
  205. delete(clients, client)
  206. if len(clients) == 0 {
  207. delete(h.conversations, client.conversationID)
  208. }
  209. }
  210. h.mu.Unlock()
  211. }
  212. }
  213. }
  214. // GetOnlineAgentIDs 获取所有在线客服的用户ID列表(去重)
  215. // 返回一个 map,key 是 agentID,value 是 true(用于快速查找)
  216. func (h *Hub) GetOnlineAgentIDs() map[uint]bool {
  217. h.mu.RLock()
  218. defer h.mu.RUnlock()
  219. agentIDs := make(map[uint]bool)
  220. for _, clients := range h.conversations {
  221. for client := range clients {
  222. if !client.isVisitor && client.agentID > 0 {
  223. agentIDs[client.agentID] = true
  224. }
  225. }
  226. }
  227. return agentIDs
  228. }