package wsc import ( "encoding/json" "errors" "io" "sync" "time" "github.com/gorilla/websocket" "github.com/sirupsen/logrus" ) // ============================================================ // 消息帧格式 // ============================================================ // Message 统一消息结构(客户端 <-> 服务端)。 type Message struct { Action string `json:"action"` Payload json.RawMessage `json:"payload,omitempty"` } // ============================================================ // 错误定义 // ============================================================ var ( ErrConnectionClosed = errors.New("wsc: connection closed") ErrSendBufferFull = errors.New("wsc: send buffer full") ) // ============================================================ // Session — 单连接会话 // ============================================================ const ( writeWait = 10 * time.Second pongWait = 60 * time.Second pingPeriod = (pongWait * 9) / 10 maxMessageSize = 65536 ) // Session 管理单个 WebSocket 连接。 type Session struct { wsCtx *Context conn *websocket.Conn server *Server send chan []byte closeCh chan struct{} once sync.Once } // writeJSON 序列化并发送 JSON 到写通道(单次 marshal)。 func (s *Session) writeJSON(v any) error { data, err := json.Marshal(v) if err != nil { return err } return s.sendRaw(data) } // sendRaw 将已序列化的字节写入发送通道。 func (s *Session) sendRaw(data []byte) error { select { case s.send <- data: return nil case <-s.closeCh: return ErrConnectionClosed default: logrus.Warnf("[wsc] send buffer full, dropping message for session=%s", s.wsCtx.SessionID[:8]) return ErrSendBufferFull } } // writePump 写协程,从 send 通道取数据写入 WebSocket。 func (s *Session) writePump() { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() s.close() }() for { select { case message := <-s.send: s.conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := s.conn.WriteMessage(websocket.TextMessage, message); err != nil { return } case <-ticker.C: s.conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := s.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } case <-s.closeCh: return } } } // readPump 读协程,读取消息并路由。 func (s *Session) readPump(router *Router) { defer func() { if r := recover(); r != nil { logrus.Errorf("[wsc] readPump panic recovered: %v, session=%s", r, s.wsCtx.SessionID[:8]) } s.close() }() s.conn.SetReadLimit(maxMessageSize) s.conn.SetReadDeadline(time.Now().Add(pongWait)) s.conn.SetPongHandler(func(string) error { s.conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) for { _, data, err := s.conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) && !errors.Is(err, io.EOF) { logrus.Errorf("[wsc] read error: %v", err) } return } var msg Message if err := json.Unmarshal(data, &msg); err != nil { s.wsCtx.WriteError(400, "invalid message format") continue } router.dispatch(s.wsCtx, &msg) } } // close 安全关闭连接。 func (s *Session) close() { s.once.Do(func() { close(s.closeCh) s.wsCtx.cancelCtx() // 离开所有房间 // 先快照房间列表再释放读锁,避免 leave 触发 onEmpty → RemoveRoom 抢写锁死锁 s.server.mu.RLock() rooms := make([]*Room, 0, len(s.server.rooms)) for _, room := range s.server.rooms { rooms = append(rooms, room) } s.server.mu.RUnlock() for _, room := range rooms { room.Leave(s.wsCtx) } // 触发断开回调 if s.server.OnDisconnect != nil { s.server.OnDisconnect(s.wsCtx) } s.server.onSessionClosed() s.conn.Close() // s.send 不关闭:readPump 的 dispatch 可能在 close 生效后仍并发回包, // 向已关闭 channel 发送会 panic。writePump 退出由 closeCh 驱动, // 残留消息随 Session 对象一起被 GC 回收。 }) }