mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
fix: serialize websocket writes in realtime server (#510)
This commit is contained in:
@@ -30,23 +30,28 @@ type broadcastMessage struct {
|
||||
Data string `json:"data"`
|
||||
}
|
||||
|
||||
type connWrap struct {
|
||||
conn *websocket.Conn
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
type nodeSession struct {
|
||||
nodeID int64
|
||||
secret string
|
||||
conn *websocket.Conn
|
||||
conn *connWrap
|
||||
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
|
||||
}
|
||||
|
||||
type adminSession struct {
|
||||
userID int64
|
||||
claims auth.Claims
|
||||
conn *websocket.Conn
|
||||
conn *connWrap
|
||||
}
|
||||
|
||||
type monitorSession struct {
|
||||
userID int64
|
||||
claims auth.Claims
|
||||
conn *websocket.Conn
|
||||
conn *connWrap
|
||||
}
|
||||
|
||||
type commandResponse struct {
|
||||
@@ -199,13 +204,14 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request, userID int6
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
cw := &connWrap{conn: conn}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
session := &adminSession{userID: userID, claims: claims, conn: conn}
|
||||
go startKeepalive(conn, done, func() bool {
|
||||
session := &adminSession{userID: userID, claims: claims, conn: cw}
|
||||
go startKeepalive(cw, done, func() bool {
|
||||
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 {
|
||||
return
|
||||
}
|
||||
cw := &connWrap{conn: conn}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
session := &monitorSession{userID: userID, claims: claims, conn: conn}
|
||||
go startKeepalive(conn, done, func() bool {
|
||||
session := &monitorSession{userID: userID, claims: claims, conn: cw}
|
||||
go startKeepalive(cw, done, func() bool {
|
||||
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 {
|
||||
return
|
||||
}
|
||||
cw := &connWrap{conn: conn}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
go startKeepalive(conn, done, nil)
|
||||
go startKeepalive(cw, done, nil)
|
||||
|
||||
version := r.URL.Query().Get("version")
|
||||
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()
|
||||
if old, ok := s.nodes[nodeID]; ok {
|
||||
_ = old.conn.Close()
|
||||
delete(s.byConn, old.conn)
|
||||
_ = old.conn.conn.Close()
|
||||
delete(s.byConn, old.conn.conn)
|
||||
}
|
||||
// 初始化 AES 加密器并缓存(仅创建一次)
|
||||
var nodeCrypto *security.AESCrypto
|
||||
if strings.TrimSpace(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.byConn[conn] = ns
|
||||
s.mu.Unlock()
|
||||
@@ -309,7 +317,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
needOfflineBroadcast := false
|
||||
s.mu.Lock()
|
||||
current, ok := s.nodes[nodeID]
|
||||
if ok && current.conn == conn {
|
||||
if ok && current.conn != nil && current.conn.conn == conn {
|
||||
delete(s.nodes, nodeID)
|
||||
needOfflineBroadcast = true
|
||||
}
|
||||
@@ -439,7 +447,7 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
|
||||
s.mu.RLock()
|
||||
ns, ok := s.nodes[nodeID]
|
||||
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("节点不在线")
|
||||
}
|
||||
|
||||
@@ -489,9 +497,7 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
|
||||
}
|
||||
}
|
||||
|
||||
_ = ns.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err = ns.conn.WriteMessage(websocket.TextMessage, messageData)
|
||||
_ = ns.conn.SetWriteDeadline(time.Time{})
|
||||
err = writeWSMessage(ns.conn, websocket.TextMessage, messageData)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return CommandResult{}, err
|
||||
@@ -635,23 +641,19 @@ func (s *Server) broadcastToRealtime(message string) {
|
||||
s.mu.RUnlock()
|
||||
|
||||
for _, c := range admins {
|
||||
if c == nil || c.conn == nil {
|
||||
if c == nil || c.conn == nil || c.conn.conn == nil {
|
||||
continue
|
||||
}
|
||||
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
|
||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
||||
err := writeWSMessage(c.conn, websocket.TextMessage, []byte(message))
|
||||
if err != nil {
|
||||
log.Printf("websocket broadcast failed: %v", err)
|
||||
}
|
||||
}
|
||||
for _, c := range monitors {
|
||||
if c == nil || c.conn == nil {
|
||||
if c == nil || c.conn == nil || c.conn.conn == nil {
|
||||
continue
|
||||
}
|
||||
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
|
||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
||||
err := writeWSMessage(c.conn, websocket.TextMessage, []byte(message))
|
||||
if err != nil {
|
||||
log.Printf("websocket broadcast failed: %v", err)
|
||||
}
|
||||
@@ -728,8 +730,21 @@ func (s *Server) validateMonitorSession(userID int64, claims auth.Claims) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func startKeepalive(conn *websocket.Conn, done <-chan struct{}, validate func() bool) {
|
||||
if conn == nil {
|
||||
func writeWSMessage(cw *connWrap, messageType int, payload []byte) error {
|
||||
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
|
||||
}
|
||||
ticker := time.NewTicker(wsPingPeriod)
|
||||
@@ -741,14 +756,12 @@ func startKeepalive(conn *websocket.Conn, done <-chan struct{}, validate func()
|
||||
return
|
||||
case <-ticker.C:
|
||||
if validate != nil && !validate() {
|
||||
_ = conn.Close()
|
||||
_ = cw.conn.Close()
|
||||
return
|
||||
}
|
||||
_ = conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := conn.WriteMessage(websocket.PingMessage, nil)
|
||||
_ = conn.SetWriteDeadline(time.Time{})
|
||||
err := writeWSMessage(cw, websocket.PingMessage, nil)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
_ = cw.conn.Close()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
@@ -149,3 +150,66 @@ func TestServeHTTPAllowsMonitorTokenWithPermission(t *testing.T) {
|
||||
}
|
||||
_ = 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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user