fix(hub): honor SSH request cancellation (#2450)

Co-authored-by: hank <hank@henrygd.me>
This commit is contained in:
user01010111
2026-09-29 11:22:17 -04:00
committed by GitHub
co-authored by hank
parent 24881fbfdf
commit 97db8bd199
2 changed files with 410 additions and 43 deletions
+99 -43
View File
@@ -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
+311
View File
@@ -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))
}