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:
@@ -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 ////////////////////////////
|
||||
/////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user