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

@@ -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 ////////////////////////////
/////////////////////////////////////////////////////////////////