room.go 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392
  1. // 房间管理模块 -- 负责创建、管理聊天房间,处理客户端的加入、离开和消息转发
  2. package wsClient
  3. import (
  4. "encoding/base64"
  5. "encoding/json"
  6. "fmt"
  7. "log"
  8. "math/rand"
  9. "net/http"
  10. "net/url"
  11. "strings"
  12. "sync"
  13. "github.com/gorilla/websocket"
  14. )
  15. // 全局房间 map,存储所有创建的房间
  16. // key: 房间名称,value: 房间对象指针
  17. var rooms = make(map[string]*Room)
  18. // Room 表示一个聊天房间
  19. type Room struct {
  20. // holds all current clientConns in the room
  21. // 保存房间中所有当前活跃的客户端连接
  22. clientConns map[*Client]bool
  23. // join is a channel for all clients wishing to join the room
  24. // 用于接收希望加入房间的客户端。新客户端通过这个 channel 请求加入
  25. join chan *Client
  26. // leave is a channel for all clients wishing to leave the room
  27. // 用于接收希望离开房间的客户端。断开连接的客户端通过这个 channel 通知离开
  28. leave chan *Client
  29. // forward is a channel that holds incoming messages that should be forwarded to the other clients.
  30. // 用于接收需要转发给其他客户端的消息
  31. // 所有向此房间发送的消息都会通过这个 channel 进行分发
  32. forward chan *message
  33. // name 房间名称,用作Redis存储的键
  34. name string
  35. }
  36. // message 包含要转发的消息内容和发送者信息
  37. type message struct {
  38. content []byte // 消息内容
  39. sender *Client // 发送者 -- // 用处:1.用于在转发时排除发送者自己
  40. }
  41. // RoomParams 表示从客户端 WebSocket子协议中解析出的房间参数
  42. type RoomParams struct {
  43. RoomName string `json:"roomName,omitempty"`
  44. UserId int32 `json:"u_id"` // 用户ID: 只是表示买家客户端是 userId
  45. CustomId int32 `json:"c_id"`
  46. ShopId int32 `json:"s_id"`
  47. Username string `json:"u_name"`
  48. Platform string `json:"platform"`
  49. Type string `json:"type,omitempty"` // 客户端类型:shop-门店端, customer-客户端
  50. StaffId int32 `json:"staff_id"` // 只是表示门店客户端的员工Id
  51. }
  52. // 全局互斥锁,用于保护 rooms 字典的并发访问,确保多个 goroutine 同时访问 rooms 时的线程安全
  53. var mu sync.Mutex
  54. // newRoom 创建并返回一个新的房间实例
  55. // 初始化所有必要的 channel 和 map
  56. func newRoom(name string) *Room {
  57. return &Room{
  58. clientConns: make(map[*Client]bool),
  59. join: make(chan *Client),
  60. leave: make(chan *Client),
  61. forward: make(chan *message), // 现在传输 message 结构体
  62. name: name, // 设置房间名称
  63. }
  64. }
  65. // GetRoomAndRun 获取指定名称的房间,如果不存在则创建新房间并启动
  66. // 这是线程安全的函数,使用互斥锁保护全局房间字典
  67. func GetRoomAndRun(name string) *Room {
  68. mu.Lock()
  69. defer mu.Unlock()
  70. if r, ok := rooms[name]; ok {
  71. return r
  72. }
  73. r := newRoom(name) // 传入房间名称
  74. rooms[name] = r
  75. // 启动房间的消息处理循环(在新的 goroutine 中)
  76. go r.run()
  77. return r
  78. }
  79. // run 是房间的核心消息处理循环
  80. // 在独立的 goroutine 中运行,使用 select 语句处理三种类型的事件
  81. func (r *Room) run() {
  82. for {
  83. select {
  84. // 处理客户端加入事件
  85. case client := <-r.join:
  86. // 检查channel是否已关闭(收到nil值)
  87. if client == nil {
  88. return // channel已关闭,退出goroutine
  89. }
  90. // 将新客户端添加到房间的客户端集合中
  91. r.clientConns[client] = true // 使用 true 作为值,实际上只使用 key(client 指针)
  92. log.Printf("客户端 %v 加入房间 %v", client.name, r.name)
  93. // 处理客户端离开事件
  94. case client := <-r.leave:
  95. // 检查channel是否已关闭(收到nil值)
  96. if client == nil {
  97. return // channel已关闭,退出goroutine
  98. }
  99. // 从房间的客户端集合中删除客户端
  100. delete(r.clientConns, client)
  101. // 关闭客户端的接收 channel,通知客户端的 write() 方法退出
  102. close(client.receive)
  103. log.Printf("客户端 %v 离开房间 %v", client.name, r.name)
  104. if len(r.clientConns) == 0 {
  105. log.Printf("房间 %v 没有客户端,关闭房间", r.name)
  106. close(r.forward)
  107. close(r.join)
  108. close(r.leave)
  109. delete(rooms, r.name)
  110. return // 退出run() goroutine,避免继续从已关闭的channel读取数据
  111. }
  112. // 处理消息转发事件
  113. case msg := <-r.forward:
  114. // 检查channel是否已关闭(收到nil值)
  115. if msg == nil {
  116. return // channel已关闭,退出goroutine
  117. }
  118. // 打印消息内容与发送者信息
  119. log.Printf("Received message from %v: --- %s", msg.sender, string(msg.content))
  120. // 直接存储原始消息内容到Redis,使用房间名作为键
  121. go StoreRawMessage(r.name, msg.content) // 异步存储,不阻塞消息转发
  122. // 创建最近联系人列表 或 更新最近联系人列表
  123. go UpdateRecentContacts(msg.sender)
  124. online := false
  125. // 将消息发送给房间中的所有客户端,但排除发送者自己
  126. for client := range r.clientConns {
  127. if client != msg.sender { // 高效的指针比较,排除发送者 // && client.userInfo.Type != msg.sender.userInfo.Type -- 排除同一类型的客户端
  128. // 将消息发送到每个客户端的接收 channel
  129. client.receive <- msg.content // 只发送消息内容,不包含发送者信息
  130. online = true
  131. }
  132. }
  133. // 本消息转发,目的是让客户与门店互相发消息,当有另一方不在线时,要累计未读消息数
  134. if !online {
  135. go IncrementUnreadCount(msg.sender, 1)
  136. }
  137. }
  138. }
  139. }
  140. // 查询当前房间的客户端数量
  141. func (r *Room)clientCount() int {
  142. return len(r.clientConns)
  143. }
  144. // 统计所有房间里的客户端总数量
  145. func GetTotalClientCount() int {
  146. totalCount := 0
  147. for _, room := range rooms {
  148. totalCount += len(room.clientConns)
  149. }
  150. return totalCount
  151. }
  152. // 用 Client.name 来查找房间
  153. func GetRoomByClientName(clientName string) *Room {
  154. for _, room := range rooms {
  155. for client := range room.clientConns {
  156. if client.name == clientName {
  157. return room
  158. }
  159. }
  160. }
  161. return nil
  162. }
  163. // 用 Client.name 查询用户是否在线
  164. func IsClientOnline(clientName string) bool {
  165. room := GetRoomByClientName(clientName)
  166. return room != nil
  167. }
  168. // ======================================== webSocket服务相关 ========================================
  169. // WebSocket 相关常量定义
  170. const (
  171. // WebSocket 连接的读写缓冲区大小(字节)
  172. socketBufferSize = 1024
  173. // 每个客户端消息接收 channel 的缓冲区大小
  174. messageBufferSize = 256
  175. )
  176. // WebSocket 升级器,用于将 HTTP 连接升级为 WebSocket 连接
  177. // 设置读写缓冲区大小以优化性能
  178. var upgrader = &websocket.Upgrader{
  179. ReadBufferSize: socketBufferSize, // 读缓冲区大小
  180. WriteBufferSize: socketBufferSize, // 写缓冲区大小
  181. CheckOrigin: func(r *http.Request) bool {
  182. return true // 允许跨域连接
  183. },
  184. }
  185. // parseWebSocketProtocols 解析WebSocket子协议中的参数
  186. // 客户端通过子协议传递多个Base64编码的JSON参数片段
  187. func ParseWebSocketProtocols(protocols []string) (*RoomParams, error) {
  188. //log.Println("-------- ParseWebSocketProtocols ---------")
  189. // 查找包含房间参数的子协议
  190. for _, protocol := range protocols {
  191. //log.Printf("Parsing protocol: %s", protocol)
  192. if strings.HasPrefix(protocol, "chat,") {
  193. // 初始化参数结构体
  194. params := RoomParams{}
  195. // 分割协议字符串,获取各个参数部分
  196. parts := strings.Split(protocol, ",")
  197. if len(parts) < 2 {
  198. log.Printf("协议格式错误,参数不足: %s", protocol)
  199. continue
  200. }
  201. // 遍历所有参数部分(跳过第一个"chat")
  202. for i := 1; i < len(parts); i++ {
  203. part := strings.TrimSpace(parts[i])
  204. if err := parseProtocolPart(part, &params); err != nil {
  205. log.Printf("解析协议部分失败 %s: %v", part, err)
  206. continue
  207. }
  208. }
  209. log.Printf("成功解析WebSocket协议参数: %+v\n", params)
  210. return &params, nil
  211. } else {
  212. log.Printf("跳过非chat协议: %s", protocol)
  213. }
  214. }
  215. return nil, fmt.Errorf("未找到有效的房间参数协议")
  216. }
  217. // parseProtocolPart 解析单个协议参数部分
  218. func parseProtocolPart(part string, params *RoomParams) error {
  219. // 检查参数格式:p0-..., p1-..., p2-...
  220. if len(part) < 3 || !strings.Contains(part, "-") {
  221. return fmt.Errorf("参数格式错误: %s", part)
  222. }
  223. // 分割前缀和base64内容
  224. sepIndex := strings.Index(part, "-")
  225. if sepIndex == -1 {
  226. return fmt.Errorf("缺少分隔符: %s", part)
  227. }
  228. prefix := part[:sepIndex]
  229. encodedData := part[sepIndex+1:]
  230. // Base64解码 - 按优先级尝试不同的编码方式
  231. decodedBytes, err := base64.RawStdEncoding.DecodeString(encodedData)
  232. if err != nil {
  233. // 如果RawStdEncoding失败,尝试RawURLEncoding(URL安全的base64)
  234. decodedBytes, err = base64.RawURLEncoding.DecodeString(encodedData)
  235. if err != nil {
  236. // 如果RawURLEncoding也失败,尝试标准的StdEncoding(带填充)
  237. decodedBytes, err = base64.StdEncoding.DecodeString(encodedData)
  238. if err != nil {
  239. // 最后尝试URLEncoding
  240. decodedBytes, err = base64.URLEncoding.DecodeString(encodedData)
  241. if err != nil {
  242. return fmt.Errorf("Base64解码失败: %v", err)
  243. }
  244. }
  245. }
  246. }
  247. // URL解码
  248. decodedStr, err := url.QueryUnescape(string(decodedBytes))
  249. if err != nil {
  250. return fmt.Errorf("URL解码失败: %v", err)
  251. }
  252. // 根据前缀类型解析不同的参数
  253. switch prefix {
  254. case "p0":
  255. // 解析用户ID和自定义ID
  256. var p0Data struct {
  257. UserId int32 `json:"u_id"`
  258. CustomId int32 `json:"c_id"`
  259. }
  260. if err := json.Unmarshal([]byte(decodedStr), &p0Data); err != nil {
  261. return fmt.Errorf("p0参数JSON解析失败: %v", err)
  262. }
  263. params.UserId = p0Data.UserId
  264. params.CustomId = p0Data.CustomId
  265. case "p1":
  266. // 解析店铺ID和用户名
  267. var p1Data struct {
  268. ShopId int32 `json:"s_id"`
  269. Username string `json:"u_name"`
  270. }
  271. if err := json.Unmarshal([]byte(decodedStr), &p1Data); err != nil {
  272. return fmt.Errorf("p1参数JSON解析失败: %v", err)
  273. }
  274. params.ShopId = p1Data.ShopId
  275. params.Username = p1Data.Username
  276. case "p2":
  277. // 解析平台和类型
  278. var p2Data struct {
  279. Platform string `json:"platform"`
  280. Type string `json:"type"`
  281. }
  282. if err := json.Unmarshal([]byte(decodedStr), &p2Data); err != nil {
  283. return fmt.Errorf("p2参数JSON解析失败: %v", err)
  284. }
  285. params.Platform = p2Data.Platform
  286. params.Type = p2Data.Type
  287. case "p3":
  288. // 解析平台和类型
  289. var p3Data struct {
  290. StaffId int32 `json:"staff_id"`
  291. }
  292. if err := json.Unmarshal([]byte(decodedStr), &p3Data); err != nil {
  293. return fmt.Errorf("p2参数JSON解析失败: %v", err)
  294. }
  295. params.StaffId = p3Data.StaffId
  296. // 当有 p4 参数时,解析 p4 参数
  297. default:
  298. log.Printf("未知的参数前缀: %s", prefix)
  299. }
  300. return nil
  301. }
  302. func (r *Room) HttpServe(w http.ResponseWriter, req *http.Request,roomParams *RoomParams) {
  303. // 将 HTTP 连接升级为 WebSocket 连接,指定接受的子协议
  304. responseHeader := http.Header{}
  305. // 这行代码的作用是在 WebSocket握手响应中添加 Sec-WebSocket-Protocol头,值为"chat"。这表示服务器接受并选择了"chat"子协议作为通信协议。当客户端在请求中提供多个子协议选项时,服务器需要在响应中指定它选择使用哪一个。这行代码确认服务器将使用基础的"chat"协议与客户端通信,这样客户端就知道应该使用哪种协议格式来解释后续的WebSocket消息。
  306. responseHeader.Add("Sec-WebSocket-Protocol", "chat") // 响应基础的 chat 协议
  307. socket, err := upgrader.Upgrade(w, req, responseHeader)
  308. if err != nil {
  309. log.Println("Upgrade error:", err)
  310. return
  311. }
  312. // ----------------------------------- 生成客户端对象 ----------------------------------
  313. // 创建新的客户端对象,使用解析出的参数
  314. client := &Client{
  315. name: clientName(roomParams),
  316. socket: socket, // WebSocket 连接
  317. receive: make(chan []byte, messageBufferSize), // 接收消息的缓冲 channel
  318. room: r, // 使用当前房间实例
  319. userInfo: UserInfo{
  320. UserId: roomParams.UserId,
  321. Username: roomParams.Username,
  322. CustomId: roomParams.CustomId,
  323. ShopId: roomParams.ShopId,
  324. Platform: roomParams.Platform,
  325. Type: roomParams.Type,
  326. StaffId: roomParams.StaffId,
  327. },
  328. }
  329. // 如果用户名为空,生成一个默认用户名
  330. if client.name == "" {
  331. client.name = fmt.Sprintf("user_%d", rand.Intn(9000000))
  332. }
  333. // 将客户端加入当前房间
  334. r.join <- client
  335. defer func() { r.leave <- client }()
  336. // 启动客户端的写消息 goroutine
  337. go client.write()
  338. // 在当前 goroutine 中处理客户端的读消息(阻塞直到连接断开)
  339. client.read()
  340. }