fix(hub): synchronise WebSocket connection access (#2453)

Co-authored-by: user01010111 <lapses.50.booster@icloud.com>
This commit is contained in:
user01010111
2026-09-29 12:12:00 +13:00
committed by GitHub
parent 09a277f074
commit 921ded1af0
3 changed files with 81 additions and 13 deletions

View File

@@ -3,6 +3,7 @@ package ws
import ( import (
"context" "context"
"errors" "errors"
"sync/atomic"
"time" "time"
"weak" "weak"
@@ -26,7 +27,7 @@ type Handler struct {
// WsConn represents a WebSocket connection to an agent. // WsConn represents a WebSocket connection to an agent.
type WsConn struct { type WsConn struct {
conn *gws.Conn conn atomic.Pointer[gws.Conn]
requestManager *RequestManager requestManager *RequestManager
DownChan chan struct{} DownChan chan struct{}
agentVersion semver.Version agentVersion semver.Version
@@ -54,12 +55,13 @@ func GetUpgrader() *gws.Upgrader {
// NewWsConnection creates a new WebSocket connection wrapper with agent version. // NewWsConnection creates a new WebSocket connection wrapper with agent version.
func NewWsConnection(conn *gws.Conn, agentVersion semver.Version) *WsConn { func NewWsConnection(conn *gws.Conn, agentVersion semver.Version) *WsConn {
return &WsConn{ ws := &WsConn{
conn: conn,
requestManager: NewRequestManager(conn), requestManager: NewRequestManager(conn),
DownChan: make(chan struct{}, 1), DownChan: make(chan struct{}, 1),
agentVersion: agentVersion, agentVersion: agentVersion,
} }
ws.conn.Store(conn)
return ws
} }
// OnOpen sets a deadline for the WebSocket connection and extracts agent version. // OnOpen sets a deadline for the WebSocket connection and extracts agent version.
@@ -87,7 +89,7 @@ func (h *Handler) OnClose(conn *gws.Conn, err error) {
if !ok { if !ok {
return return
} }
wsConn.(*WsConn).conn = nil wsConn.(*WsConn).conn.Store(nil)
// wait 5 seconds to allow reconnection before setting system down // wait 5 seconds to allow reconnection before setting system down
// use a weak pointer to avoid keeping references if the system is removed // use a weak pointer to avoid keeping references if the system is removed
go func(downChan weak.Pointer[chan struct{}]) { go func(downChan weak.Pointer[chan struct{}]) {
@@ -101,8 +103,8 @@ func (h *Handler) OnClose(conn *gws.Conn, err error) {
// Close terminates the WebSocket connection gracefully. // Close terminates the WebSocket connection gracefully.
func (ws *WsConn) Close(msg []byte) { func (ws *WsConn) Close(msg []byte) {
if ws.IsConnected() { if conn := ws.conn.Load(); conn != nil {
ws.conn.WriteClose(1000, msg) conn.WriteClose(1000, msg)
} }
if ws.requestManager != nil { if ws.requestManager != nil {
ws.requestManager.Close() ws.requestManager.Close()
@@ -111,24 +113,26 @@ func (ws *WsConn) Close(msg []byte) {
// Ping sends a ping frame to keep the connection alive. // Ping sends a ping frame to keep the connection alive.
func (ws *WsConn) Ping() error { func (ws *WsConn) Ping() error {
if ws.conn == nil { conn := ws.conn.Load()
if conn == nil {
return gws.ErrConnClosed return gws.ErrConnClosed
} }
ws.conn.SetDeadline(time.Now().Add(deadline)) conn.SetDeadline(time.Now().Add(deadline))
return ws.conn.WritePing(nil) return conn.WritePing(nil)
} }
// sendMessage encodes data to CBOR and sends it as a binary message to the agent. // sendMessage encodes data to CBOR and sends it as a binary message to the agent.
// This is kept for backwards compatibility but new actions should use RequestManager. // This is kept for backwards compatibility but new actions should use RequestManager.
func (ws *WsConn) sendMessage(data common.HubRequest[any]) error { func (ws *WsConn) sendMessage(data common.HubRequest[any]) error {
if ws.conn == nil { conn := ws.conn.Load()
if conn == nil {
return gws.ErrConnClosed return gws.ErrConnClosed
} }
bytes, err := cbor.Marshal(data) bytes, err := cbor.Marshal(data)
if err != nil { if err != nil {
return err return err
} }
return ws.conn.WriteMessage(gws.OpcodeBinary, bytes) return conn.WriteMessage(gws.OpcodeBinary, bytes)
} }
// handleAgentRequest processes a request to the agent, handling both legacy and new formats. // handleAgentRequest processes a request to the agent, handling both legacy and new formats.
@@ -163,7 +167,7 @@ func (ws *WsConn) handleAgentRequest(req *PendingRequest, handler ResponseHandle
// IsConnected returns true if the WebSocket connection is active. // IsConnected returns true if the WebSocket connection is active.
func (ws *WsConn) IsConnected() bool { func (ws *WsConn) IsConnected() bool {
return ws.conn != nil return ws.conn.Load() != nil
} }
// AgentVersion returns the connected agent's version (as reported during handshake). // AgentVersion returns the connected agent's version (as reported during handshake).

View File

@@ -0,0 +1,64 @@
//go:build testing
package ws
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/blang/semver"
"github.com/henrygd/beszel/internal/common"
"github.com/lxzan/gws"
"github.com/stretchr/testify/require"
)
func TestWsConnConcurrentClose(t *testing.T) {
connections := make(chan *WsConn, 1)
serverDone := make(chan struct{})
upgrader := gws.NewUpgrader(&Handler{}, &gws.ServerOption{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer close(serverDone)
conn, err := upgrader.Upgrade(w, r)
if err != nil {
t.Error(err)
return
}
ws := NewWsConnection(conn, semver.MustParse("0.12.10"))
conn.Session().Store("wsConn", ws)
connections <- ws
conn.ReadLoop()
}))
defer server.Close()
client, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, &gws.ClientOption{
Addr: "ws" + strings.TrimPrefix(server.URL, "http"),
})
require.NoError(t, err)
defer client.WriteClose(1000, nil)
ws := <-connections
require.True(t, ws.IsConnected())
readerReady := make(chan struct{})
readerDone := make(chan struct{})
go func() {
defer close(readerDone)
close(readerReady)
for {
select {
case <-serverDone:
return
default:
ws.IsConnected()
}
}
}()
<-readerReady
require.NoError(t, client.WriteClose(1000, nil))
<-readerDone
require.False(t, ws.IsConnected())
require.ErrorIs(t, ws.Ping(), gws.ErrConnClosed)
require.ErrorIs(t, ws.sendMessage(common.HubRequest[any]{}), gws.ErrConnClosed)
ws.Close(nil)
}

View File

@@ -39,7 +39,7 @@ func TestNewWsConnection(t *testing.T) {
wsConn := NewWsConnection(nil, semver.MustParse("0.12.10")) wsConn := NewWsConnection(nil, semver.MustParse("0.12.10"))
assert.NotNil(t, wsConn, "WebSocket connection should not be nil") assert.NotNil(t, wsConn, "WebSocket connection should not be nil")
assert.Nil(t, wsConn.conn, "Connection should be nil as passed") assert.Nil(t, wsConn.conn.Load(), "Connection should be nil as passed")
assert.NotNil(t, wsConn.requestManager, "Request manager should be initialized") assert.NotNil(t, wsConn.requestManager, "Request manager should be initialized")
assert.NotNil(t, wsConn.DownChan, "Down channel should be initialized") assert.NotNil(t, wsConn.DownChan, "Down channel should be initialized")
assert.Equal(t, 1, cap(wsConn.DownChan), "Down channel should have capacity of 1") assert.Equal(t, 1, cap(wsConn.DownChan), "Down channel should have capacity of 1")