fix: serialize websocket writes in realtime server (#510)

This commit is contained in:
sagit
2026-05-17 21:52:54 +08:00
committed by GitHub
2 changed files with 108 additions and 31 deletions
+44 -31
View File
@@ -30,23 +30,28 @@ type broadcastMessage struct {
Data string `json:"data"` Data string `json:"data"`
} }
type connWrap struct {
conn *websocket.Conn
mu sync.Mutex
}
type nodeSession struct { type nodeSession struct {
nodeID int64 nodeID int64
secret string secret string
conn *websocket.Conn conn *connWrap
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建 crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
} }
type adminSession struct { type adminSession struct {
userID int64 userID int64
claims auth.Claims claims auth.Claims
conn *websocket.Conn conn *connWrap
} }
type monitorSession struct { type monitorSession struct {
userID int64 userID int64
claims auth.Claims claims auth.Claims
conn *websocket.Conn conn *connWrap
} }
type commandResponse struct { type commandResponse struct {
@@ -199,13 +204,14 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request, userID int6
if err != nil { if err != nil {
return return
} }
cw := &connWrap{conn: conn}
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait)) _ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
conn.SetPongHandler(func(string) error { conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsPongWait)) return conn.SetReadDeadline(time.Now().Add(wsPongWait))
}) })
done := make(chan struct{}) done := make(chan struct{})
session := &adminSession{userID: userID, claims: claims, conn: conn} session := &adminSession{userID: userID, claims: claims, conn: cw}
go startKeepalive(conn, done, func() bool { go startKeepalive(cw, done, func() bool {
return s.validateAdminSession(session.userID, session.claims) return s.validateAdminSession(session.userID, session.claims)
}) })
@@ -233,13 +239,14 @@ func (s *Server) handleMonitor(w http.ResponseWriter, r *http.Request, userID in
if err != nil { if err != nil {
return return
} }
cw := &connWrap{conn: conn}
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait)) _ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
conn.SetPongHandler(func(string) error { conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsPongWait)) return conn.SetReadDeadline(time.Now().Add(wsPongWait))
}) })
done := make(chan struct{}) done := make(chan struct{})
session := &monitorSession{userID: userID, claims: claims, conn: conn} session := &monitorSession{userID: userID, claims: claims, conn: cw}
go startKeepalive(conn, done, func() bool { go startKeepalive(cw, done, func() bool {
return s.validateMonitorSession(session.userID, session.claims) return s.validateMonitorSession(session.userID, session.claims)
}) })
@@ -267,12 +274,13 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
if err != nil { if err != nil {
return return
} }
cw := &connWrap{conn: conn}
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait)) _ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
conn.SetPongHandler(func(string) error { conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsPongWait)) return conn.SetReadDeadline(time.Now().Add(wsPongWait))
}) })
done := make(chan struct{}) done := make(chan struct{})
go startKeepalive(conn, done, nil) go startKeepalive(cw, done, nil)
version := r.URL.Query().Get("version") version := r.URL.Query().Get("version")
httpVal := parseIntDefault(r.URL.Query().Get("http"), 0) httpVal := parseIntDefault(r.URL.Query().Get("http"), 0)
@@ -281,15 +289,15 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
s.mu.Lock() s.mu.Lock()
if old, ok := s.nodes[nodeID]; ok { if old, ok := s.nodes[nodeID]; ok {
_ = old.conn.Close() _ = old.conn.conn.Close()
delete(s.byConn, old.conn) delete(s.byConn, old.conn.conn)
} }
// 初始化 AES 加密器并缓存(仅创建一次) // 初始化 AES 加密器并缓存(仅创建一次)
var nodeCrypto *security.AESCrypto var nodeCrypto *security.AESCrypto
if strings.TrimSpace(secret) != "" { if strings.TrimSpace(secret) != "" {
nodeCrypto, _ = security.NewAESCrypto(secret) nodeCrypto, _ = security.NewAESCrypto(secret)
} }
ns := &nodeSession{nodeID: nodeID, secret: secret, conn: conn, crypto: nodeCrypto} ns := &nodeSession{nodeID: nodeID, secret: secret, conn: cw, crypto: nodeCrypto}
s.nodes[nodeID] = ns s.nodes[nodeID] = ns
s.byConn[conn] = ns s.byConn[conn] = ns
s.mu.Unlock() s.mu.Unlock()
@@ -309,7 +317,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
needOfflineBroadcast := false needOfflineBroadcast := false
s.mu.Lock() s.mu.Lock()
current, ok := s.nodes[nodeID] current, ok := s.nodes[nodeID]
if ok && current.conn == conn { if ok && current.conn != nil && current.conn.conn == conn {
delete(s.nodes, nodeID) delete(s.nodes, nodeID)
needOfflineBroadcast = true needOfflineBroadcast = true
} }
@@ -439,7 +447,7 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
s.mu.RLock() s.mu.RLock()
ns, ok := s.nodes[nodeID] ns, ok := s.nodes[nodeID]
s.mu.RUnlock() s.mu.RUnlock()
if !ok || ns == nil || ns.conn == nil { if !ok || ns == nil || ns.conn == nil || ns.conn.conn == nil {
return CommandResult{}, errors.New("节点不在线") return CommandResult{}, errors.New("节点不在线")
} }
@@ -489,9 +497,7 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
} }
} }
_ = ns.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) err = writeWSMessage(ns.conn, websocket.TextMessage, messageData)
err = ns.conn.WriteMessage(websocket.TextMessage, messageData)
_ = ns.conn.SetWriteDeadline(time.Time{})
if err != nil { if err != nil {
cleanup() cleanup()
return CommandResult{}, err return CommandResult{}, err
@@ -635,23 +641,19 @@ func (s *Server) broadcastToRealtime(message string) {
s.mu.RUnlock() s.mu.RUnlock()
for _, c := range admins { for _, c := range admins {
if c == nil || c.conn == nil { if c == nil || c.conn == nil || c.conn.conn == nil {
continue continue
} }
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) err := writeWSMessage(c.conn, websocket.TextMessage, []byte(message))
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
_ = c.conn.SetWriteDeadline(time.Time{})
if err != nil { if err != nil {
log.Printf("websocket broadcast failed: %v", err) log.Printf("websocket broadcast failed: %v", err)
} }
} }
for _, c := range monitors { for _, c := range monitors {
if c == nil || c.conn == nil { if c == nil || c.conn == nil || c.conn.conn == nil {
continue continue
} }
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) err := writeWSMessage(c.conn, websocket.TextMessage, []byte(message))
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
_ = c.conn.SetWriteDeadline(time.Time{})
if err != nil { if err != nil {
log.Printf("websocket broadcast failed: %v", err) log.Printf("websocket broadcast failed: %v", err)
} }
@@ -728,8 +730,21 @@ func (s *Server) validateMonitorSession(userID int64, claims auth.Claims) bool {
return true return true
} }
func startKeepalive(conn *websocket.Conn, done <-chan struct{}, validate func() bool) { func writeWSMessage(cw *connWrap, messageType int, payload []byte) error {
if conn == nil { if cw == nil || cw.conn == nil {
return errors.New("websocket connection not initialized")
}
cw.mu.Lock()
defer cw.mu.Unlock()
_ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
err := cw.conn.WriteMessage(messageType, payload)
_ = cw.conn.SetWriteDeadline(time.Time{})
return err
}
func startKeepalive(cw *connWrap, done <-chan struct{}, validate func() bool) {
if cw == nil || cw.conn == nil {
return return
} }
ticker := time.NewTicker(wsPingPeriod) ticker := time.NewTicker(wsPingPeriod)
@@ -741,14 +756,12 @@ func startKeepalive(conn *websocket.Conn, done <-chan struct{}, validate func()
return return
case <-ticker.C: case <-ticker.C:
if validate != nil && !validate() { if validate != nil && !validate() {
_ = conn.Close() _ = cw.conn.Close()
return return
} }
_ = conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) err := writeWSMessage(cw, websocket.PingMessage, nil)
err := conn.WriteMessage(websocket.PingMessage, nil)
_ = conn.SetWriteDeadline(time.Time{})
if err != nil { if err != nil {
_ = conn.Close() _ = cw.conn.Close()
return return
} }
} }
+64
View File
@@ -5,6 +5,7 @@ import (
"net/http/httptest" "net/http/httptest"
"net/url" "net/url"
"strings" "strings"
"sync"
"testing" "testing"
"go-backend/internal/auth" "go-backend/internal/auth"
@@ -149,3 +150,66 @@ func TestServeHTTPAllowsMonitorTokenWithPermission(t *testing.T) {
} }
_ = conn.Close() _ = conn.Close()
} }
func TestConnWrapSerializesConcurrentWrites(t *testing.T) {
serverConn, clientConn := websocketTestPipe(t)
defer serverConn.Close()
defer clientConn.Close()
cw := &connWrap{conn: serverConn}
readerDone := make(chan struct{})
go func() {
defer close(readerDone)
for i := 0; i < 64; i++ {
if _, _, err := clientConn.ReadMessage(); err != nil {
return
}
}
}()
var wg sync.WaitGroup
for i := 0; i < 64; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if err := writeWSMessage(cw, websocket.TextMessage, []byte("x")); err != nil {
t.Errorf("writeWSMessage() error = %v", err)
}
}()
}
wg.Wait()
_ = clientConn.Close()
<-readerDone
}
func websocketTestPipe(t *testing.T) (*websocket.Conn, *websocket.Conn) {
t.Helper()
upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
serverConnCh := make(chan *websocket.Conn, 1)
serverErrCh := make(chan error, 1)
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
serverErrCh <- err
return
}
serverConnCh <- conn
}))
t.Cleanup(ts.Close)
clientConn, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(ts.URL, "http"), nil)
if err != nil {
t.Fatalf("dial websocket: %v", err)
}
select {
case err := <-serverErrCh:
t.Fatalf("upgrade websocket: %v", err)
case serverConn := <-serverConnCh:
return serverConn, clientConn
}
return nil, nil
}