You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
nnet/internal/server/websocket_server.go

344 lines
8.2 KiB
Go

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

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
}