diff --git a/internal/hub/ws/ws.go b/internal/hub/ws/ws.go index def433a25..7d61e596c 100644 --- a/internal/hub/ws/ws.go +++ b/internal/hub/ws/ws.go @@ -3,6 +3,7 @@ package ws import ( "context" "errors" + "sync/atomic" "time" "weak" @@ -26,7 +27,7 @@ type Handler struct { // WsConn represents a WebSocket connection to an agent. type WsConn struct { - conn *gws.Conn + conn atomic.Pointer[gws.Conn] requestManager *RequestManager DownChan chan struct{} agentVersion semver.Version @@ -54,12 +55,13 @@ func GetUpgrader() *gws.Upgrader { // NewWsConnection creates a new WebSocket connection wrapper with agent version. func NewWsConnection(conn *gws.Conn, agentVersion semver.Version) *WsConn { - return &WsConn{ - conn: conn, + ws := &WsConn{ requestManager: NewRequestManager(conn), DownChan: make(chan struct{}, 1), agentVersion: agentVersion, } + ws.conn.Store(conn) + return ws } // 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 { return } - wsConn.(*WsConn).conn = nil + wsConn.(*WsConn).conn.Store(nil) // wait 5 seconds to allow reconnection before setting system down // use a weak pointer to avoid keeping references if the system is removed 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. func (ws *WsConn) Close(msg []byte) { - if ws.IsConnected() { - ws.conn.WriteClose(1000, msg) + if conn := ws.conn.Load(); conn != nil { + conn.WriteClose(1000, msg) } if ws.requestManager != nil { ws.requestManager.Close() @@ -111,24 +113,26 @@ func (ws *WsConn) Close(msg []byte) { // Ping sends a ping frame to keep the connection alive. func (ws *WsConn) Ping() error { - if ws.conn == nil { + conn := ws.conn.Load() + if conn == nil { return gws.ErrConnClosed } - ws.conn.SetDeadline(time.Now().Add(deadline)) - return ws.conn.WritePing(nil) + conn.SetDeadline(time.Now().Add(deadline)) + return conn.WritePing(nil) } // 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. func (ws *WsConn) sendMessage(data common.HubRequest[any]) error { - if ws.conn == nil { + conn := ws.conn.Load() + if conn == nil { return gws.ErrConnClosed } bytes, err := cbor.Marshal(data) if err != nil { 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. @@ -163,7 +167,7 @@ func (ws *WsConn) handleAgentRequest(req *PendingRequest, handler ResponseHandle // IsConnected returns true if the WebSocket connection is active. 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). diff --git a/internal/hub/ws/ws_lifecycle_test.go b/internal/hub/ws/ws_lifecycle_test.go new file mode 100644 index 000000000..abcdc0762 --- /dev/null +++ b/internal/hub/ws/ws_lifecycle_test.go @@ -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) +} diff --git a/internal/hub/ws/ws_test.go b/internal/hub/ws/ws_test.go index bdbc4cb23..9970cf47b 100644 --- a/internal/hub/ws/ws_test.go +++ b/internal/hub/ws/ws_test.go @@ -39,7 +39,7 @@ func TestNewWsConnection(t *testing.T) { wsConn := NewWsConnection(nil, semver.MustParse("0.12.10")) 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.DownChan, "Down channel should be initialized") assert.Equal(t, 1, cap(wsConn.DownChan), "Down channel should have capacity of 1")