refactor(hub): make SSHTransport the single owner of each system's SSH connection

The updater and on-demand requests kept separate copies of the SSH client
(sys.client and the transport's client) and synced them after each request.
That allowed a closed client to be reinstalled over a newer one and leaked
connections that were replaced without being closed.

The updater now dials, opens sessions and tears down timed-out connections
through the transport. The dial keeps the TCP keepalive and handshake
deadline, and an OnConnect callback handles the per-connection resets.
The transport is created under a lock and agentVersion is now atomic.
This commit is contained in:
henrygd
2026-09-29 11:51:13 -04:00
parent 97db8bd199
commit 1997984325
11 changed files with 376 additions and 394 deletions

View File

@@ -42,12 +42,14 @@ func TestSSHNetworkMonitorReconnectSync(t *testing.T) {
t.Cleanup(sys.closeSSHConnection)
requests := make(chan monitor.SyncRequest, 10)
var failSync atomic.Bool
var connections atomic.Int32
go func() {
for {
conn, err := listener.Accept()
if err != nil {
return
}
connections.Add(1)
go func() {
server, channels, reqs, err := ssh.NewServerConn(conn, config)
if err != nil {
@@ -138,15 +140,17 @@ func TestSSHNetworkMonitorReconnectSync(t *testing.T) {
require.False(t, sys.monitorsNeedSync.Load())
fetch()
require.Empty(t, requests, "steady-state fetch must not resync")
require.Equal(t, int32(1), connections.Load(), "stats and monitor sync must share one connection")
// Simulate loss of the agent process/connection and its in-memory monitors.
require.NoError(t, sys.client.Load().Close())
require.NoError(t, sys.sshTransport.GetClient().Close())
fetch()
require.ElementsMatch(t, configs, receive().Configs)
require.False(t, sys.monitorsNeedSync.Load())
require.Equal(t, int32(2), connections.Load(), "reconnect must open exactly one new connection")
// Failed replacements are retried on the next successful stats fetch.
require.NoError(t, sys.client.Load().Close())
require.NoError(t, sys.sshTransport.GetClient().Close())
failSync.Store(true)
fetch()
require.ElementsMatch(t, configs, receive().Configs)
@@ -160,7 +164,7 @@ func TestSSHNetworkMonitorReconnectSync(t *testing.T) {
probe.Set("enabled", false)
require.NoError(t, app.SaveNoValidate(probe))
}
require.NoError(t, sys.client.Load().Close())
require.NoError(t, sys.sshTransport.GetClient().Close())
fetch()
require.Empty(t, receive().Configs, "empty replacement must clear stale monitors")
}

View File

@@ -62,7 +62,8 @@ func TestNetworkMonitorSyncSkipsOlderAgents(t *testing.T) {
for _, version := range []string{"0.0.0", "0.18.0", "0.19.0"} {
t.Run(version, func(t *testing.T) {
// No transport: attempting to send any request would fail.
sys := &System{agentVersion: semver.MustParse(version)}
sys := &System{}
sys.setAgentVersion(semver.MustParse(version))
require.NoError(t, sys.SyncNetworkMonitors(nil))
result, err := sys.UpsertNetworkMonitor(monitor.Config{ID: "test"}, true)
require.NoError(t, err)

View File

@@ -64,7 +64,7 @@ func (sys *System) DeleteNetworkMonitor(id string) error {
}
func (sys *System) syncNetworkMonitors(req monitor.SyncRequest) (monitor.SyncResponse, error) {
if sys.agentVersion.LT(beszel.MinVersionNetworkMonitors) {
if sys.getAgentVersion().LT(beszel.MinVersionNetworkMonitors) {
return monitor.SyncResponse{}, nil
}
timeout := 5 * time.Second

View File

@@ -3,18 +3,12 @@
package systems
import (
"crypto/ed25519"
"crypto/rand"
"errors"
"net"
"sync"
"testing"
"testing/synctest"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
)
// TestRunWithTimeout covers the guard added for issue #2041: the per-system SSH
@@ -60,155 +54,3 @@ func TestRunWithTimeout(t *testing.T) {
})
})
}
// closedConn stands in for a connection whose peer has gone away: opening a
// channel fails rather than succeeding, which is what NewSession does on a
// client that closeSSHConnection has already closed.
type closedConn struct{ ssh.Conn }
func (closedConn) OpenChannel(string, []byte) (ssh.Channel, <-chan *ssh.Request, error) {
return nil, nil, errors.New("use of closed network connection")
}
func (closedConn) Close() error { return nil }
// TestCreateSessionDuringClose covers issue #2157: the background SMART fetch
// creates a session while the updater can be tearing the same connection down,
// so session creation must not read the client field after it is cleared.
func TestCreateSessionDuringClose(t *testing.T) {
for range 500 {
sys := &System{ctx: t.Context()}
sys.client.Store(&ssh.Client{Conn: closedConn{}})
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
session, err := sys.createSessionWithTimeout(time.Second)
assert.Nil(t, session)
assert.Error(t, err, "a closed connection must surface an error, not a session")
}()
go func() {
defer wg.Done()
sys.closeSSHConnection()
}()
wg.Wait()
}
}
// TestDialSSHHandshakeTimeout covers a peer that accepts the TCP connection but
// never sends an SSH banner. Without a handshake deadline the dial blocks the
// updater forever (GHSA-h9jh-29rh-w464).
func TestDialSSHHandshakeTimeout(t *testing.T) {
prev := sshHandshakeTimeout
sshHandshakeTimeout = 200 * time.Millisecond
t.Cleanup(func() { sshHandshakeTimeout = prev })
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer ln.Close()
accepted := make(chan net.Conn, 1)
go func() {
conn, err := ln.Accept()
if err == nil {
accepted <- conn
}
}()
config := &ssh.ClientConfig{
User: "u",
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
done := make(chan error, 1)
go func() {
client, err := dialSSHWithKeepAlive("tcp", ln.Addr().String(), config)
if client != nil {
client.Close()
}
done <- err
}()
select {
case err := <-done:
assert.Error(t, err, "a silent peer must fail the handshake")
case <-time.After(5 * time.Second):
t.Fatal("dial blocked on a peer that never sends an SSH banner")
}
// the hub must close its side of the connection
conn := <-accepted
defer conn.Close()
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err = conn.Read(make([]byte, 256))
for err == nil {
_, err = conn.Read(make([]byte, 256))
}
var netErr net.Error
assert.False(t, errors.As(err, &netErr) && netErr.Timeout(), "hub should close the connection, got %v", err)
}
// TestDialSSHClearsHandshakeDeadline ensures the handshake deadline does not
// carry over to the established connection, which is reused for many updates.
func TestDialSSHClearsHandshakeDeadline(t *testing.T) {
prev := sshHandshakeTimeout
sshHandshakeTimeout = 200 * time.Millisecond
t.Cleanup(func() { sshHandshakeTimeout = prev })
_, hostPriv, err := ed25519.GenerateKey(rand.Reader)
require.NoError(t, err)
hostSigner, err := ssh.NewSignerFromKey(hostPriv)
require.NoError(t, err)
_, clientPriv, err := ed25519.GenerateKey(rand.Reader)
require.NoError(t, err)
clientSigner, err := ssh.NewSignerFromKey(clientPriv)
require.NoError(t, err)
serverConfig := &ssh.ServerConfig{
PublicKeyCallback: func(ssh.ConnMetadata, ssh.PublicKey) (*ssh.Permissions, error) { return nil, nil },
}
serverConfig.AddHostKey(hostSigner)
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer ln.Close()
go func() {
conn, err := ln.Accept()
if err != nil {
return
}
defer conn.Close()
_, chans, reqs, err := ssh.NewServerConn(conn, serverConfig)
if err != nil {
return
}
go ssh.DiscardRequests(reqs)
for newChan := range chans {
ch, chReqs, err := newChan.Accept()
if err != nil {
continue
}
go ssh.DiscardRequests(chReqs)
ch.Close()
}
}()
config := &ssh.ClientConfig{
User: "u",
Auth: []ssh.AuthMethod{ssh.PublicKeys(clientSigner)},
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
client, err := dialSSHWithKeepAlive("tcp", ln.Addr().String(), config)
require.NoError(t, err)
defer client.Close()
// wait past the handshake deadline; the connection must still be usable
time.Sleep(3 * sshHandshakeTimeout)
session, err := client.NewSession()
require.NoError(t, err, "connection should outlive the handshake deadline")
session.Close()
}

View File

@@ -8,7 +8,6 @@ import (
"fmt"
"hash/fnv"
"math/rand"
"net"
"strings"
"sync"
"sync/atomic"
@@ -39,25 +38,25 @@ import (
)
type System struct {
Id string `db:"id"`
Host string `db:"host"`
Port string `db:"port"`
Status string `db:"status"` // Use GetStatus/swapStatus after publishing the system.
statusMu sync.RWMutex // Protects Status and exchanges used by alert transitions.
manager *SystemManager // Manager that this system belongs to
client atomic.Pointer[ssh.Client] // SSH client for fetching data
sshTransport *transport.SSHTransport // SSH transport for requests
data *system.CombinedData // system data from agent
ctx context.Context // Context for stopping the updater
cancel context.CancelFunc // Stops and removes system from updater
WsConn *ws.WsConn // Handler for agent WebSocket connection
agentVersion semver.Version // Agent version
updateTicker *time.Ticker // Ticker for updating the system
detailsFetched atomic.Bool // True if static system details have been fetched and saved
smartFetching atomic.Bool // True if SMART devices are currently being fetched
smartInterval time.Duration // Interval for periodic SMART data updates
zfsFetching atomic.Bool // True if ZFS pools are currently being fetched
zfsInterval time.Duration // Interval for periodic ZFS detail data updates
Id string `db:"id"`
Host string `db:"host"`
Port string `db:"port"`
Status string `db:"status"` // Use GetStatus/swapStatus after publishing the system.
statusMu sync.RWMutex // Protects Status and exchanges used by alert transitions.
manager *SystemManager // Manager that this system belongs to
sshMu sync.Mutex // Protects sshTransport creation
sshTransport *transport.SSHTransport // Owns the SSH connection to the agent
data *system.CombinedData // system data from agent
ctx context.Context // Context for stopping the updater
cancel context.CancelFunc // Stops and removes system from updater
WsConn *ws.WsConn // Handler for agent WebSocket connection
agentVersion atomic.Pointer[semver.Version] // Use getAgentVersion/setAgentVersion
updateTicker *time.Ticker // Ticker for updating the system
detailsFetched atomic.Bool // True if static system details have been fetched and saved
smartFetching atomic.Bool // True if SMART devices are currently being fetched
smartInterval time.Duration // Interval for periodic SMART data updates
zfsFetching atomic.Bool // True if ZFS pools are currently being fetched
zfsInterval time.Duration // Interval for periodic ZFS detail data updates
// A fresh connection needs a full monitor configuration sync.
monitorsNeedSync atomic.Bool
@@ -684,19 +683,11 @@ func (sys *System) request(ctx context.Context, action common.WebSocketAction, r
}
// Fall back to SSH if WebSocket fails
if err := sys.ensureSSHTransport(); err != nil {
sshTransport, err := sys.getSSHTransport()
if err != nil {
return err
}
err := sys.sshTransport.RequestWithRetry(ctx, action, req, dest, 1)
// Keep legacy SSH client/version fields in sync for other code paths.
if sys.sshTransport != nil {
client := sys.sshTransport.GetClient()
if previous := sys.client.Swap(client); client != nil && client != previous {
sys.monitorsNeedSync.Store(true)
}
sys.agentVersion = sys.sshTransport.GetAgentVersion()
}
return err
return sshTransport.RequestWithRetry(ctx, action, req, dest, 1)
}
func shouldFallbackToSSH(err error) bool {
@@ -719,27 +710,49 @@ func shouldCloseWebSocket(err error) bool {
return errors.Is(err, gws.ErrConnClosed) || errors.Is(err, transport.ErrWebSocketNotConnected)
}
// ensureSSHTransport ensures the SSH transport is initialized and connected.
func (sys *System) ensureSSHTransport() error {
if sys.sshTransport == nil {
if sys.manager.sshConfig == nil {
if err := sys.manager.createSSHClientConfig(); err != nil {
return err
}
// getSSHTransport returns the system's SSH transport, creating it on first use.
// The transport owns the only SSH connection to the agent; it is shared by the
// updater and on-demand requests and connects lazily.
func (sys *System) getSSHTransport() (*transport.SSHTransport, error) {
sys.sshMu.Lock()
defer sys.sshMu.Unlock()
if sys.sshTransport != nil {
return sys.sshTransport, nil
}
if sys.manager.sshConfig == nil {
if err := sys.manager.createSSHClientConfig(); err != nil {
return nil, err
}
sys.sshTransport = transport.NewSSHTransport(transport.SSHTransportConfig{
Host: sys.Host,
Port: sys.Port,
Config: sys.manager.sshConfig,
Timeout: 4 * time.Second,
})
}
// Sync client state with transport
if client := sys.client.Load(); client != nil {
sys.sshTransport.SetClient(client)
sys.sshTransport.SetAgentVersion(sys.agentVersion)
sys.sshTransport = transport.NewSSHTransport(transport.SSHTransportConfig{
Host: sys.Host,
Port: sys.Port,
Config: sys.manager.sshConfig,
Timeout: sessionTimeout,
OnConnect: sys.onSSHConnect,
})
return sys.sshTransport, nil
}
// onSSHConnect resets per-connection state after a new SSH connection is made.
func (sys *System) onSSHConnect(agentVersion semver.Version) {
sys.setAgentVersion(agentVersion)
sys.monitorsNeedSync.Store(true)
sys.manager.resetFailedSmartFetchState(sys.Id)
sys.manager.resetFailedZfsFetchState(sys.Id)
}
// getAgentVersion returns the connected agent's version, or zero if unknown.
func (sys *System) getAgentVersion() semver.Version {
if v := sys.agentVersion.Load(); v != nil {
return *v
}
return nil
return semver.Version{}
}
// setAgentVersion records the connected agent's version.
func (sys *System) setAgentVersion(v semver.Version) {
sys.agentVersion.Store(&v)
}
// fetchDataFromAgent attempts to fetch data from the agent, prioritizing WebSocket if available.
@@ -828,7 +841,7 @@ func (sys *System) FetchSystemdLogsFromAgent(serviceName string) (string, error)
func (sys *System) FetchSmartDataFromAgent() (smart.SmartDataResponse, error) {
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
if sys.agentVersion.LT(beszel.MinVersionAgentResponse) {
if sys.getAgentVersion().LT(beszel.MinVersionAgentResponse) {
var data map[string]smart.SmartData
err := sys.request(ctx, common.GetSmartData, nil, &data)
return smart.SmartDataResponse{Data: data}, err
@@ -867,7 +880,7 @@ func MakeStableHashId(strings ...string) string {
// fetchDataViaSSH handles fetching data using SSH.
func (sys *System) fetchDataViaSSH(options common.DataRequestOptions) (*system.CombinedData, error) {
data := &system.CombinedData{}
err := sys.runSSHOperation(4*time.Second, 1, func(session *ssh.Session) (bool, error) {
err := sys.runSSHOperation(1, func(session *ssh.Session) (bool, error) {
stdout, err := session.StdoutPipe()
if err != nil {
return false, err
@@ -880,7 +893,7 @@ func (sys *System) fetchDataViaSSH(options common.DataRequestOptions) (*system.C
// reset in case of retry after a partial decode
*data = system.CombinedData{}
if sys.agentVersion.GTE(beszel.MinVersionAgentResponse) && stdinErr == nil {
if sys.getAgentVersion().GTE(beszel.MinVersionAgentResponse) && stdinErr == nil {
req := common.HubRequest[any]{Action: common.GetData, Data: options}
_ = cbor.NewEncoder(stdin).Encode(req)
_ = stdin.Close()
@@ -896,7 +909,7 @@ func (sys *System) fetchDataViaSSH(options common.DataRequestOptions) (*system.C
}
var decodeErr error
if sys.agentVersion.GTE(beszel.MinVersionCbor) {
if sys.getAgentVersion().GTE(beszel.MinVersionCbor) {
decodeErr = cbor.NewDecoder(stdout).Decode(data)
} else {
decodeErr = json.NewDecoder(stdout).Decode(data)
@@ -919,23 +932,31 @@ func (sys *System) fetchDataViaSSH(options common.DataRequestOptions) (*system.C
return data, nil
}
// runSSHOperation establishes an SSH session and executes the provided operation.
// The operation can request a retry by returning true as the first return value.
func (sys *System) runSSHOperation(timeout time.Duration, retries int, operation func(*ssh.Session) (bool, error)) error {
// runSSHOperation opens a session on the system's SSH connection and executes
// the provided operation. The operation can request a retry by returning true
// as the first return value.
func (sys *System) runSSHOperation(retries int, operation func(*ssh.Session) (bool, error)) error {
sshTransport, err := sys.getSSHTransport()
if err != nil {
return err
}
for attempt := 0; attempt <= retries; attempt++ {
if sys.client.Load() == nil || sys.GetStatus() == down {
if err := sys.createSSHClient(); err != nil {
return err
}
// A down system may still hold a dead connection, so always re-dial.
if sys.GetStatus() == down {
sshTransport.Close()
}
client, err := sshTransport.Connect(sys.ctx)
if err != nil {
return err
}
session, err := sys.createSessionWithTimeout(timeout)
session, err := sshTransport.NewSession(sys.ctx, client)
if err != nil {
if attempt >= retries {
return err
}
sys.manager.hub.Logger().Warn("Session closed. Retrying...", "host", sys.Host, "port", sys.Port, "err", err)
sys.closeSSHConnection()
sshTransport.CloseClient(client)
continue
}
@@ -949,14 +970,14 @@ func (sys *System) runSSHOperation(timeout time.Duration, retries int, operation
retry, opErr := runWithTimeout(sshOperationTimeout, func() (bool, error) {
defer session.Close()
return operation(session)
}, sys.closeSSHConnection)
}, func() { sshTransport.CloseClient(client) })
if opErr == nil {
return nil
}
if retry {
sys.closeSSHConnection()
sshTransport.CloseClient(client)
if attempt < retries {
continue
}
@@ -1005,108 +1026,13 @@ func runWithTimeout(timeout time.Duration, op func() (bool, error), onTimeout fu
}
}
// createSSHClient creates a new SSH client for the system
func (s *System) createSSHClient() error {
if s.manager.sshConfig == nil {
if err := s.manager.createSSHClientConfig(); err != nil {
return err
}
}
network := "tcp"
host := s.Host
if strings.HasPrefix(host, "/") {
network = "unix"
} else {
host = net.JoinHostPort(host, s.Port)
}
client, err := dialSSHWithKeepAlive(network, host, s.manager.sshConfig)
s.client.Store(client)
if err != nil {
return err
}
s.agentVersion, _ = extractAgentVersion(string(client.Conn.ServerVersion()))
s.monitorsNeedSync.Store(true)
s.manager.resetFailedSmartFetchState(s.Id)
s.manager.resetFailedZfsFetchState(s.Id)
return nil
}
// sshKeepAliveInterval is the TCP keep-alive idle interval for SSH connections
// to agents. Enabling OS-level keep-alives lets the hub eventually detect a
// dead peer on an otherwise idle connection instead of trusting it forever.
// This is a backstop for genuine network death; an application-level wedge
// (agent process hung while its kernel keeps ACKing) is caught by the
// per-operation timeout in runSSHOperation instead (see issue #2041).
const sshKeepAliveInterval = 30 * time.Second
// sshHandshakeTimeout bounds the SSH handshake after the TCP connection is
// established. ssh.ClientConfig.Timeout only covers the TCP connect, so a peer
// that accepts the connection but never sends an SSH banner would otherwise
// block the updater forever.
var sshHandshakeTimeout = 10 * time.Second
// dialSSHWithKeepAlive dials an SSH connection like ssh.Dial, but enables TCP
// keep-alive on the underlying connection so half-open connections are
// eventually detected by the operating system.
func dialSSHWithKeepAlive(network, addr string, config *ssh.ClientConfig) (*ssh.Client, error) {
dialer := net.Dialer{
Timeout: config.Timeout,
KeepAlive: sshKeepAliveInterval,
}
conn, err := dialer.Dial(network, addr)
if err != nil {
return nil, err
}
_ = conn.SetDeadline(time.Now().Add(sshHandshakeTimeout))
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
if err != nil {
_ = conn.Close()
return nil, err
}
// clear the handshake deadline so it doesn't apply to the long-lived connection
_ = conn.SetDeadline(time.Time{})
return ssh.NewClient(sshConn, chans, reqs), nil
}
// createSessionWithTimeout creates a new SSH session with a timeout to avoid hanging
// in case of network issues
func (sys *System) createSessionWithTimeout(timeout time.Duration) (*ssh.Session, error) {
client := sys.client.Load()
if client == nil {
return nil, fmt.Errorf("client not initialized")
}
ctx, cancel := context.WithTimeout(sys.ctx, timeout)
defer cancel()
sessionChan := make(chan *ssh.Session, 1)
errChan := make(chan error, 1)
go func() {
if session, err := client.NewSession(); err != nil {
errChan <- err
} else {
sessionChan <- session
}
}()
select {
case session := <-sessionChan:
return session, nil
case err := <-errChan:
return nil, err
case <-ctx.Done():
return nil, fmt.Errorf("timeout")
}
}
// closeSSHConnection closes the SSH connection but keeps the system in the manager
func (sys *System) closeSSHConnection() {
if sys.sshTransport != nil {
sys.sshTransport.Close()
}
if client := sys.client.Swap(nil); client != nil {
client.Close()
sys.sshMu.Lock()
sshTransport := sys.sshTransport
sys.sshMu.Unlock()
if sshTransport != nil {
sshTransport.Close()
}
}
@@ -1119,12 +1045,6 @@ func (sys *System) closeWebSocketConnection() {
}
}
// extractAgentVersion extracts the beszel version from SSH server version string
func extractAgentVersion(versionString string) (semver.Version, error) {
_, after, _ := strings.Cut(versionString, "_")
return semver.Parse(after)
}
// getJitter returns a channel that will be triggered after a random delay
// between 51% and 95% of the interval.
// This is used to stagger the initial WebSocket connections to prevent clustering.

View File

@@ -354,7 +354,7 @@ func (sm *SystemManager) AddWebSocketSystem(systemId string, agentVersion semver
system := sm.NewSystem(systemId)
system.WsConn = wsConn
system.agentVersion = agentVersion
system.setAgentVersion(agentVersion)
system.monitorsNeedSync.Store(true)
if err := sm.AddRecord(systemRecord, system); err != nil {

View File

@@ -21,7 +21,7 @@ type zfsFetchState struct {
}
func (sys *System) supportsZfsData() bool {
return sys.agentVersion.GTE(beszel.MinVersionZfsData)
return sys.getAgentVersion().GTE(beszel.MinVersionZfsData)
}
// FetchAndSaveZfsPools fetches ZFS detail data from the agent and saves it to

View File

@@ -16,10 +16,11 @@ import (
)
func TestSupportsZfsData(t *testing.T) {
sys := &System{agentVersion: semver.MustParse("0.18.8")}
sys := &System{}
sys.setAgentVersion(semver.MustParse("0.18.8"))
assert.False(t, sys.supportsZfsData())
sys.agentVersion = semver.MustParse("0.18.9")
sys.setAgentVersion(semver.MustParse("0.18.9"))
assert.True(t, sys.supportsZfsData())
}

View File

@@ -16,24 +16,43 @@ import (
"golang.org/x/crypto/ssh"
)
// SSHTransport implements Transport over SSH connections.
// sshKeepAliveInterval is the TCP keep-alive idle interval for SSH connections
// to agents. Enabling OS-level keep-alives lets the hub eventually detect a
// dead peer on an otherwise idle connection instead of trusting it forever.
// This is a backstop for genuine network death; an application-level wedge
// (agent process hung while its kernel keeps ACKing) is caught by per-operation
// timeouts instead (see issue #2041).
const sshKeepAliveInterval = 30 * time.Second
// sshHandshakeTimeout bounds the SSH handshake after the TCP connection is
// established. ssh.ClientConfig.Timeout only covers the TCP connect, so a peer
// that accepts the connection but never sends an SSH banner would otherwise
// block the caller forever (GHSA-h9jh-29rh-w464).
var sshHandshakeTimeout = 10 * time.Second
// SSHTransport implements Transport over SSH connections. It owns the single
// SSH connection to an agent, which is shared by all requests and sessions.
type SSHTransport struct {
mu sync.Mutex
client *ssh.Client
config *ssh.ClientConfig
host string
port string
agentVersion semver.Version
timeout time.Duration
mu sync.Mutex
client *ssh.Client
config *ssh.ClientConfig
host string
port string
timeout time.Duration
onConnect func(agentVersion semver.Version)
}
// SSHTransportConfig holds configuration for creating an SSH transport.
type SSHTransportConfig struct {
Host string
Port string
Config *ssh.ClientConfig
AgentVersion semver.Version
Timeout time.Duration
Host string
Port string
Config *ssh.ClientConfig
Timeout time.Duration
// OnConnect is called when a new connection is established, before it is
// available to other callers, with the agent version from its SSH server
// version string. It runs under the transport lock and must not call back
// into the transport.
OnConnect func(agentVersion semver.Version)
}
// NewSSHTransport creates a new SSH transport with the given configuration.
@@ -43,48 +62,27 @@ func NewSSHTransport(cfg SSHTransportConfig) *SSHTransport {
timeout = 4 * time.Second
}
return &SSHTransport{
config: cfg.Config,
host: cfg.Host,
port: cfg.Port,
agentVersion: cfg.AgentVersion,
timeout: timeout,
config: cfg.Config,
host: cfg.Host,
port: cfg.Port,
timeout: timeout,
onConnect: cfg.OnConnect,
}
}
// SetClient sets the SSH client for reuse across requests.
func (t *SSHTransport) SetClient(client *ssh.Client) {
t.mu.Lock()
defer t.mu.Unlock()
t.client = client
}
// SetAgentVersion sets the agent version (extracted from SSH handshake).
func (t *SSHTransport) SetAgentVersion(version semver.Version) {
t.mu.Lock()
defer t.mu.Unlock()
t.agentVersion = version
}
// GetClient returns the current SSH client (for connection management).
// GetClient returns the current SSH client, or nil if not connected.
func (t *SSHTransport) GetClient() *ssh.Client {
t.mu.Lock()
defer t.mu.Unlock()
return t.client
}
// GetAgentVersion returns the agent version.
func (t *SSHTransport) GetAgentVersion() semver.Version {
t.mu.Lock()
defer t.mu.Unlock()
return t.agentVersion
}
// Request sends a request to the agent via SSH and unmarshals the response.
func (t *SSHTransport) Request(ctx context.Context, action common.WebSocketAction, req any, dest any) (err error) {
if err := ctx.Err(); err != nil {
return err
}
client, err := t.connect(ctx)
client, err := t.Connect(ctx)
if err != nil {
return err
}
@@ -92,18 +90,18 @@ func (t *SSHTransport) Request(ctx context.Context, action common.WebSocketActio
// Closing only the session still depends on the peer processing SSH packets.
// Close the captured connection to release every blocked read/write, including
// concurrent sessions; subsequent requests can reconnect.
stop := closeOnCancellation(ctx, func() { t.closeClient(client) })
stop := closeOnCancellation(ctx, func() { t.CloseClient(client) })
defer func() {
stop()
if err != nil && ctx.Err() != nil {
err = ctx.Err()
}
if isConnectionError(err) {
t.closeClient(client)
t.CloseClient(client)
}
}()
session, err := t.createSessionWithTimeout(ctx, client)
session, err := t.NewSession(ctx, client)
if err != nil {
return err
}
@@ -152,11 +150,12 @@ func (t *SSHTransport) IsConnected() bool {
// Close terminates the SSH connection.
func (t *SSHTransport) Close() {
t.closeClient(t.GetClient())
t.CloseClient(t.GetClient())
}
// closeClient removes only the connection owned by the completed request.
func (t *SSHTransport) closeClient(client *ssh.Client) {
// CloseClient closes client and clears it if it is still the current
// connection, so a late close never discards a replacement connection.
func (t *SSHTransport) CloseClient(client *ssh.Client) {
t.mu.Lock()
if t.client == client {
t.client = nil
@@ -182,8 +181,9 @@ func closeOnCancellation(ctx context.Context, closeConn func()) func() {
}
}
// connect reuses the current client or establishes a cancellable SSH connection.
func (t *SSHTransport) connect(ctx context.Context) (*ssh.Client, error) {
// Connect returns the current client or establishes a new cancellable SSH
// connection, calling OnConnect when a new connection is stored.
func (t *SSHTransport) Connect(ctx context.Context) (*ssh.Client, error) {
if client := t.GetClient(); client != nil {
return client, nil
}
@@ -199,11 +199,12 @@ func (t *SSHTransport) connect(ctx context.Context) (*ssh.Client, error) {
host = net.JoinHostPort(host, t.port)
}
dialer := net.Dialer{Timeout: t.config.Timeout}
dialer := net.Dialer{Timeout: t.config.Timeout, KeepAlive: sshKeepAliveInterval}
conn, err := dialer.DialContext(ctx, network, host)
if err != nil {
return nil, err
}
_ = conn.SetDeadline(time.Now().Add(sshHandshakeTimeout))
stop := closeOnCancellation(ctx, func() { conn.Close() })
sshConn, chans, reqs, err := ssh.NewClientConn(conn, host, t.config)
stop()
@@ -215,6 +216,8 @@ func (t *SSHTransport) connect(ctx context.Context) (*ssh.Client, error) {
conn.Close()
return nil, err
}
// clear the handshake deadline so it doesn't apply to the long-lived connection
_ = conn.SetDeadline(time.Time{})
client := ssh.NewClient(sshConn, chans, reqs)
t.mu.Lock()
@@ -223,17 +226,23 @@ func (t *SSHTransport) connect(ctx context.Context) (*ssh.Client, error) {
client.Close()
return existing, nil
}
// Initialize per-connection state (e.g. the agent version, which selects the
// protocol) before other callers can reuse the client.
if t.onConnect != nil {
agentVersion, _ := extractAgentVersion(string(client.Conn.ServerVersion()))
t.onConnect(agentVersion)
}
t.client = client
t.agentVersion, _ = extractAgentVersion(string(client.Conn.ServerVersion()))
t.mu.Unlock()
return client, nil
}
// createSessionWithTimeout bounds session creation independently of the request.
func (t *SSHTransport) createSessionWithTimeout(ctx context.Context, client *ssh.Client) (*ssh.Session, error) {
// NewSession opens a session on client, bounded by the transport timeout
// independently of ctx. The connection is closed if session creation stalls.
func (t *SSHTransport) NewSession(ctx context.Context, client *ssh.Client) (*ssh.Session, error) {
ctx, cancel := context.WithTimeout(ctx, t.timeout)
defer cancel()
stop := closeOnCancellation(ctx, func() { t.closeClient(client) })
stop := closeOnCancellation(ctx, func() { t.CloseClient(client) })
session, err := client.NewSession()
stop()
if ctx.Err() != nil {

View File

@@ -0,0 +1,175 @@
package transport
import (
"context"
"crypto/ed25519"
"crypto/rand"
"errors"
"net"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
)
// newDialTestTransport returns a transport that dials ln with config.
func newDialTestTransport(t *testing.T, ln net.Listener, config *ssh.ClientConfig) *SSHTransport {
t.Helper()
host, port, err := net.SplitHostPort(ln.Addr().String())
require.NoError(t, err)
return NewSSHTransport(SSHTransportConfig{Host: host, Port: port, Config: config})
}
// closedConn stands in for a connection whose peer has gone away: opening a
// channel fails rather than succeeding, which is what NewSession does on a
// client that has already been closed.
type closedConn struct{ ssh.Conn }
func (closedConn) OpenChannel(string, []byte) (ssh.Channel, <-chan *ssh.Request, error) {
return nil, nil, errors.New("use of closed network connection")
}
func (closedConn) Close() error { return nil }
// TestNewSessionDuringClose covers issue #2157: a background request creates a
// session while the updater can be tearing the same connection down, so session
// creation must use the captured client rather than the cleared field.
func TestNewSessionDuringClose(t *testing.T) {
for range 500 {
transport := NewSSHTransport(SSHTransportConfig{})
transport.client = &ssh.Client{Conn: closedConn{}}
var wg sync.WaitGroup
wg.Go(func() {
client, err := transport.Connect(t.Context())
if err != nil {
return // already closed; no config to re-dial
}
session, err := transport.NewSession(t.Context(), client)
assert.Nil(t, session)
assert.Error(t, err, "a closed connection must surface an error, not a session")
})
wg.Go(transport.Close)
wg.Wait()
}
}
// TestConnectHandshakeTimeout covers a peer that accepts the TCP connection but
// never sends an SSH banner. Without a handshake deadline the dial blocks the
// caller forever (GHSA-h9jh-29rh-w464).
func TestConnectHandshakeTimeout(t *testing.T) {
prev := sshHandshakeTimeout
sshHandshakeTimeout = 200 * time.Millisecond
t.Cleanup(func() { sshHandshakeTimeout = prev })
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer ln.Close()
accepted := make(chan net.Conn, 1)
go func() {
conn, err := ln.Accept()
if err == nil {
accepted <- conn
}
}()
config := &ssh.ClientConfig{
User: "u",
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
transport := newDialTestTransport(t, ln, config)
done := make(chan error, 1)
go func() {
_, err := transport.Connect(context.Background())
transport.Close()
done <- err
}()
select {
case err := <-done:
assert.Error(t, err, "a silent peer must fail the handshake")
case <-time.After(5 * time.Second):
t.Fatal("dial blocked on a peer that never sends an SSH banner")
}
// the hub must close its side of the connection
conn := <-accepted
defer conn.Close()
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err = conn.Read(make([]byte, 256))
for err == nil {
_, err = conn.Read(make([]byte, 256))
}
var netErr net.Error
assert.False(t, errors.As(err, &netErr) && netErr.Timeout(), "hub should close the connection, got %v", err)
}
// TestConnectClearsHandshakeDeadline ensures the handshake deadline does not
// carry over to the established connection, which is reused for many updates.
func TestConnectClearsHandshakeDeadline(t *testing.T) {
prev := sshHandshakeTimeout
sshHandshakeTimeout = 200 * time.Millisecond
t.Cleanup(func() { sshHandshakeTimeout = prev })
_, hostPriv, err := ed25519.GenerateKey(rand.Reader)
require.NoError(t, err)
hostSigner, err := ssh.NewSignerFromKey(hostPriv)
require.NoError(t, err)
_, clientPriv, err := ed25519.GenerateKey(rand.Reader)
require.NoError(t, err)
clientSigner, err := ssh.NewSignerFromKey(clientPriv)
require.NoError(t, err)
serverConfig := &ssh.ServerConfig{
PublicKeyCallback: func(ssh.ConnMetadata, ssh.PublicKey) (*ssh.Permissions, error) { return nil, nil },
}
serverConfig.AddHostKey(hostSigner)
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer ln.Close()
go func() {
conn, err := ln.Accept()
if err != nil {
return
}
defer conn.Close()
_, chans, reqs, err := ssh.NewServerConn(conn, serverConfig)
if err != nil {
return
}
go ssh.DiscardRequests(reqs)
for newChan := range chans {
ch, chReqs, err := newChan.Accept()
if err != nil {
continue
}
go ssh.DiscardRequests(chReqs)
ch.Close()
}
}()
config := &ssh.ClientConfig{
User: "u",
Auth: []ssh.AuthMethod{ssh.PublicKeys(clientSigner)},
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
transport := newDialTestTransport(t, ln, config)
client, err := transport.Connect(context.Background())
require.NoError(t, err)
defer transport.Close()
// wait past the handshake deadline; the connection must still be usable
time.Sleep(3 * sshHandshakeTimeout)
session, err := client.NewSession()
require.NoError(t, err, "connection should outlive the handshake deadline")
session.Close()
}

View File

@@ -12,6 +12,7 @@ import (
"testing"
"time"
"github.com/blang/semver"
"github.com/fxamacker/cbor/v2"
"github.com/henrygd/beszel/internal/common"
"github.com/stretchr/testify/require"
@@ -166,7 +167,7 @@ func TestSSHRequestCancellation(t *testing.T) {
// timing the later SSH phases on slower hosts.
if stage != "handshake" {
setupCtx, stopSetup := context.WithTimeout(t.Context(), 2*time.Second)
_, err := transport.connect(setupCtx)
_, err := transport.Connect(setupCtx)
stopSetup()
require.NoError(t, err)
}
@@ -270,7 +271,7 @@ func TestSSHCancelledSharedConnection(t *testing.T) {
transport, reached, _ := newSSHTestTransport(t, "response")
ctx, cancel := context.WithTimeout(t.Context(), 3*time.Second)
defer cancel()
client, err := transport.connect(ctx)
client, err := transport.Connect(ctx)
require.NoError(t, err)
// Keep another session on the shared connection waiting for an exit.
session, err := client.NewSession()
@@ -305,7 +306,36 @@ func TestSSHCancelledSharedConnection(t *testing.T) {
replacement := transport.GetClient()
require.NotSame(t, client, replacement)
// Late cleanup of the old connection must not discard its replacement.
transport.closeClient(client)
transport.CloseClient(client)
require.Same(t, replacement, transport.GetClient())
require.NoError(t, transport.Request(ctx, common.GetContainerLogs, nil, &result))
}
func TestConnectInitializesBeforePublishing(t *testing.T) {
transport, _, _ := newSSHTestTransport(t, "")
entered, release := make(chan struct{}), make(chan struct{})
transport.onConnect = func(semver.Version) {
close(entered)
<-release
}
connected := make(chan error, 1)
go func() {
_, err := transport.Connect(t.Context())
connected <- err
}()
select {
case <-entered:
case <-time.After(2 * time.Second):
t.Fatal("OnConnect was not called")
}
reused := make(chan *ssh.Client, 1)
go func() { reused <- transport.GetClient() }()
select {
case <-reused:
t.Fatal("client was exposed before OnConnect completed")
case <-time.After(100 * time.Millisecond):
}
close(release)
require.NoError(t, <-connected)
require.NotNil(t, <-reused)
}