mirror of
https://github.com/henrygd/beszel.git
synced 2026-10-02 14:27: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,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.
|
||||
|
||||
Reference in New Issue
Block a user