mirror of
https://github.com/henrygd/beszel.git
synced 2026-09-30 19:56:21 +00:00
fix(hub): honor SSH request cancellation (#2450)
Co-authored-by: hank <hank@henrygd.me>
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
Reference in New Issue
Block a user