room.go 13 KB

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