From 1997984325de5057bee0bd5b309fa534dac5d639 Mon Sep 17 00:00:00 2001 From: henrygd Date: Tue, 29 Sep 2026 11:51:13 -0400 Subject: [PATCH] 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. --- .../hub/systems/network_monitor_ssh_test.go | 10 +- .../hub/systems/network_monitor_sync_test.go | 3 +- internal/hub/systems/network_monitors.go | 2 +- internal/hub/systems/ssh_timeout_test.go | 158 ----------- internal/hub/systems/system.go | 260 ++++++------------ internal/hub/systems/system_manager.go | 2 +- internal/hub/systems/system_zfs.go | 2 +- internal/hub/systems/system_zfs_test.go | 5 +- internal/hub/transport/ssh.go | 117 ++++---- internal/hub/transport/ssh_dial_test.go | 175 ++++++++++++ internal/hub/transport/ssh_test.go | 36 ++- 11 files changed, 376 insertions(+), 394 deletions(-) create mode 100644 internal/hub/transport/ssh_dial_test.go diff --git a/internal/hub/systems/network_monitor_ssh_test.go b/internal/hub/systems/network_monitor_ssh_test.go index c825c37db..89c05e3c0 100644 --- a/internal/hub/systems/network_monitor_ssh_test.go +++ b/internal/hub/systems/network_monitor_ssh_test.go @@ -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") } diff --git a/internal/hub/systems/network_monitor_sync_test.go b/internal/hub/systems/network_monitor_sync_test.go index 4e0c26656..a89077afc 100644 --- a/internal/hub/systems/network_monitor_sync_test.go +++ b/internal/hub/systems/network_monitor_sync_test.go @@ -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) diff --git a/internal/hub/systems/network_monitors.go b/internal/hub/systems/network_monitors.go index ec669ce58..fa40cc6d7 100644 --- a/internal/hub/systems/network_monitors.go +++ b/internal/hub/systems/network_monitors.go @@ -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 diff --git a/internal/hub/systems/ssh_timeout_test.go b/internal/hub/systems/ssh_timeout_test.go index 217ac12dd..2e7e1373c 100644 --- a/internal/hub/systems/ssh_timeout_test.go +++ b/internal/hub/systems/ssh_timeout_test.go @@ -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() -} diff --git a/internal/hub/systems/system.go b/internal/hub/systems/system.go index 55a5878f0..f5888f75b 100644 --- a/internal/hub/systems/system.go +++ b/internal/hub/systems/system.go @@ -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. diff --git a/internal/hub/systems/system_manager.go b/internal/hub/systems/system_manager.go index 7bd27ab71..e00c3924f 100644 --- a/internal/hub/systems/system_manager.go +++ b/internal/hub/systems/system_manager.go @@ -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 { diff --git a/internal/hub/systems/system_zfs.go b/internal/hub/systems/system_zfs.go index 934657117..9d04d444e 100644 --- a/internal/hub/systems/system_zfs.go +++ b/internal/hub/systems/system_zfs.go @@ -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 diff --git a/internal/hub/systems/system_zfs_test.go b/internal/hub/systems/system_zfs_test.go index b035dd80c..07ffa0eee 100644 --- a/internal/hub/systems/system_zfs_test.go +++ b/internal/hub/systems/system_zfs_test.go @@ -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()) } diff --git a/internal/hub/transport/ssh.go b/internal/hub/transport/ssh.go index 1a2491ae7..8d4964684 100644 --- a/internal/hub/transport/ssh.go +++ b/internal/hub/transport/ssh.go @@ -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 { diff --git a/internal/hub/transport/ssh_dial_test.go b/internal/hub/transport/ssh_dial_test.go new file mode 100644 index 000000000..9307d081d --- /dev/null +++ b/internal/hub/transport/ssh_dial_test.go @@ -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() +} diff --git a/internal/hub/transport/ssh_test.go b/internal/hub/transport/ssh_test.go index 98be4c2dd..f6a01feee 100644 --- a/internal/hub/transport/ssh_test.go +++ b/internal/hub/transport/ssh_test.go @@ -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) +}