mirror of
https://github.com/henrygd/beszel.git
synced 2026-09-29 12:57:50 +02:00
fix(hub): synchronise WebSocket connection access (#2453)
Co-authored-by: user01010111 <lapses.50.booster@icloud.com>
This commit is contained in:
@@ -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).
|
||||||
|
|||||||
64
internal/hub/ws/ws_lifecycle_test.go
Normal file
64
internal/hub/ws/ws_lifecycle_test.go
Normal 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)
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user