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,9 +12,9 @@ import (
"syscall"
"time"
"github.com/gliderlabs/ssh"
"github.com/henrygd/beszel/agent/health"
"github.com/henrygd/beszel/agent/utils"
"github.com/henrygd/beszel/internal/entities/system"
)
// ConnectionManager manages the connection state and events for the agent.
@@ -22,16 +22,16 @@ import (
// them based on availability and managing reconnection attempts.
type ConnectionManager struct {
agent *Agent // Reference to the parent agent
// mu guards State and isConnecting, which are read and written from both
// the main event loop and the goroutine spawned by connect().
// mu guards state shared by the event loop, connection attempts and SSH callbacks.
mu sync.Mutex
State ConnectionState // Current connection state
eventChan chan ConnectionEvent // Channel for connection events
sshChanged chan struct{} // Coalesced, nonblocking SSH connection notifications
wsClient *WebSocketClient // WebSocket client for hub communication
serverOptions ServerOptions // Configuration for SSH server
wsTicker *time.Ticker // Ticker for WebSocket connection attempts
isConnecting bool // Prevents multiple simultaneous reconnection attempts
ConnectionType system.ConnectionType
sshConnections int // Authenticated SSH TCP connections, not sessions
}
// ConnectionState represents the current connection state of the agent.
@@ -60,8 +60,9 @@ const wsTickerInterval = 10 * time.Second
// newConnectionManager creates a new connection manager for the given agent.
func newConnectionManager(agent *Agent) *ConnectionManager {
cm := &ConnectionManager{
agent: agent,
State: Disconnected,
agent: agent,
State: Disconnected,
sshChanged: make(chan struct{}, 1),
}
return cm
}
@@ -89,6 +90,58 @@ func (c *ConnectionManager) getState() ConnectionState {
return c.State
}
func (c *ConnectionManager) hasSSHConnection() bool {
c.mu.Lock()
defer c.mu.Unlock()
return c.sshConnections > 0
}
func (c *ConnectionManager) notifySSHChange() {
select {
case c.sshChanged <- struct{}{}:
default:
}
}
// sshConnectionTrackedKey marks an SSH connection context as already counted.
type sshConnectionTrackedKey struct{}
// sshConnectionOpened tracks the authenticated TCP connection. Individual SSH
// sessions are short-lived and must not trigger a return to WebSocket.
//
// It is called from the session handler rather than the public key handler,
// which runs when a key is offered and before the client has proven it holds
// the private key. A connection is counted once however many sessions it opens.
func (c *ConnectionManager) sshConnectionOpened(ctx ssh.Context) {
ctx.Lock()
tracked := ctx.Value(sshConnectionTrackedKey{}) != nil
if !tracked {
ctx.SetValue(sshConnectionTrackedKey{}, true)
}
ctx.Unlock()
if tracked {
return
}
c.mu.Lock()
c.sshConnections++
first := c.sshConnections == 1
c.mu.Unlock()
if first {
c.notifySSHChange()
}
go func() {
<-ctx.Done()
c.mu.Lock()
c.sshConnections--
last := c.sshConnections == 0
c.mu.Unlock()
if last {
c.notifySSHChange()
}
}()
}
// setConnecting sets the isConnecting flag and reports its previous value.
func (c *ConnectionManager) setConnecting(v bool) (previous bool) {
c.mu.Lock()
@@ -148,10 +201,14 @@ func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
select {
case connectionEvent := <-c.eventChan:
c.handleEvent(connectionEvent)
case <-c.sshChanged:
c.handleSSHChange()
case <-c.wsTicker.C:
// skip if connect() is still running its own attempt
if !c.isConnectingNow() {
_ = c.startWebSocketConnection()
if err := c.startWebSocketConnection(); err != nil {
c.startSSHServer()
}
}
case <-healthTicker:
_ = health.Update()
@@ -162,6 +219,14 @@ func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
}
}
func (c *ConnectionManager) handleSSHChange() {
if c.hasSSHConnection() {
c.handleEvent(SSHConnect)
} else {
c.handleEvent(SSHDisconnect)
}
}
// stop does not stop the connection manager itself, just any active connections. The manager will attempt to reconnect after stopping, so this should only be called immediately before shutting down the entire agent.
//
// If we need or want to expose a graceful Stop method in the future, do something like this to actually stop the manager:
@@ -193,17 +258,29 @@ func (c *ConnectionManager) stop() error {
func (c *ConnectionManager) handleEvent(event ConnectionEvent) {
switch event {
case WebSocketConnect:
if c.wsClient == nil || !c.wsClient.isVerified() {
return // a superseded connection authenticated after a new attempt began
}
// WebSocket is preferred, so it takes over even if an attempt that was
// already in flight authenticates after SSH has connected.
c.handleStateChange(WebSocketConnected)
case SSHConnect:
if c.getState() == Disconnected {
if c.getState() == Disconnected && c.hasSSHConnection() {
c.handleStateChange(SSHConnected)
}
case WebSocketDisconnect:
if c.wsClient != nil && c.wsClient.getConn() != nil {
return // an older connection closed after its replacement was installed
}
if c.getState() == WebSocketConnected {
c.handleStateChange(Disconnected)
} else if c.getState() == Disconnected {
// The WebSocket upgrade can succeed before authentication fails.
// In that case Connect returned nil, so its error path cannot start SSH.
c.startSSHServer()
}
case SSHDisconnect:
if c.getState() == SSHConnected {
if c.getState() == SSHConnected && !c.hasSSHConnection() {
c.handleStateChange(Disconnected)
}
}
@@ -223,18 +300,17 @@ func (c *ConnectionManager) handleStateChange(newState ConnectionState) {
switch newState {
case WebSocketConnected:
slog.Info("WebSocket connected", "host", c.wsClient.hubURL.Host)
c.ConnectionType = system.ConnectionTypeWebSocket
c.stopWsTicker()
_ = c.agent.StopServer()
c.setConnecting(false)
case SSHConnected:
// stop new ws connection attempts
slog.Info("SSH connection established")
c.ConnectionType = system.ConnectionTypeSSH
c.stopWsTicker()
c.setConnecting(false)
case Disconnected:
c.ConnectionType = system.ConnectionTypeNone
// Listen for SSH whenever disconnected so the hub can fall back to it
// or redial straight away. WebSocket is still tried first below and
// stops the server if it connects.
c.startSSHServer()
// Always keep the ticker running while disconnected. A pending WebSocket
// handshake started by connect() can fail asynchronously (e.g. the hub
// closes the socket, or the deadline set in OnOpen expires) after
@@ -274,7 +350,6 @@ func (c *ConnectionManager) connect() {
}
if c.getState() == Disconnected {
c.startSSHServer()
c.startWsTicker()
}
}
}
@@ -301,9 +376,28 @@ func (c *ConnectionManager) startWebSocketConnection() error {
// startSSHServer starts the SSH server if the agent is currently disconnected.
func (c *ConnectionManager) startSSHServer() {
if c.getState() == Disconnected {
go c.agent.StartServer(c.serverOptions)
c.mu.Lock()
if c.State != Disconnected {
c.mu.Unlock()
return
}
if disabled, _ := utils.GetEnv("DISABLE_SSH"); disabled == "true" {
c.mu.Unlock()
return
}
server, listener, err := c.agent.prepareSSHServer(c.serverOptions)
c.mu.Unlock()
if err != nil {
if !errors.Is(err, errSSHServerRunning) {
slog.Warn("SSH server failed to start", "err", err)
}
return
}
go func() {
if err := c.agent.serveSSHServer(server, listener); err != nil && !errors.Is(err, ssh.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
slog.Warn("SSH server stopped", "err", err)
}
}()
}
// closeWebSocket closes the WebSocket connection if it exists.