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 (
"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).

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"))
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")