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"`
|
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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user