fix(agent): improve WebSocket and SSH fallback handling (#2441)

Make WebSocket reconnect and SSH fallback transitions reliable across asynchronous disconnects, stale callbacks, and overlapping connections. Keep the SSH listener available while disconnected and allow a verified WebSocket connection to take precedence when it recovers.

Co-authored-by: henrygd <hank@henrygd.me>
This commit is contained in:
spatiumstas
2026-10-01 20:33:19 +03:00
committed by henrygd
parent d8c2b1f310
commit 4ffd83677d
10 changed files with 688 additions and 65 deletions

View File

@@ -12,11 +12,13 @@ import (
"os"
"path"
"strings"
"sync"
"time"
"github.com/henrygd/beszel"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/common"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/fxamacker/cbor/v2"
"github.com/lxzan/gws"
@@ -53,6 +55,7 @@ type WebSocketClient struct {
gws.BuiltinEventHandler
options *gws.ClientOption // WebSocket client configuration options
agent *Agent // Reference to the parent agent
connMu sync.RWMutex // Guards Conn and hubVerified across callbacks
Conn *gws.Conn // Active WebSocket connection
hubURL *url.URL // Parsed hub URL for connection
token string // Authentication token for hub registration
@@ -203,12 +206,16 @@ func (client *WebSocketClient) Connect() (err error) {
// make sure previous connection is closed
client.Close()
client.Conn, _, err = gws.NewClient(client, client.getOptions())
conn, _, err := gws.NewClient(client, client.getOptions())
if err != nil {
return err
}
client.connMu.Lock()
client.Conn = conn
client.hubVerified = false
client.connMu.Unlock()
go client.Conn.ReadLoop()
go conn.ReadLoop()
return nil
}
@@ -222,6 +229,14 @@ func (client *WebSocketClient) OnOpen(conn *gws.Conn) {
// OnClose handles WebSocket connection closure.
// It logs the closure reason and notifies the connection manager.
func (client *WebSocketClient) OnClose(conn *gws.Conn, err error) {
client.connMu.Lock()
if client.Conn != conn {
client.connMu.Unlock()
return
}
client.Conn = nil
client.hubVerified = false
client.connMu.Unlock()
if err != nil {
slog.Warn("Connection closed", "err", strings.TrimPrefix(err.Error(), "gws: "))
}
@@ -232,6 +247,9 @@ func (client *WebSocketClient) OnClose(conn *gws.Conn, err error) {
// It decodes CBOR messages and routes them to appropriate handlers.
func (client *WebSocketClient) OnMessage(conn *gws.Conn, message *gws.Message) {
defer message.Close()
if client.getConn() != conn {
return
}
conn.SetDeadline(time.Now().Add(wsDeadline))
if message.Opcode != gws.OpcodeBinary {
@@ -246,7 +264,7 @@ func (client *WebSocketClient) OnMessage(conn *gws.Conn, message *gws.Message) {
return
}
if err := client.handleHubRequest(&HubRequest, HubRequest.Id); err != nil {
if err := client.handleHubRequest(&HubRequest, HubRequest.Id, conn); err != nil {
slog.Error("Error handling message", "err", err)
}
}
@@ -259,7 +277,7 @@ func (client *WebSocketClient) OnPing(conn *gws.Conn, message []byte) {
}
// handleAuthChallenge verifies the authenticity of the hub and returns the system's fingerprint.
func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.RawMessage], requestID *uint32) (err error) {
func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.RawMessage], requestID *uint32, conn *gws.Conn) (err error) {
var authRequest common.FingerprintRequest
if err := cbor.Unmarshal(msg.Data, &authRequest); err != nil {
return err
@@ -269,7 +287,13 @@ func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.R
return err
}
client.connMu.Lock()
if conn != nil && client.Conn != conn {
client.connMu.Unlock()
return gws.ErrConnClosed
}
client.hubVerified = true
client.connMu.Unlock()
client.agent.connectionManager.eventChan <- WebSocketConnect
response := &common.FingerprintResponse{
@@ -283,6 +307,9 @@ func (client *WebSocketClient) handleAuthChallenge(msg *common.HubRequest[cbor.R
_, response.Port, _ = net.SplitHostPort(serverAddr)
}
if conn != nil {
return client.sendResponseOnConn(conn, response, requestID)
}
return client.sendResponse(response, requestID)
}
@@ -303,35 +330,65 @@ func (client *WebSocketClient) verifySignature(signature []byte) (err error) {
// Close closes the WebSocket connection gracefully.
// This method is safe to call multiple times.
func (client *WebSocketClient) Close() {
if client.Conn != nil {
_ = client.Conn.WriteClose(1000, nil)
if conn := client.getConn(); conn != nil {
_ = conn.WriteClose(1000, nil)
}
}
func (client *WebSocketClient) getConn() *gws.Conn {
client.connMu.RLock()
defer client.connMu.RUnlock()
return client.Conn
}
func (client *WebSocketClient) isVerified() bool {
client.connMu.RLock()
defer client.connMu.RUnlock()
return client.Conn != nil && client.hubVerified
}
// handleHubRequest routes the request to the appropriate handler using the handler registry.
func (client *WebSocketClient) handleHubRequest(msg *common.HubRequest[cbor.RawMessage], requestID *uint32) error {
func (client *WebSocketClient) handleHubRequest(msg *common.HubRequest[cbor.RawMessage], requestID *uint32, conn *gws.Conn) error {
client.connMu.RLock()
verified := client.hubVerified
client.connMu.RUnlock()
sendResponse := client.sendResponse
if conn != nil {
sendResponse = func(data any, requestID *uint32) error {
return client.sendResponseOnConn(conn, data, requestID)
}
}
ctx := &HandlerContext{
Client: client,
Agent: client.agent,
Request: msg,
RequestID: requestID,
HubVerified: client.hubVerified,
SendResponse: client.sendResponse,
Client: client,
Conn: conn,
Agent: client.agent,
Request: msg,
RequestID: requestID,
HubVerified: verified,
ConnectionType: system.ConnectionTypeWebSocket,
SendResponse: sendResponse,
}
return client.agent.handlerRegistry.Handle(ctx)
}
// sendMessage encodes the given data to CBOR and sends it as a binary message over the WebSocket connection to the hub.
func (client *WebSocketClient) sendMessage(data any) error {
return client.sendMessageOnConn(client.getConn(), data)
}
func (client *WebSocketClient) sendMessageOnConn(conn *gws.Conn, data any) error {
bytes, err := cbor.Marshal(data)
if err != nil {
return err
}
err = client.Conn.WriteMessage(gws.OpcodeBinary, bytes)
if conn == nil {
return gws.ErrConnClosed
}
err = conn.WriteMessage(gws.OpcodeBinary, bytes)
if err != nil {
// If writing fails (e.g., broken pipe due to network issues),
// close the connection to trigger reconnection logic (#1263)
client.Close()
_ = conn.WriteClose(1000, nil)
}
return err
}
@@ -340,12 +397,16 @@ func (client *WebSocketClient) sendMessage(data any) error {
// For ID-based requests, we must populate legacy typed fields for backward
// compatibility with older hubs (<= 0.17) that don't read the generic Data field.
func (client *WebSocketClient) sendResponse(data any, requestID *uint32) error {
return client.sendResponseOnConn(client.getConn(), data, requestID)
}
func (client *WebSocketClient) sendResponseOnConn(conn *gws.Conn, data any, requestID *uint32) error {
if requestID != nil {
response := newAgentResponse(data, requestID)
return client.sendMessage(response)
return client.sendMessageOnConn(conn, response)
}
// Legacy format - send data directly
return client.sendMessage(data)
return client.sendMessageOnConn(conn, data)
}
// getUserAgent returns one of two User-Agent strings based on current time.