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:
@@ -6,6 +6,7 @@ package agent
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -45,6 +46,8 @@ type Agent struct {
|
||||
connectionManager *ConnectionManager // Channel to signal connection events
|
||||
handlerRegistry *HandlerRegistry // Registry for routing incoming messages
|
||||
server *ssh.Server // SSH server
|
||||
serverListener net.Listener // SSH listener, also closed if Serve has not started yet
|
||||
serverMu sync.Mutex // Guards server and serverListener
|
||||
dataDir string // Directory for persisting data
|
||||
keys []gossh.PublicKey // SSH public keys
|
||||
smartManager *SmartManager // Manages SMART data
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -474,7 +474,7 @@ func TestWebSocketClient_HandleHubRequest(t *testing.T) {
|
||||
Data: cbor.RawMessage{},
|
||||
}
|
||||
|
||||
err := client.handleHubRequest(hubRequest, nil)
|
||||
err := client.handleHubRequest(hubRequest, nil, nil)
|
||||
|
||||
if tc.expectError {
|
||||
assert.Error(t, err)
|
||||
@@ -536,6 +536,18 @@ func TestWebSocketClient_Close(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestWebSocketClient_IgnoresStaleClose(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
agent.connectionManager.eventChan = make(chan ConnectionEvent, 1)
|
||||
current := &gws.Conn{}
|
||||
client := &WebSocketClient{agent: agent, Conn: current, hubVerified: true}
|
||||
|
||||
client.OnClose(&gws.Conn{}, nil)
|
||||
assert.Same(t, current, client.getConn())
|
||||
assert.True(t, client.hubVerified)
|
||||
assert.Empty(t, agent.connectionManager.eventChan)
|
||||
}
|
||||
|
||||
// TestWebSocketClient_ConnectRateLimit tests connection rate limiting
|
||||
func TestWebSocketClient_ConnectRateLimit(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"github.com/lxzan/gws"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/crypto/ssh"
|
||||
@@ -95,6 +96,7 @@ func TestConnectionManager_StateTransitions(t *testing.T) {
|
||||
func TestConnectionManager_EventHandling(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
cm := agent.connectionManager
|
||||
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "true")
|
||||
cm.wsClient = &WebSocketClient{
|
||||
hubURL: &url.URL{
|
||||
Host: "localhost:8080",
|
||||
@@ -112,6 +114,12 @@ func TestConnectionManager_EventHandling(t *testing.T) {
|
||||
event: WebSocketConnect,
|
||||
expectedState: WebSocketConnected,
|
||||
},
|
||||
{
|
||||
name: "WebSocket connect from SSH connected",
|
||||
initialState: SSHConnected,
|
||||
event: WebSocketConnect,
|
||||
expectedState: WebSocketConnected,
|
||||
},
|
||||
{
|
||||
name: "SSH connect from disconnected",
|
||||
initialState: Disconnected,
|
||||
@@ -157,6 +165,20 @@ func TestConnectionManager_EventHandling(t *testing.T) {
|
||||
// and the goroutine would otherwise race with the direct field
|
||||
// writes here and in later subtests.
|
||||
cm.setConnecting(true)
|
||||
cm.mu.Lock()
|
||||
cm.sshConnections = 0
|
||||
if tc.event == SSHConnect {
|
||||
cm.sshConnections = 1
|
||||
}
|
||||
cm.mu.Unlock()
|
||||
cm.wsClient.connMu.Lock()
|
||||
cm.wsClient.Conn = nil
|
||||
cm.wsClient.hubVerified = false
|
||||
if tc.event == WebSocketConnect {
|
||||
cm.wsClient.Conn = &gws.Conn{}
|
||||
cm.wsClient.hubVerified = true
|
||||
}
|
||||
cm.wsClient.connMu.Unlock()
|
||||
cm.State = tc.initialState
|
||||
cm.handleEvent(tc.event)
|
||||
assert.Equal(t, tc.expectedState, cm.State, "State should match expected after event")
|
||||
@@ -238,6 +260,24 @@ func TestConnectionManager_ReconnectionLogic(t *testing.T) {
|
||||
assert.True(t, cm.isConnectingNow(), "Should set isConnecting flag")
|
||||
}
|
||||
|
||||
func TestWebSocketDisconnectStartsSSH(t *testing.T) {
|
||||
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
|
||||
agent := createTestAgent(t)
|
||||
cm := agent.connectionManager
|
||||
cm.serverOptions = createTestServerOptions(t)
|
||||
cm.State = WebSocketConnected
|
||||
cm.setConnecting(true) // keep this test focused on the synchronous fallback
|
||||
defer cm.stopWsTicker()
|
||||
|
||||
cm.handleEvent(WebSocketDisconnect)
|
||||
require.Equal(t, Disconnected, cm.getState())
|
||||
agent.serverMu.Lock()
|
||||
listener := agent.serverListener
|
||||
agent.serverMu.Unlock()
|
||||
require.NotNil(t, listener, "SSH should be ready as soon as an established WS closes")
|
||||
require.NoError(t, agent.StopServer())
|
||||
}
|
||||
|
||||
// TestConnectionManager_TickerSurvivesStaleDisconnect reproduces the freeze from
|
||||
// https://github.com/henrygd/beszel/issues/2326: a reconnect attempt's handshake
|
||||
// can fail asynchronously (after connect() already returned with a nil error)
|
||||
|
||||
@@ -10,17 +10,20 @@ import (
|
||||
"github.com/henrygd/beszel/internal/entities/monitor"
|
||||
"github.com/henrygd/beszel/internal/entities/smart"
|
||||
"github.com/henrygd/beszel/internal/entities/system"
|
||||
"github.com/lxzan/gws"
|
||||
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
// HandlerContext provides context for request handlers
|
||||
type HandlerContext struct {
|
||||
Client *WebSocketClient
|
||||
Agent *Agent
|
||||
Request *common.HubRequest[cbor.RawMessage]
|
||||
RequestID *uint32
|
||||
HubVerified bool
|
||||
Client *WebSocketClient
|
||||
Conn *gws.Conn // WebSocket that carried this request, if any
|
||||
Agent *Agent
|
||||
Request *common.HubRequest[cbor.RawMessage]
|
||||
RequestID *uint32
|
||||
HubVerified bool
|
||||
ConnectionType system.ConnectionType // Transport that carried this request
|
||||
// SendResponse abstracts how a handler sends responses (WS or SSH)
|
||||
SendResponse func(data any, requestID *uint32) error
|
||||
}
|
||||
@@ -101,7 +104,11 @@ func (h *GetDataHandler) Handle(hctx *HandlerContext) error {
|
||||
_ = cbor.Unmarshal(hctx.Request.Data, &options)
|
||||
|
||||
sysStats := hctx.Agent.gatherStats(options)
|
||||
return hctx.SendResponse(sysStats, hctx.RequestID)
|
||||
// Cached stats may be shared by concurrent SSH and WebSocket requests.
|
||||
// Set the transport on the response copy, not on the cached data.
|
||||
response := *sysStats
|
||||
response.Info.ConnectionType = hctx.ConnectionType
|
||||
return hctx.SendResponse(&response, hctx.RequestID)
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////
|
||||
@@ -111,7 +118,7 @@ func (h *GetDataHandler) Handle(hctx *HandlerContext) error {
|
||||
type CheckFingerprintHandler struct{}
|
||||
|
||||
func (h *CheckFingerprintHandler) Handle(hctx *HandlerContext) error {
|
||||
return hctx.Client.handleAuthChallenge(hctx.Request, hctx.RequestID)
|
||||
return hctx.Client.handleAuthChallenge(hctx.Request, hctx.RequestID, hctx.Conn)
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/henrygd/beszel/agent/zfs"
|
||||
"github.com/henrygd/beszel/internal/common"
|
||||
"github.com/henrygd/beszel/internal/entities/smart"
|
||||
"github.com/henrygd/beszel/internal/entities/system"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
@@ -32,6 +33,30 @@ func TestNewAgentResponseSmartData(t *testing.T) {
|
||||
assert.True(t, response.SmartComplete)
|
||||
}
|
||||
|
||||
func TestGetDataHandlerReportsRequestTransport(t *testing.T) {
|
||||
cache := NewSystemDataCache()
|
||||
cached := &system.CombinedData{}
|
||||
cache.Set(cached, defaultDataCacheTimeMs)
|
||||
agent := &Agent{cache: cache}
|
||||
options, err := cbor.Marshal(common.DataRequestOptions{CacheTimeMs: defaultDataCacheTimeMs})
|
||||
assert.NoError(t, err)
|
||||
request := &common.HubRequest[cbor.RawMessage]{Action: common.GetData, Data: options}
|
||||
for _, transport := range []system.ConnectionType{system.ConnectionTypeSSH, system.ConnectionTypeWebSocket} {
|
||||
ctx := &HandlerContext{
|
||||
Agent: agent,
|
||||
Request: request,
|
||||
ConnectionType: transport,
|
||||
SendResponse: func(data any, _ *uint32) error {
|
||||
response := data.(*system.CombinedData)
|
||||
assert.Equal(t, transport, response.Info.ConnectionType)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
assert.NoError(t, (&GetDataHandler{}).Handle(ctx))
|
||||
assert.Equal(t, system.ConnectionTypeNone, cached.Info.ConnectionType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetZfsDataHandlerForceRefresh(t *testing.T) {
|
||||
poolCalls := 0
|
||||
zm := &StoragePoolManager{detailInterval: time.Hour, backends: []*poolBackend{{name: "zfs"}}}
|
||||
|
||||
@@ -34,11 +34,26 @@ type ServerOptions struct {
|
||||
// and begins listening for connections. Returns an error if the server
|
||||
// is already running or if there's an issue starting the server.
|
||||
func (a *Agent) StartServer(opts ServerOptions) error {
|
||||
server, listener, err := a.prepareSSHServer(opts)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return a.serveSSHServer(server, listener)
|
||||
}
|
||||
|
||||
var errSSHServerRunning = errors.New("server already started")
|
||||
|
||||
// prepareSSHServer binds the listener before Serve starts so a concurrent stop
|
||||
// can always close it, including when the WebSocket wins the connection race.
|
||||
func (a *Agent) prepareSSHServer(opts ServerOptions) (*ssh.Server, net.Listener, error) {
|
||||
a.serverMu.Lock()
|
||||
defer a.serverMu.Unlock()
|
||||
|
||||
if disableSSH, _ := utils.GetEnv("DISABLE_SSH"); disableSSH == "true" {
|
||||
return errors.New("SSH disabled")
|
||||
return nil, nil, errors.New("SSH disabled")
|
||||
}
|
||||
if a.server != nil {
|
||||
return errors.New("server already started")
|
||||
return nil, nil, errSSHServerRunning
|
||||
}
|
||||
|
||||
slog.Info("Starting SSH server", "addr", opts.Addr, "network", opts.Network)
|
||||
@@ -46,21 +61,18 @@ func (a *Agent) StartServer(opts ServerOptions) error {
|
||||
if opts.Network == "unix" {
|
||||
// remove existing socket file if it exists
|
||||
if err := os.Remove(opts.Addr); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// start listening on the address
|
||||
ln, err := net.Listen(opts.Network, opts.Addr)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, nil, err
|
||||
}
|
||||
defer ln.Close()
|
||||
|
||||
// set default handler
|
||||
ssh.Handle(a.handleSession)
|
||||
|
||||
a.server = &ssh.Server{
|
||||
server := &ssh.Server{
|
||||
Handler: a.handleSession,
|
||||
ServerConfigCallback: newSSHServerConfig,
|
||||
// check public key(s)
|
||||
PublicKeyHandler: func(ctx ssh.Context, key ssh.PublicKey) bool {
|
||||
@@ -82,8 +94,20 @@ func (a *Agent) StartServer(opts ServerOptions) error {
|
||||
IdleTimeout: 70 * time.Second,
|
||||
}
|
||||
|
||||
// Start SSH server on the listener
|
||||
return a.server.Serve(ln)
|
||||
a.server = server
|
||||
a.serverListener = ln
|
||||
return server, ln, nil
|
||||
}
|
||||
|
||||
func (a *Agent) serveSSHServer(server *ssh.Server, listener net.Listener) error {
|
||||
err := server.Serve(listener)
|
||||
a.serverMu.Lock()
|
||||
if a.server == server {
|
||||
a.server = nil
|
||||
a.serverListener = nil
|
||||
}
|
||||
a.serverMu.Unlock()
|
||||
return err
|
||||
}
|
||||
|
||||
// newSSHServerConfig returns a separate config for each connection because
|
||||
@@ -115,9 +139,8 @@ func (a *Agent) getHubVersion(sessionCtx ssh.Context) semver.Version {
|
||||
// appropriate encoding format based on hub version, and exits with appropriate
|
||||
// status codes.
|
||||
func (a *Agent) handleSession(s ssh.Session) {
|
||||
a.connectionManager.eventChan <- SSHConnect
|
||||
|
||||
sessionCtx := s.Context()
|
||||
a.connectionManager.sshConnectionOpened(sessionCtx)
|
||||
|
||||
hubVersion := a.getHubVersion(sessionCtx)
|
||||
|
||||
@@ -164,12 +187,13 @@ func (a *Agent) handleSSHRequest(w io.Writer, req *common.HubRequest[cbor.RawMes
|
||||
}
|
||||
|
||||
ctx := &HandlerContext{
|
||||
Client: nil,
|
||||
Agent: a,
|
||||
Request: req,
|
||||
RequestID: nil,
|
||||
HubVerified: true,
|
||||
SendResponse: sshResponder,
|
||||
Client: nil,
|
||||
Agent: a,
|
||||
Request: req,
|
||||
RequestID: nil,
|
||||
HubVerified: true,
|
||||
ConnectionType: system.ConnectionTypeSSH,
|
||||
SendResponse: sshResponder,
|
||||
}
|
||||
|
||||
if handler, ok := a.handlerRegistry.GetHandler(req.Action); ok {
|
||||
@@ -184,7 +208,9 @@ func (a *Agent) handleSSHRequest(w io.Writer, req *common.HubRequest[cbor.RawMes
|
||||
// handleLegacyStats serves the legacy one-shot stats payload for older hubs
|
||||
func (a *Agent) handleLegacyStats(w io.Writer, hubVersion semver.Version) error {
|
||||
stats := a.gatherStats(common.DataRequestOptions{CacheTimeMs: defaultDataCacheTimeMs})
|
||||
return a.writeToSession(w, stats, hubVersion)
|
||||
response := *stats
|
||||
response.Info.ConnectionType = system.ConnectionTypeSSH
|
||||
return a.writeToSession(w, &response, hubVersion)
|
||||
}
|
||||
|
||||
// writeToSession encodes and writes system statistics to the session.
|
||||
@@ -261,12 +287,21 @@ func GetNetwork(addr string) string {
|
||||
// StopServer stops the SSH server if it's running.
|
||||
// It returns an error if the server is not running or if there's an error stopping it.
|
||||
func (a *Agent) StopServer() error {
|
||||
a.serverMu.Lock()
|
||||
if a.server == nil {
|
||||
a.serverMu.Unlock()
|
||||
return errors.New("SSH server not running")
|
||||
}
|
||||
server := a.server
|
||||
listener := a.serverListener
|
||||
a.server = nil
|
||||
a.serverListener = nil
|
||||
a.serverMu.Unlock()
|
||||
|
||||
slog.Info("Stopping SSH server")
|
||||
_ = a.server.Close()
|
||||
a.server = nil
|
||||
if listener != nil {
|
||||
_ = listener.Close()
|
||||
}
|
||||
_ = server.Close()
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -7,7 +7,11 @@ import (
|
||||
"crypto/ed25519"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -21,6 +25,7 @@ import (
|
||||
"github.com/blang/semver"
|
||||
"github.com/fxamacker/cbor/v2"
|
||||
"github.com/gliderlabs/ssh"
|
||||
"github.com/lxzan/gws"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
gossh "golang.org/x/crypto/ssh"
|
||||
@@ -220,6 +225,348 @@ func TestStopServerDoesNotBlockWhenEventQueueFull(t *testing.T) {
|
||||
assert.Equal(t, WebSocketConnect, <-agent.connectionManager.eventChan)
|
||||
}
|
||||
|
||||
func TestSSHConnectionFallbackLifecycle(t *testing.T) {
|
||||
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
|
||||
agent := createTestAgent(t)
|
||||
cm := agent.connectionManager
|
||||
cm.eventChan = make(chan ConnectionEvent, 4)
|
||||
|
||||
_, privateKey, err := ed25519.GenerateKey(nil)
|
||||
require.NoError(t, err)
|
||||
signer, err := gossh.NewSignerFromKey(privateKey)
|
||||
require.NoError(t, err)
|
||||
cm.serverOptions = ServerOptions{
|
||||
Network: "tcp",
|
||||
Addr: "127.0.0.1:0",
|
||||
Keys: []gossh.PublicKey{signer.PublicKey()},
|
||||
}
|
||||
|
||||
// A WebSocket that closed after its upgrade returned nil must start SSH.
|
||||
cm.handleEvent(WebSocketDisconnect)
|
||||
agent.serverMu.Lock()
|
||||
require.NotNil(t, agent.serverListener)
|
||||
addr := agent.serverListener.Addr().String()
|
||||
agent.serverMu.Unlock()
|
||||
defer func() { _ = agent.StopServer() }()
|
||||
|
||||
clientConfig := &gossh.ClientConfig{
|
||||
User: "hub",
|
||||
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
|
||||
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
|
||||
Timeout: 4 * time.Second,
|
||||
}
|
||||
client, err := gossh.Dial("tcp", addr, clientConfig)
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
// A connection is counted when it starts its first session.
|
||||
startSession := func(c *gossh.Client) *gossh.Session {
|
||||
session, err := c.NewSession()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, session.Shell())
|
||||
return session
|
||||
}
|
||||
session := startSession(client)
|
||||
|
||||
select {
|
||||
case <-cm.sshChanged:
|
||||
cm.handleSSHChange()
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("SSH connection did not notify the manager")
|
||||
}
|
||||
require.Equal(t, SSHConnected, cm.getState())
|
||||
|
||||
wsAttempt := make(chan struct{}, 1)
|
||||
releaseWS := make(chan struct{})
|
||||
var releaseOnce sync.Once
|
||||
release := func() { releaseOnce.Do(func() { close(releaseWS) }) }
|
||||
hub := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
wsAttempt <- struct{}{}
|
||||
<-releaseWS
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
}))
|
||||
defer hub.Close()
|
||||
defer release()
|
||||
t.Setenv("BESZEL_AGENT_HUB_URL", hub.URL)
|
||||
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
|
||||
cm.wsClient, err = newWebSocketClient(agent)
|
||||
require.NoError(t, err)
|
||||
|
||||
// A normal short-lived session must not be mistaken for a lost connection,
|
||||
// and further sessions must not count the same connection again.
|
||||
_ = session.Close()
|
||||
_ = startSession(client).Close()
|
||||
select {
|
||||
case <-cm.sshChanged:
|
||||
t.Fatal("session close unexpectedly changed SSH connection state")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
cm.mu.Lock()
|
||||
assert.Equal(t, 1, cm.sshConnections, "sessions should not be counted as connections")
|
||||
cm.mu.Unlock()
|
||||
|
||||
secondClient, err := gossh.Dial("tcp", addr, clientConfig)
|
||||
require.NoError(t, err)
|
||||
defer secondClient.Close()
|
||||
defer startSession(secondClient).Close()
|
||||
require.Eventually(t, func() bool {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
return cm.sshConnections == 2
|
||||
}, 5*time.Second, 10*time.Millisecond, "second SSH connection was not counted")
|
||||
require.NoError(t, client.Close())
|
||||
select {
|
||||
case <-cm.sshChanged:
|
||||
t.Fatal("closing one of two SSH connections changed SSH connection state")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
require.Equal(t, SSHConnected, cm.getState())
|
||||
require.NoError(t, secondClient.Close())
|
||||
select {
|
||||
case <-cm.sshChanged:
|
||||
cm.handleSSHChange()
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("SSH TCP close did not notify the manager")
|
||||
}
|
||||
require.Equal(t, Disconnected, cm.getState())
|
||||
require.NotNil(t, cm.wsTicker)
|
||||
select {
|
||||
case <-wsAttempt:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("agent did not retry WebSocket after SSH disconnected")
|
||||
}
|
||||
// The hub may redial straight away, so the listener must stay open while
|
||||
// the WebSocket attempt is pending and after it fails.
|
||||
requireSameListener := func(msg string) {
|
||||
agent.serverMu.Lock()
|
||||
defer agent.serverMu.Unlock()
|
||||
require.NotNil(t, agent.serverListener, msg)
|
||||
assert.Equal(t, addr, agent.serverListener.Addr().String(), msg)
|
||||
}
|
||||
requireSameListener("SSH listener should stay open during the WebSocket attempt")
|
||||
thirdClient, err := gossh.Dial("tcp", addr, clientConfig)
|
||||
require.NoError(t, err, "SSH should accept a redial during the WebSocket attempt")
|
||||
require.NoError(t, thirdClient.Close())
|
||||
release()
|
||||
require.Eventually(t, func() bool {
|
||||
return !cm.isConnectingNow()
|
||||
}, 5*time.Second, 10*time.Millisecond, "reconnect attempt did not finish")
|
||||
requireSameListener("SSH listener should stay open after the WebSocket attempt fails")
|
||||
cm.stopWsTicker()
|
||||
}
|
||||
|
||||
// offeredKeySigner offers an authorized public key without proving possession
|
||||
// of its private key: Sign blocks until released, then signs with another key.
|
||||
type offeredKeySigner struct {
|
||||
gossh.Signer
|
||||
publicKey gossh.PublicKey
|
||||
signing chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (s *offeredKeySigner) PublicKey() gossh.PublicKey { return s.publicKey }
|
||||
|
||||
func (s *offeredKeySigner) Sign(rand io.Reader, data []byte) (*gossh.Signature, error) {
|
||||
close(s.signing)
|
||||
<-s.release
|
||||
return s.Signer.Sign(rand, data)
|
||||
}
|
||||
|
||||
// The public key handler runs when a key is offered, before the client signs
|
||||
// anything, so it must not be what marks an SSH connection as established.
|
||||
func TestSSHPublicKeyOfferIsNotAConnection(t *testing.T) {
|
||||
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
|
||||
agent := createTestAgent(t)
|
||||
cm := agent.connectionManager
|
||||
cm.eventChan = make(chan ConnectionEvent, 4)
|
||||
|
||||
newSigner := func() gossh.Signer {
|
||||
_, privateKey, err := ed25519.GenerateKey(nil)
|
||||
require.NoError(t, err)
|
||||
signer, err := gossh.NewSignerFromKey(privateKey)
|
||||
require.NoError(t, err)
|
||||
return signer
|
||||
}
|
||||
hubKey := newSigner().PublicKey()
|
||||
cm.serverOptions = ServerOptions{
|
||||
Network: "tcp",
|
||||
Addr: "127.0.0.1:0",
|
||||
Keys: []gossh.PublicKey{hubKey},
|
||||
}
|
||||
|
||||
cm.handleEvent(WebSocketDisconnect)
|
||||
agent.serverMu.Lock()
|
||||
require.NotNil(t, agent.serverListener)
|
||||
addr := agent.serverListener.Addr().String()
|
||||
agent.serverMu.Unlock()
|
||||
defer func() { _ = agent.StopServer() }()
|
||||
|
||||
signer := &offeredKeySigner{
|
||||
Signer: newSigner(),
|
||||
publicKey: hubKey,
|
||||
signing: make(chan struct{}),
|
||||
release: make(chan struct{}),
|
||||
}
|
||||
dialErr := make(chan error, 1)
|
||||
go func() {
|
||||
client, err := gossh.Dial("tcp", addr, &gossh.ClientConfig{
|
||||
User: "hub",
|
||||
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
|
||||
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
|
||||
Timeout: 4 * time.Second,
|
||||
})
|
||||
if client != nil {
|
||||
client.Close()
|
||||
}
|
||||
dialErr <- err
|
||||
}()
|
||||
|
||||
// The server has accepted the offered key and is waiting for a signature.
|
||||
select {
|
||||
case <-signer.signing:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server did not accept the offered public key")
|
||||
}
|
||||
select {
|
||||
case <-cm.sshChanged:
|
||||
t.Fatal("offering a public key changed SSH connection state")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
assert.False(t, cm.hasSSHConnection())
|
||||
assert.Equal(t, Disconnected, cm.getState())
|
||||
|
||||
close(signer.release)
|
||||
require.Error(t, <-dialErr, "a signature from another key must be rejected")
|
||||
assert.False(t, cm.hasSSHConnection())
|
||||
}
|
||||
|
||||
// startSSHFallbackServer starts the fallback SSH server for a disconnected
|
||||
// agent and returns its address and a client config that can authenticate.
|
||||
func startSSHFallbackServer(t *testing.T) (*Agent, string, *gossh.ClientConfig) {
|
||||
t.Helper()
|
||||
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
|
||||
agent := createTestAgent(t)
|
||||
cm := agent.connectionManager
|
||||
cm.eventChan = make(chan ConnectionEvent, 4)
|
||||
|
||||
_, privateKey, err := ed25519.GenerateKey(nil)
|
||||
require.NoError(t, err)
|
||||
signer, err := gossh.NewSignerFromKey(privateKey)
|
||||
require.NoError(t, err)
|
||||
cm.serverOptions = ServerOptions{
|
||||
Network: "tcp",
|
||||
Addr: "127.0.0.1:0",
|
||||
Keys: []gossh.PublicKey{signer.PublicKey()},
|
||||
}
|
||||
|
||||
cm.startSSHServer()
|
||||
agent.serverMu.Lock()
|
||||
require.NotNil(t, agent.serverListener)
|
||||
addr := agent.serverListener.Addr().String()
|
||||
agent.serverMu.Unlock()
|
||||
t.Cleanup(func() { _ = agent.StopServer() })
|
||||
|
||||
return agent, addr, &gossh.ClientConfig{
|
||||
User: "hub",
|
||||
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
|
||||
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
|
||||
Timeout: 4 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
// connectSSHSession dials the agent and starts a session, which is what marks
|
||||
// the connection as established.
|
||||
func connectSSHSession(t *testing.T, addr string, config *gossh.ClientConfig) *gossh.Client {
|
||||
t.Helper()
|
||||
client, err := gossh.Dial("tcp", addr, config)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = client.Close() })
|
||||
session, err := client.NewSession()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, session.Shell())
|
||||
return client
|
||||
}
|
||||
|
||||
// handleNextSSHChange applies the next SSH connection notification, as the
|
||||
// connection manager's event loop would.
|
||||
func handleNextSSHChange(t *testing.T, cm *ConnectionManager) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-cm.sshChanged:
|
||||
cm.handleSSHChange()
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("SSH connection change did not notify the manager")
|
||||
}
|
||||
}
|
||||
|
||||
// An agent without a WebSocket client only has SSH, so losing the hub's SSH
|
||||
// connection must leave the listener in place for it to reconnect.
|
||||
func TestSSHDisconnectKeepsListenerWithoutWebSocket(t *testing.T) {
|
||||
agent, addr, clientConfig := startSSHFallbackServer(t)
|
||||
cm := agent.connectionManager
|
||||
require.Nil(t, cm.wsClient)
|
||||
defer cm.stopWsTicker()
|
||||
|
||||
client := connectSSHSession(t, addr, clientConfig)
|
||||
handleNextSSHChange(t, cm)
|
||||
require.Equal(t, SSHConnected, cm.getState())
|
||||
|
||||
require.NoError(t, client.Close())
|
||||
handleNextSSHChange(t, cm)
|
||||
require.Equal(t, Disconnected, cm.getState())
|
||||
require.Eventually(t, func() bool {
|
||||
return !cm.isConnectingNow()
|
||||
}, 5*time.Second, 10*time.Millisecond, "reconnect attempt did not finish")
|
||||
|
||||
agent.serverMu.Lock()
|
||||
require.NotNil(t, agent.serverListener, "SSH listener should stay open")
|
||||
assert.Equal(t, addr, agent.serverListener.Addr().String(), "SSH listener should not be restarted")
|
||||
agent.serverMu.Unlock()
|
||||
|
||||
connectSSHSession(t, addr, clientConfig)
|
||||
handleNextSSHChange(t, cm)
|
||||
assert.Equal(t, SSHConnected, cm.getState())
|
||||
}
|
||||
|
||||
// A WebSocket attempt that was already in flight can authenticate after SSH has
|
||||
// connected. WebSocket is preferred, so it takes over and SSH is shut down.
|
||||
func TestWebSocketTakesOverFromSSH(t *testing.T) {
|
||||
agent, addr, clientConfig := startSSHFallbackServer(t)
|
||||
cm := agent.connectionManager
|
||||
|
||||
client := connectSSHSession(t, addr, clientConfig)
|
||||
handleNextSSHChange(t, cm)
|
||||
require.Equal(t, SSHConnected, cm.getState())
|
||||
|
||||
cm.wsClient = &WebSocketClient{
|
||||
agent: agent,
|
||||
hubURL: &url.URL{Host: "localhost:8080"},
|
||||
Conn: &gws.Conn{},
|
||||
hubVerified: true,
|
||||
}
|
||||
cm.handleEvent(WebSocketConnect)
|
||||
require.Equal(t, WebSocketConnected, cm.getState())
|
||||
agent.serverMu.Lock()
|
||||
assert.Nil(t, agent.serverListener, "SSH listener should close once WebSocket takes over")
|
||||
agent.serverMu.Unlock()
|
||||
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
_ = client.Wait()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("SSH connection was not closed when WebSocket took over")
|
||||
}
|
||||
|
||||
// The SSH connection closing must not disturb the WebSocket state.
|
||||
handleNextSSHChange(t, cm)
|
||||
assert.False(t, cm.hasSSHConnection())
|
||||
assert.Equal(t, WebSocketConnected, cm.getState())
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////
|
||||
//////////////////// ParseKeys Tests ////////////////////////////
|
||||
/////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -277,7 +277,6 @@ func (a *Agent) getSystemStats(cacheTimeMs uint16) system.Stats {
|
||||
systemStats.WiFi = wifi.Signals(a.systemInfo.WiFi)
|
||||
|
||||
// update system info
|
||||
a.systemInfo.ConnectionType = a.connectionManager.ConnectionType
|
||||
a.systemInfo.Cpu = systemStats.Cpu
|
||||
a.systemInfo.LoadAvg = systemStats.LoadAvg
|
||||
a.systemInfo.MemPct = systemStats.MemPct
|
||||
|
||||
Reference in New Issue
Block a user