diff --git a/agent/server.go b/agent/server.go index 2c209fa0b..2fde95ed8 100644 --- a/agent/server.go +++ b/agent/server.go @@ -57,21 +57,11 @@ func (a *Agent) StartServer(opts ServerOptions) error { } defer ln.Close() - // base config (limit to allowed algorithms) - config := &gossh.ServerConfig{ - ServerVersion: fmt.Sprintf("SSH-2.0-%s_%s", beszel.AppName, beszel.Version), - } - config.KeyExchanges = common.DefaultKeyExchanges - config.MACs = common.DefaultMACs - config.Ciphers = common.DefaultCiphers - // set default handler ssh.Handle(a.handleSession) a.server = &ssh.Server{ - ServerConfigCallback: func(ctx ssh.Context) *gossh.ServerConfig { - return config - }, + ServerConfigCallback: newSSHServerConfig, // check public key(s) PublicKeyHandler: func(ctx ssh.Context, key ssh.PublicKey) bool { remoteAddr := ctx.RemoteAddr() @@ -96,6 +86,19 @@ func (a *Agent) StartServer(opts ServerOptions) error { return a.server.Serve(ln) } +// newSSHServerConfig returns a separate config for each connection because +// gliderlabs adds host keys and connection-specific callbacks to it. +func newSSHServerConfig(ssh.Context) *gossh.ServerConfig { + return &gossh.ServerConfig{ + Config: gossh.Config{ + KeyExchanges: common.DefaultKeyExchanges, + MACs: common.DefaultMACs, + Ciphers: common.DefaultCiphers, + }, + ServerVersion: fmt.Sprintf("SSH-2.0-%s_%s", beszel.AppName, beszel.Version), + } +} + // getHubVersion extracts the hub version from the SSH client version string // for a given session. Returns a zero version if parsing fails. func (a *Agent) getHubVersion(sessionCtx ssh.Context) semver.Version { diff --git a/agent/server_config_test.go b/agent/server_config_test.go new file mode 100644 index 000000000..acfe96821 --- /dev/null +++ b/agent/server_config_test.go @@ -0,0 +1,111 @@ +//go:build testing + +package agent + +import ( + "crypto/ed25519" + "fmt" + "net" + "sync" + "testing" + "time" + + "github.com/henrygd/beszel" + "github.com/henrygd/beszel/internal/common" + + "github.com/gliderlabs/ssh" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + gossh "golang.org/x/crypto/ssh" +) + +func TestSSHServerConfigConcurrentConnections(t *testing.T) { + _, key, err := ed25519.GenerateKey(nil) + require.NoError(t, err) + signer, err := gossh.NewSignerFromKey(key) + require.NoError(t, err) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + const connections = 2 + configs := make(chan *gossh.ServerConfig, connections) + release := make(chan struct{}) + var releaseOnce sync.Once + unblock := func() { releaseOnce.Do(func() { close(release) }) } + server := &ssh.Server{ + HostSigners: []ssh.Signer{signer}, + ServerConfigCallback: func(ctx ssh.Context) *gossh.ServerConfig { + config := newSSHServerConfig(ctx) + configs <- config + // Both connections must obtain their configuration before either + // lets gliderlabs add host keys and connection-specific callbacks. + <-release + return config + }, + PublicKeyHandler: func(_ ssh.Context, key ssh.PublicKey) bool { + return ssh.KeysEqual(key, signer.PublicKey()) + }, + Handler: func(session ssh.Session) { _ = session.Exit(0) }, + } + served := make(chan error, 1) + go func() { served <- server.Serve(listener) }() + t.Cleanup(func() { + unblock() + _ = listener.Close() + _ = server.Close() + select { + case <-served: + case <-time.After(5 * time.Second): + t.Error("SSH test server did not stop") + } + }) + + results := make(chan error, connections) + for range connections { + go func() { + conn, err := net.DialTimeout("tcp", listener.Addr().String(), 5*time.Second) + if err != nil { + results <- err + return + } + defer conn.Close() + _ = conn.SetDeadline(time.Now().Add(5 * time.Second)) + client, _, _, err := gossh.NewClientConn(conn, listener.Addr().String(), &gossh.ClientConfig{ + User: "test", + Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)}, + HostKeyCallback: gossh.FixedHostKey(signer.PublicKey()), + }) + if err == nil { + err = client.Close() + } + results <- err + }() + } + + var first *gossh.ServerConfig + for range connections { + select { + case config := <-configs: + assert.Equal(t, fmt.Sprintf("SSH-2.0-%s_%s", beszel.AppName, beszel.Version), config.ServerVersion) + assert.Equal(t, common.DefaultKeyExchanges, config.KeyExchanges) + assert.Equal(t, common.DefaultMACs, config.MACs) + assert.Equal(t, common.DefaultCiphers, config.Ciphers) + if first == nil { + first = config + } else { + assert.NotSame(t, first, config, "SSH connections must not share mutable configuration") + } + case <-time.After(5 * time.Second): + t.Fatal("SSH connections did not reach their config callbacks") + } + } + unblock() + for range connections { + select { + case err := <-results: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("SSH handshake did not finish") + } + } +}