package server import ( "context" "fmt" "net" "net/http" "strings" "sync" "time" "git.noahlan.cn/noahlan/nnet/v2/internal/connection" "git.noahlan.cn/noahlan/nnet/v2/pkg/config" pkgerrors "git.noahlan.cn/noahlan/nnet/v2/pkg/errors" protocolpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/protocol" unpackerpkg "git.noahlan.cn/noahlan/nnet/v2/pkg/unpacker" "github.com/gorilla/websocket" ) // WebSocketServer is a real WebSocket transport server. It reuses the same // router, codec, protocol and connection manager stack as the TCP server. type WebSocketServer struct { *Server addr string isWSS bool httpServer *http.Server listener net.Listener upgrader websocket.Upgrader ctx context.Context cancel context.CancelFunc mu sync.RWMutex started bool msgHandler *messageHandler handlerTimeout time.Duration } // NewWebSocketServer 创建WebSocket服务器(ws或wss) func NewWebSocketServer(cfg *config.Config) (*WebSocketServer, error) { if cfg == nil { cfg = config.DefaultConfig() cfg.Addr = "ws://:6995" } addr, isWSS := normalizeWebSocketAddr(cfg.Addr) if isWSS && !cfg.TLSEnabled { cfg.TLSEnabled = true } gnetSrv, err := newGnetServer(cfg, ProtocolTCP) if err != nil { return nil, err } base := newServerWithGnet(gnetSrv, cfg) configHelper := newServerConfigHelper(cfg) ctx, cancel := context.WithCancel(context.Background()) return &WebSocketServer{ Server: base, addr: addr, isWSS: isWSS, upgrader: websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, }, ctx: ctx, cancel: cancel, msgHandler: gnetSrv.eventHandler.msgHandler, handlerTimeout: configHelper.HandlerTimeout(), }, nil } // NewWSSServer 创建WSS服务器(WebSocket Secure) func NewWSSServer(cfg *config.Config) (*WebSocketServer, error) { if cfg == nil { cfg = config.DefaultConfig() } cfg.TLSEnabled = true addr := cfg.Addr if strings.HasPrefix(addr, "ws://") { cfg.Addr = "wss://" + strings.TrimPrefix(addr, "ws://") } else if !strings.HasPrefix(addr, "wss://") { cfg.Addr = "wss://" + addr } return NewWebSocketServer(cfg) } func normalizeWebSocketAddr(addr string) (string, bool) { switch { case strings.HasPrefix(addr, "wss://"): addr = strings.TrimPrefix(addr, "wss://") return defaultListenAddr(addr), true case strings.HasPrefix(addr, "ws://"): addr = strings.TrimPrefix(addr, "ws://") return defaultListenAddr(addr), false default: return defaultListenAddr(addr), false } } func defaultListenAddr(addr string) string { if addr == "" || addr == ":" { return ":6995" } return addr } // Start 启动WebSocket服务器。与TCP服务器保持一致:启动监听后返回。 func (s *WebSocketServer) Start() error { s.mu.Lock() if s.started { s.mu.Unlock() return pkgerrors.ErrServerAlreadyStarted } s.ctx, s.cancel = context.WithCancel(context.Background()) mux := http.NewServeMux() mux.HandleFunc("/", s.handleWebSocket) httpServer := &http.Server{Handler: mux} listener, err := net.Listen("tcp", s.addr) if err != nil { s.mu.Unlock() return fmt.Errorf("failed to listen on %s: %w", s.addr, err) } s.listener = listener s.httpServer = httpServer s.started = true s.gnetServer.mu.Lock() s.gnetServer.started = true s.gnetServer.mu.Unlock() s.mu.Unlock() for _, hook := range s.serverLifecycleHooks { if err := hook.OnInit(); err != nil { _ = listener.Close() return pkgerrors.New("failed to execute OnInit hook").WithCause(err) } } for _, hook := range s.serverLifecycleHooks { if err := hook.OnStart(); err != nil { s.gnetServer.logger.Error("OnStart hook error: %v", err) } } s.gnetServer.logger.Info("WebSocket server started on %s", listener.Addr().String()) go s.serve(listener, httpServer) return nil } func (s *WebSocketServer) serve(listener net.Listener, httpServer *http.Server) { var err error if s.isWSS { if s.config.TLS == nil || s.config.TLS.CertFile == "" || s.config.TLS.KeyFile == "" { err = pkgerrors.New("wss requires TLS cert and key files") } else { err = httpServer.ServeTLS(listener, s.config.TLS.CertFile, s.config.TLS.KeyFile) } } else { err = httpServer.Serve(listener) } if err != nil && err != http.ErrServerClosed { s.gnetServer.logger.Error("WebSocket server error: %v", err) } } // Stop 停止WebSocket服务器。 func (s *WebSocketServer) Stop() error { s.mu.Lock() if !s.started { s.mu.Unlock() return pkgerrors.ErrServerNotStarted } httpServer := s.httpServer shutdownTimeout := s.config.ShutdownTimeout if shutdownTimeout <= 0 { shutdownTimeout = 30 * time.Second } s.started = false s.gnetServer.mu.Lock() s.gnetServer.started = false s.gnetServer.mu.Unlock() s.cancel() s.mu.Unlock() s.gnetServer.logger.Info("Stopping WebSocket server...") ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) defer cancel() var stopErr error if httpServer != nil { stopErr = httpServer.Shutdown(ctx) } for _, conn := range s.connManager.GetAll() { _ = conn.Close() } ticker := time.NewTicker(100 * time.Millisecond) defer ticker.Stop() timeout := time.NewTimer(shutdownTimeout) defer timeout.Stop() waitConnections: for { if s.connManager.Count() == 0 { break } select { case <-ticker.C: case <-timeout.C: s.forceCloseAllConnections() break waitConnections } if s.connManager.Count() == 0 { break } } for _, hook := range s.serverLifecycleHooks { if err := hook.OnStop(); err != nil { s.gnetServer.logger.Error("OnStop hook error: %v", err) } } s.gnetServer.logger.Info("WebSocket server stopped") return stopErr } // Started 检查服务器是否已启动。 func (s *WebSocketServer) Started() bool { s.mu.RLock() defer s.mu.RUnlock() return s.started } func (s *WebSocketServer) handleWebSocket(w http.ResponseWriter, r *http.Request) { wsConn, err := s.upgrader.Upgrade(w, r, nil) if err != nil { s.gnetServer.logger.Debug("WebSocket upgrade failed: %v", err) return } conn := connection.NewWebSocketConnection("", wsConn) connID := conn.ID() if err := s.connManager.Add(conn); err != nil { s.gnetServer.logger.Error("Failed to add WebSocket connection: %v", err) _ = wsConn.Close() return } for _, hook := range s.connLifecycleHooks { if err := hook.OnOpen(connID, conn.RemoteAddr()); err != nil { s.gnetServer.logger.Error("OnOpen hook error: %v", err) } } protocol := s.resolveApplicationProtocol() unpacker := newServerProtocolUnpacker(protocol) ctxConn := toContextConnection(conn) defer func() { for _, hook := range s.connLifecycleHooks { if err := hook.OnClose(connID, nil); err != nil { s.gnetServer.logger.Error("OnClose hook error: %v", err) } } _ = s.connManager.Remove(connID) }() for { select { case <-s.ctx.Done(): return default: } _, data, err := wsConn.ReadMessage() if err != nil { return } conn.UpdateActive() messages, err := splitWebSocketMessages(data, unpacker) if err != nil { s.gnetServer.logger.Debug("WebSocket unpack failed: %v", err) return } if len(messages) == 0 { continue } for _, message := range messages { s.msgHandler.handleMessageWithContext( s.ctx, ctxConn, message, protocol, s.codecRegistry, s.handlerTimeout, ) } } } func (s *WebSocketServer) resolveApplicationProtocol() protocolpkg.Protocol { configHelper := newServerConfigHelper(s.config) if !configHelper.IsProtocolEncodeEnabled() { return nil } protocol, err := s.protocolManager.Get(s.config.ApplicationProtocol, "") if err != nil { s.gnetServer.logger.Debug("Failed to resolve WebSocket application protocol: %v", err) return nil } return protocol } type serverUnpackingProtocol interface { Unpacker() unpackerpkg.Unpacker } func newServerProtocolUnpacker(protocol protocolpkg.Protocol) unpackerpkg.Unpacker { if protocol == nil { return nil } if provider, ok := protocol.(serverUnpackingProtocol); ok { return provider.Unpacker() } return nil } func splitWebSocketMessages(data []byte, unpacker unpackerpkg.Unpacker) ([][]byte, error) { if unpacker == nil { msg := make([]byte, len(data)) copy(msg, data) return [][]byte{msg}, nil } messages, _, _, err := unpacker.Unpack(data) return messages, err }