mirror of
https://github.com/henrygd/beszel.git
synced 2026-10-02 06:17:47 +02:00
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:
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user