package websocket import ( "context" "errors" "fmt" "sync" "sync/atomic" "time" "github.com/gogf/gf/v2/container/gmap" "github.com/gogf/gf/v2/encoding/gjson" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/net/ghttp" "github.com/gogf/gf/v2/os/glog" "github.com/gogf/gf/v2/os/grpool" "github.com/google/uuid" "github.com/gorilla/websocket" ) // WsServer 泛化 WebSocket 服务器 type WsServer struct { connections *gmap.StrAnyMap upgrader websocket.Upgrader workerPool *grpool.Pool handlers map[string]MessageHandler handlerMu sync.RWMutex opts ServerOptions svcClosed int32 closeOnce sync.Once } // NewWsServer 创建泛化 WebSocket 服务器 func NewWsServer(opts ...ServerOption) *WsServer { o := defaultOptions() for _, opt := range opts { opt(&o) } return &WsServer{ connections: gmap.NewStrAnyMap(true), upgrader: websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, CheckOrigin: o.checkOrigin, }, workerPool: grpool.New(o.workerPoolSize), handlers: make(map[string]MessageHandler), opts: o, svcClosed: 0, } } // OnMessage 注册业务消息处理器 func (s *WsServer) OnMessage(msgType string, handler MessageHandler) { s.handlerMu.Lock() defer s.handlerMu.Unlock() s.handlers[msgType] = handler } // Upgrade 将 HTTP 连接升级为 WebSocket 并注册到连接池 func (s *WsServer) Upgrade(ctx context.Context, r *ghttp.Request, sessionId string) (*WsConnection, error) { if g.IsEmpty(sessionId) { sessionId = uuid.NewString() } if atomic.LoadInt32(&s.svcClosed) == 1 { return nil, errors.New("websocket server is closed") } if s.connections.Size() >= s.opts.maxConnections { return nil, errors.New("too many online websocket connections") } wsConn, err := s.upgrader.Upgrade(r.Response.Writer, r.Request, nil) if err != nil { return nil, fmt.Errorf("upgrade failed: %w", err) } headers := make(map[string]string) for k, v := range r.Request.Header { if len(v) > 0 { headers[k] = v[0] } } key := s.opts.connKeyPrefix + sessionId // 踢下线旧连接 s.kickOld(key) baseCtx := context.WithoutCancel(ctx) closeCtx, closeCancel := context.WithCancel(baseCtx) wc := &WsConnection{ SessionId: sessionId, Conn: wsConn, Headers: headers, closeCancel: closeCancel, closed: 0, } s.connections.Set(key, wc) // 连接成功回执 _ = s.writeJSON(closeCtx, wc, &WsPushMsg{Type: "ack", Message: "WebSocket连接成功", Data: map[string]any{ "sessionId": sessionId, }}) // Pong 心跳回调,重置读超时 wsConn.SetPongHandler(func(string) error { _ = wsConn.SetReadDeadline(time.Now().Add(s.opts.readTimeout)) return nil }) go s.handleConnection(closeCtx, key, wc) return wc, nil } // PushToSession 向指定会话推送消息 func (s *WsServer) PushToSession(ctx context.Context, sessionId string, msg *WsPushMsg) { key := s.opts.connKeyPrefix + sessionId val := s.connections.Get(key) if val == nil { return } wc, ok := val.(*WsConnection) if !ok || wc.IsClosed() { return } _ = s.writeJSON(ctx, wc, msg) } // GetOnlineSessions 获取在线会话列表 func (s *WsServer) GetOnlineSessions() []string { var sessions []string prefixLen := len(s.opts.connKeyPrefix) s.connections.Iterator(func(key string, _ interface{}) bool { if len(key) > prefixLen { sessions = append(sessions, key[prefixLen:]) } return true }) return sessions } // Close 全局优雅关闭 func (s *WsServer) Close() { s.closeOnce.Do(func() { atomic.StoreInt32(&s.svcClosed, 1) s.workerPool.Close() s.connections.LockFunc(func(m map[string]interface{}) { for _, val := range m { wc, ok := val.(*WsConnection) if !ok { continue } if atomic.CompareAndSwapInt32(&wc.closed, 0, 1) { if wc.closeCancel != nil { wc.closeCancel() } _ = wc.Conn.Close() } } }) s.connections.Clear() }) } // ====================== 内部方法 ====================== // kickOld 踢掉同session旧连接,不再主动remove,由旧连接defer清理 func (s *WsServer) kickOld(key string) { val := s.connections.Get(key) if val == nil { return } old, ok := val.(*WsConnection) if !ok { return } if atomic.CompareAndSwapInt32(&old.closed, 0, 1) { if old.closeCancel != nil { old.closeCancel() } _ = old.Conn.Close() } } // heartbeatLoop 心跳发送协程,入参改为 *WsConnection,复用写锁 func (s *WsServer) heartbeatLoop(ctx context.Context, wc *WsConnection, done <-chan struct{}) { ticker := time.NewTicker(s.opts.heartbeatInterval) defer ticker.Stop() conn := wc.Conn for { select { case <-ticker.C: wc.writeMu.Lock() _ = conn.SetWriteDeadline(time.Now().Add(s.opts.writeTimeout)) err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(s.opts.writeTimeout)) wc.writeMu.Unlock() if err != nil { glog.Debugf(ctx, "heartbeat ping failed: %v", err) return } case <-done: return case <-ctx.Done(): return } } } func (s *WsServer) handleConnection(ctx context.Context, key string, wc *WsConnection) { conn := wc.Conn defer func() { if atomic.CompareAndSwapInt32(&wc.closed, 0, 1) { if wc.closeCancel != nil { wc.closeCancel() } _ = conn.Close() } // 关键修复:只删除自身实例,防止旧连接误删新连接 s.connections.LockFunc(func(m map[string]interface{}) { if v, exist := m[key]; exist && v == wc { delete(m, key) } }) }() done := make(chan struct{}) defer close(done) go s.heartbeatLoop(ctx, wc, done) _ = conn.SetReadDeadline(time.Now().Add(s.opts.readTimeout)) for { select { case <-ctx.Done(): return default: } msgType, data, err := conn.ReadMessage() if err != nil { // 正常关闭不打error日志 if !websocket.IsUnexpectedCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway, websocket.CloseNoStatusReceived, ) { glog.Debugf(ctx, "normal close: %s, err: %v", key, err) } else { glog.Infof(ctx, "unexpected close: %s, err: %v", key, err) } break } _ = conn.SetReadDeadline(time.Now().Add(s.opts.readTimeout)) switch msgType { case websocket.PingMessage: wc.writeMu.Lock() _ = conn.SetWriteDeadline(time.Now().Add(s.opts.writeTimeout)) _ = conn.WriteMessage(websocket.PongMessage, nil) wc.writeMu.Unlock() continue case websocket.CloseMessage: return case websocket.BinaryMessage, websocket.TextMessage: default: continue } if len(data) == 0 { continue } var msg WsMessage if err := gjson.Unmarshal(data, &msg); err != nil { _ = s.writeJSON(ctx, wc, &WsPushMsg{Type: "error", Message: "消息格式错误", Error: err.Error()}) continue } s.handlerMu.RLock() handler, exists := s.handlers[msg.Type] s.handlerMu.RUnlock() if !exists { _ = s.writeJSON(ctx, wc, &WsPushMsg{Type: "error", Message: fmt.Sprintf("未知消息类型: %s", msg.Type)}) continue } // 【重要修复】投递到workerPool,避免业务阻塞读循环 taskCtx := ctx payload := msg.Payload if err := s.workerPool.Add(taskCtx, func(ctx context.Context) { handler(ctx, wc, payload) }); err != nil { _ = s.writeJSON(ctx, wc, &WsPushMsg{ Type: "error", Message: "服务繁忙,任务队列已满", }) } } } // writeJSON 统一写入消息,入参改为 *WsConnection,带并发写锁 func (s *WsServer) writeJSON(ctx context.Context, wc *WsConnection, data interface{}) error { wc.writeMu.Lock() defer wc.writeMu.Unlock() jsonBytes, err := gjson.Encode(data) if err != nil { glog.Errorf(ctx, "json encode failed: %v", err) return err } _ = wc.Conn.SetWriteDeadline(time.Now().Add(s.opts.writeTimeout)) if err = wc.Conn.WriteMessage(websocket.TextMessage, jsonBytes); err != nil { glog.Debugf(ctx, "websocket write failed: %v", err) return err } return nil }