Files
beszel-ipv6/internal/hub/transport/ssh_dial_test.go
henrygd 1997984325 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.
2026-09-29 11:59:15 -04:00

176 lines
5.1 KiB
Go

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()
}