From 97db8bd199b10f90ce854ddc913491d43fda9216 Mon Sep 17 00:00:00 2001 From: user01010111 <12504630+user01010111@users.noreply.github.com> Date: Wed, 30 Sep 2026 04:22:17 +1300 Subject: [PATCH] fix(hub): honor SSH request cancellation (#2450) Co-authored-by: hank --- internal/hub/transport/ssh.go | 142 +++++++++---- internal/hub/transport/ssh_test.go | 311 +++++++++++++++++++++++++++++ 2 files changed, 410 insertions(+), 43 deletions(-) create mode 100644 internal/hub/transport/ssh_test.go diff --git a/internal/hub/transport/ssh.go b/internal/hub/transport/ssh.go index 840c5564c..1a2491ae7 100644 --- a/internal/hub/transport/ssh.go +++ b/internal/hub/transport/ssh.go @@ -7,6 +7,7 @@ import ( "io" "net" "strings" + "sync" "time" "github.com/blang/semver" @@ -17,6 +18,7 @@ import ( // SSHTransport implements Transport over SSH connections. type SSHTransport struct { + mu sync.Mutex client *ssh.Client config *ssh.ClientConfig host string @@ -51,33 +53,57 @@ func NewSSHTransport(cfg SSHTransportConfig) *SSHTransport { // 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). 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) error { - if t.client == nil { - if err := t.connect(); err != nil { - return err - } +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) + if err != nil { + return err } - session, err := t.createSessionWithTimeout(ctx) + // 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) }) + defer func() { + stop() + if err != nil && ctx.Err() != nil { + err = ctx.Err() + } + if isConnectionError(err) { + t.closeClient(client) + } + }() + + session, err := t.createSessionWithTimeout(ctx, client) if err != nil { return err } @@ -121,21 +147,48 @@ func (t *SSHTransport) Request(ctx context.Context, action common.WebSocketActio // IsConnected returns true if the SSH connection is active. func (t *SSHTransport) IsConnected() bool { - return t.client != nil + return t.GetClient() != nil } // Close terminates the SSH connection. func (t *SSHTransport) Close() { - if t.client != nil { - t.client.Close() + t.closeClient(t.GetClient()) +} + +// closeClient removes only the connection owned by the completed request. +func (t *SSHTransport) closeClient(client *ssh.Client) { + t.mu.Lock() + if t.client == client { t.client = nil } + t.mu.Unlock() + if client != nil { + client.Close() + } } -// connect establishes a new SSH connection. -func (t *SSHTransport) connect() error { +// closeOnCancellation stops I/O when ctx is cancelled. The returned function +// waits for any in-progress close so it cannot outlive the operation. +func closeOnCancellation(ctx context.Context, closeConn func()) func() { + done := make(chan struct{}) + stop := context.AfterFunc(ctx, func() { + closeConn() + close(done) + }) + return func() { + if !stop() { + <-done + } + } +} + +// connect reuses the current client or establishes a cancellable SSH connection. +func (t *SSHTransport) connect(ctx context.Context) (*ssh.Client, error) { + if client := t.GetClient(); client != nil { + return client, nil + } if t.config == nil { - return errors.New("SSH config not set") + return nil, errors.New("SSH config not set") } network := "tcp" @@ -146,46 +199,50 @@ func (t *SSHTransport) connect() error { host = net.JoinHostPort(host, t.port) } - client, err := ssh.Dial(network, host, t.config) + dialer := net.Dialer{Timeout: t.config.Timeout} + conn, err := dialer.DialContext(ctx, network, host) if err != nil { - return err + return nil, err + } + stop := closeOnCancellation(ctx, func() { conn.Close() }) + sshConn, chans, reqs, err := ssh.NewClientConn(conn, host, t.config) + stop() + if ctx.Err() != nil { + conn.Close() + return nil, ctx.Err() + } + if err != nil { + conn.Close() + return nil, err + } + client := ssh.NewClient(sshConn, chans, reqs) + + t.mu.Lock() + if existing := t.client; existing != nil { + t.mu.Unlock() + client.Close() + return existing, nil } t.client = client - - // Extract agent version from server version string t.agentVersion, _ = extractAgentVersion(string(client.Conn.ServerVersion())) - return nil + t.mu.Unlock() + return client, nil } -// createSessionWithTimeout creates a new SSH session with a timeout. -func (t *SSHTransport) createSessionWithTimeout(ctx context.Context) (*ssh.Session, error) { - if t.client == nil { - return nil, errors.New("client not initialized") - } - +// createSessionWithTimeout bounds session creation independently of the request. +func (t *SSHTransport) createSessionWithTimeout(ctx context.Context, client *ssh.Client) (*ssh.Session, error) { ctx, cancel := context.WithTimeout(ctx, t.timeout) defer cancel() - - sessionChan := make(chan *ssh.Session, 1) - errChan := make(chan error, 1) - - go func() { - session, err := t.client.NewSession() - if err != nil { - errChan <- err - } else { - sessionChan <- session + stop := closeOnCancellation(ctx, func() { t.closeClient(client) }) + session, err := client.NewSession() + stop() + if ctx.Err() != nil { + if session != nil { + session.Close() } - }() - - select { - case session := <-sessionChan: - return session, nil - case err := <-errChan: - return nil, err - case <-ctx.Done(): - return nil, errors.New("timeout creating session") + return nil, ctx.Err() } + return session, err } // extractAgentVersion extracts the beszel version from SSH server version string. @@ -206,7 +263,6 @@ func (t *SSHTransport) RequestWithRetry(ctx context.Context, action common.WebSo // Check if it's a connection error that warrants a retry if isConnectionError(err) && attempt < retries { - t.Close() continue } return err diff --git a/internal/hub/transport/ssh_test.go b/internal/hub/transport/ssh_test.go new file mode 100644 index 000000000..98be4c2dd --- /dev/null +++ b/internal/hub/transport/ssh_test.go @@ -0,0 +1,311 @@ +package transport + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "io" + "net" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/fxamacker/cbor/v2" + "github.com/henrygd/beszel/internal/common" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" +) + +// newSSHTestTransport starts a loopback SSH server that stalls the first +// connection at the given stage; later connections behave normally so tests +// can verify reconnection. +func newSSHTestTransport(t *testing.T, stage string) (*SSHTransport, <-chan struct{}, <-chan struct{}) { + t.Helper() + _, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + signer, err := ssh.NewSignerFromKey(key) + require.NoError(t, err) + config := &ssh.ServerConfig{NoClientAuth: true, ServerVersion: "SSH-2.0-beszel_0.20.0"} + config.AddHostKey(signer) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + host, port, err := net.SplitHostPort(listener.Addr().String()) + require.NoError(t, err) + transport := NewSSHTransport(SSHTransportConfig{ + Host: host, Port: port, Timeout: 2 * time.Second, + Config: &ssh.ClientConfig{User: "test", HostKeyCallback: ssh.FixedHostKey(signer.PublicKey()), Timeout: time.Second}, + }) + reached, closed := make(chan struct{}), make(chan struct{}) + var once sync.Once + var connections atomic.Int32 + var mu sync.Mutex + conns := map[net.Conn]bool{} + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + for { + conn, err := listener.Accept() + if err != nil { + return + } + mu.Lock() + conns[conn] = true + mu.Unlock() + first := connections.Add(1) == 1 + wg.Add(1) + go func() { + defer wg.Done() + defer conn.Close() + if first { + defer close(closed) + } + if first && stage == "handshake" { + close(reached) + _, _ = io.Copy(io.Discard, conn) + return + } + server, channels, requests, err := ssh.NewServerConn(conn, config) + if err != nil { + return + } + defer server.Close() + go ssh.DiscardRequests(requests) + disconnected := make(chan struct{}) + go func() { _ = server.Wait(); close(disconnected) }() + for channel := range channels { + wg.Add(1) + go func() { + defer wg.Done() + stall := func(at string) bool { + if !first || stage != at { + return false + } + once.Do(func() { close(reached) }) + <-disconnected + return true + } + if stall("session") { + return + } + ch, requests, err := channel.Accept() + if err != nil { + return + } + defer ch.Close() + for request := range requests { + if request.Type != "shell" { + _ = request.Reply(false, nil) + continue + } + if stall("shell") { + return + } + _ = request.Reply(true, nil) + if stall("write") { + return + } + var req common.HubRequest[cbor.RawMessage] + if cbor.NewDecoder(ch).Decode(&req) != nil || stall("response") { + return + } + if stage == "slow-response" { + select { + case <-time.After(250 * time.Millisecond): + case <-disconnected: + return + } + } + data, _ := cbor.Marshal("control response") + if cbor.NewEncoder(ch).Encode(common.AgentResponse{Data: data}) != nil || stall("exit") { + return + } + _, _ = ch.SendRequest("exit-status", false, ssh.Marshal(struct{ Status uint32 }{0})) + return + } + }() + } + }() + } + }() + t.Cleanup(func() { + listener.Close() + transport.Close() + mu.Lock() + for conn := range conns { + conn.Close() + } + mu.Unlock() + done := make(chan struct{}) + go func() { wg.Wait(); close(done) }() + select { + case <-done: + case <-time.After(3 * time.Second): + t.Error("SSH server did not stop") + } + }) + return transport, reached, closed +} + +func TestSSHRequestCancellation(t *testing.T) { + for _, stage := range []string{"handshake", "session", "shell", "write", "response", "exit"} { + for _, cancellation := range []string{"deadline", "cancel"} { + t.Run(stage+"/"+cancellation, func(t *testing.T) { + transport, reached, closed := newSSHTestTransport(t, stage) + var req any + if stage == "write" { + // Larger than the SSH receive window, so an unread write blocks. + req = strings.Repeat("x", 4<<20) + } + ctx, cancel := context.WithCancel(t.Context()) + if cancellation == "deadline" { + cancel() + // Handshake has its own deadline case. Complete it before + // timing the later SSH phases on slower hosts. + if stage != "handshake" { + setupCtx, stopSetup := context.WithTimeout(t.Context(), 2*time.Second) + _, err := transport.connect(setupCtx) + stopSetup() + require.NoError(t, err) + } + ctx, cancel = context.WithTimeout(t.Context(), 250*time.Millisecond) + } + defer cancel() + done := make(chan error, 1) + result := "unchanged" + go func() { done <- transport.RequestWithRetry(ctx, common.GetContainerLogs, req, &result, 1) }() + select { + case <-reached: + case <-time.After(2 * time.Second): + t.Fatal("request did not reach stalled phase") + } + want := context.DeadlineExceeded + if cancellation == "cancel" { + want = context.Canceled + cancel() + } + select { + case err := <-done: + require.ErrorIs(t, err, want) + case <-time.After(2 * time.Second): + t.Fatal("request ignored cancellation") + } + require.Equal(t, "unchanged", result) + require.False(t, transport.IsConnected()) + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("cancelled request left its connection open") + } + require.NoError(t, transport.Request(t.Context(), common.GetContainerLogs, nil, &result)) + require.Equal(t, "control response", result) + }) + } + } +} + +func TestSSHSessionTimeout(t *testing.T) { + transport, _, _ := newSSHTestTransport(t, "session") + transport.timeout = 50 * time.Millisecond + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + var result string + require.ErrorIs(t, transport.Request(ctx, common.GetContainerLogs, nil, &result), context.DeadlineExceeded) + require.NoError(t, ctx.Err(), "the session timeout must fire before the caller's deadline") + require.False(t, transport.IsConnected()) +} + +func TestSSHRequestAlreadyCancelled(t *testing.T) { + transport := NewSSHTransport(SSHTransportConfig{}) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + var result string + require.ErrorIs(t, transport.Request(ctx, common.GetContainerLogs, nil, &result), context.Canceled) +} + +func TestSSHSlowResponse(t *testing.T) { + transport, _, _ := newSSHTestTransport(t, "slow-response") + transport.timeout = 100 * time.Millisecond + ctx, cancel := context.WithTimeout(t.Context(), 3*time.Second) + defer cancel() + var result string + require.NoError(t, transport.Request(ctx, common.GetContainerLogs, nil, &result)) + require.Equal(t, "control response", result, "session timeout must not shorten the request deadline") +} + +func TestSSHConcurrentRequests(t *testing.T) { + transport, _, _ := newSSHTestTransport(t, "") + var wg sync.WaitGroup + for range 10 { + wg.Go(func() { + ctx, cancel := context.WithTimeout(t.Context(), 3*time.Second) + defer cancel() + var result string + if err := transport.Request(ctx, common.GetContainerLogs, nil, &result); err != nil { + t.Error(err) + } else if result != "control response" { + t.Errorf("unexpected response %q", result) + } + }) + } + wg.Wait() + client := transport.GetClient() + require.NotNil(t, client) + ctx, cancel := context.WithCancel(t.Context()) + var result string + require.NoError(t, transport.Request(ctx, common.GetContainerLogs, nil, &result)) + cancel() + // Cancelling a completed request must not close the reused connection. + require.NoError(t, transport.Request(t.Context(), common.GetContainerLogs, nil, &result)) + require.Same(t, client, transport.GetClient()) + // The existing retry contract still replaces an unusable connection. + require.NoError(t, client.Close()) + require.NoError(t, transport.RequestWithRetry(t.Context(), common.GetContainerLogs, nil, &result, 1)) + require.NotSame(t, client, transport.GetClient()) +} + +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) + require.NoError(t, err) + // Keep another session on the shared connection waiting for an exit. + session, err := client.NewSession() + require.NoError(t, err) + require.NoError(t, session.Shell()) + waiting := make(chan error, 1) + go func() { waiting <- session.Wait() }() + requestCtx, cancelRequest := context.WithCancel(ctx) + defer cancelRequest() + done := make(chan error, 1) + var result string + go func() { done <- transport.Request(requestCtx, common.GetContainerLogs, nil, &result) }() + select { + case <-reached: + case <-ctx.Done(): + t.Fatal("request did not reach the response read") + } + cancelRequest() + select { + case err := <-done: + require.ErrorIs(t, err, context.Canceled) + case <-ctx.Done(): + t.Fatal("request ignored cancellation") + } + select { + case err := <-waiting: + require.Error(t, err, "closing the shared client must release other sessions") + case <-ctx.Done(): + t.Fatal("concurrent session remained blocked") + } + require.NoError(t, transport.Request(ctx, common.GetContainerLogs, nil, &result)) + replacement := transport.GetClient() + require.NotSame(t, client, replacement) + // Late cleanup of the old connection must not discard its replacement. + transport.closeClient(client) + require.Same(t, replacement, transport.GetClient()) + require.NoError(t, transport.Request(ctx, common.GetContainerLogs, nil, &result)) +}