mirror of
https://github.com/henrygd/beszel.git
synced 2026-09-30 11:46:13 +00:00
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:
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
+90
-170
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user