Reuse SSH transports for concurrent exec sessions

This commit is contained in:
Fedor Korotkov 2026-05-06 09:16:59 -04:00
parent 3083b541df
commit 2e8a6d7a58
9 changed files with 575 additions and 51 deletions

View File

@ -5,7 +5,6 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"net"
"net/http" "net/http"
"strconv" "strconv"
"time" "time"
@ -190,8 +189,14 @@ func (controller *Controller) newSSHExecSession(
sessionContext, sessionContextCancel := context.WithCancel(context.Background()) sessionContext, sessionContextCancel := context.WithCancel(context.Background())
type sshExecAttempt struct { type sshExecAttempt struct {
portForwardConn net.Conn lease *execSSHTransportLease
exec *sshexec.Exec exec sshExecRunner
}
transportKey := execSSHTransportKey{
workerName: vm.Worker,
vmUID: vm.UID,
restartCount: vm.RestartCount,
} }
attempt, err := retry.NewWithData[sshExecAttempt]( attempt, err := retry.NewWithData[sshExecAttempt](
@ -201,32 +206,49 @@ func (controller *Controller) newSSHExecSession(
retry.Attempts(0), retry.Attempts(0),
retry.LastErrorOnly(true), retry.LastErrorOnly(true),
).Do(func() (sshExecAttempt, error) { ).Do(func() (sshExecAttempt, error) {
portForwardConn, err := controller.portForwardConnection( lease, err := controller.execSSHPool.acquire(waitContext, transportKey, func() (execSSHTransport, error) {
sessionContext, portForwardConn, err := controller.portForwardConnection(
waitContext, context.Background(),
vm.Worker, waitContext,
vm.UID, vm.Worker,
22, vm.UID,
) 22,
)
if err != nil {
return nil, err
}
client, err := sshexec.NewClient(portForwardConn, vm.SSHUsername(), vm.SSHPassword())
if err != nil {
return nil, fmt.Errorf("failed to establish SSH connection to a VM: %w", err)
}
return &execSSHClientTransport{client: client}, nil
})
if err != nil { if err != nil {
return sshExecAttempt{}, err return sshExecAttempt{}, err
} }
exec, err := sshexec.New(portForwardConn, vm.SSHUsername(), vm.SSHPassword(), sshexec.Options{ exec, err := lease.transport().NewExec(sshexec.Options{
Interactive: spec.interactive, Interactive: spec.interactive,
TTY: spec.tty, TTY: spec.tty,
Rows: spec.rows, Rows: spec.rows,
Cols: spec.cols, Cols: spec.cols,
}) })
if err != nil { if err != nil {
_ = portForwardConn.Close() lease.release()
return sshExecAttempt{}, fmt.Errorf("failed to establish SSH connection to a VM: %w", err) err = fmt.Errorf("failed to create SSH session for a VM: %w", err)
if lease.reused {
return sshExecAttempt{}, retry.Unrecoverable(err)
}
return sshExecAttempt{}, err
} }
return sshExecAttempt{ return sshExecAttempt{
portForwardConn: portForwardConn, lease: lease,
exec: exec, exec: exec,
}, nil }, nil
}) })
if err != nil { if err != nil {
@ -242,7 +264,7 @@ func (controller *Controller) newSSHExecSession(
spec, spec,
runCommand, runCommand,
attempt.exec, attempt.exec,
attempt.portForwardConn, attempt.lease.release,
registry, registry,
controller.execSessionExitTTL, controller.execSessionExitTTL,
policy, policy,

View File

@ -66,6 +66,7 @@ type Controller struct {
sshNoClientAuth bool sshNoClientAuth bool
sshServer *sshserver.SSHServer sshServer *sshserver.SSHServer
execSessions *execSessionRegistry execSessions *execSessionRegistry
execSSHPool *execSSHTransportPool
single singleflight.Group single singleflight.Group
@ -80,6 +81,7 @@ func New(opts ...Option) (*Controller, error) {
execSessionExitTTL: 10 * time.Minute, execSessionExitTTL: 10 * time.Minute,
pingInterval: 30 * time.Second, pingInterval: 30 * time.Second,
execSessions: newExecSessionRegistry(), execSessions: newExecSessionRegistry(),
execSSHPool: newExecSSHTransportPool(),
single: singleflight.Group{}, single: singleflight.Group{},
} }
@ -313,6 +315,7 @@ func (controller *Controller) Run(ctx context.Context) error {
<-ctx.Done() <-ctx.Done()
controller.execSessions.closeAll() controller.execSessions.closeAll()
controller.execSSHPool.closeAll()
if err := controller.httpServer.Shutdown(ctx); err != nil { if err := controller.httpServer.Shutdown(ctx); err != nil {
controller.logger.Errorf("failed to cleanly shutdown the HTTP server: %v", err) controller.logger.Errorf("failed to cleanly shutdown the HTTP server: %v", err)

View File

@ -5,7 +5,6 @@ import (
"errors" "errors"
"io" "io"
"maps" "maps"
"net"
"sync" "sync"
"time" "time"
@ -325,14 +324,14 @@ func (subscriber *execSessionSubscriber) close() {
} }
type execSession struct { type execSession struct {
key execSessionKey key execSessionKey
spec execSessionSpec spec execSessionSpec
command string command string
exec sshExecRunner exec sshExecRunner
transport net.Conn release func()
registry *execSessionRegistry registry *execSessionRegistry
exitTTL time.Duration exitTTL time.Duration
policy execSessionPolicy policy execSessionPolicy
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
@ -348,6 +347,7 @@ type execSession struct {
expiryTimer *time.Timer expiryTimer *time.Timer
startOnce sync.Once startOnce sync.Once
closeOnce sync.Once
done chan struct{} done chan struct{}
doneOnce sync.Once doneOnce sync.Once
} }
@ -356,7 +356,7 @@ func newExecSession(
key execSessionKey, key execSessionKey,
command string, command string,
exec sshExecRunner, exec sshExecRunner,
transport net.Conn, release func(),
registry *execSessionRegistry, registry *execSessionRegistry,
exitTTL time.Duration, exitTTL time.Duration,
policy execSessionPolicy, policy execSessionPolicy,
@ -370,7 +370,7 @@ func newExecSession(
execSessionSpec{command: command}, execSessionSpec{command: command},
command, command,
exec, exec,
transport, release,
registry, registry,
exitTTL, exitTTL,
policy, policy,
@ -384,7 +384,7 @@ func newExecSessionWithContextAndSpec(
spec execSessionSpec, spec execSessionSpec,
command string, command string,
exec sshExecRunner, exec sshExecRunner,
transport net.Conn, release func(),
registry *execSessionRegistry, registry *execSessionRegistry,
exitTTL time.Duration, exitTTL time.Duration,
policy execSessionPolicy, policy execSessionPolicy,
@ -398,7 +398,7 @@ func newExecSessionWithContextAndSpec(
spec: spec.clone(), spec: spec.clone(),
command: command, command: command,
exec: exec, exec: exec,
transport: transport, release: release,
registry: registry, registry: registry,
exitTTL: exitTTL, exitTTL: exitTTL,
policy: policy, policy: policy,
@ -574,10 +574,7 @@ func (session *execSession) close() {
closeSubscribers(subscribers) closeSubscribers(subscribers)
session.cancel() session.cancel()
_ = session.exec.Close() session.closeCommandResources()
if session.transport != nil {
_ = session.transport.Close()
}
if session.registry != nil { if session.registry != nil {
session.registry.remove(session.key, session) session.registry.remove(session.key, session)
} }
@ -661,6 +658,8 @@ func (session *execSession) markFinished() {
close(session.done) close(session.done)
}) })
session.closeCommandResources()
if shouldClose { if shouldClose {
session.close() session.close()
} }
@ -693,6 +692,15 @@ func (session *execSession) dropSubscriber(subscriber *execSessionSubscriber) {
session.detachLocked(subscriber) session.detachLocked(subscriber)
} }
func (session *execSession) closeCommandResources() {
session.closeOnce.Do(func() {
_ = session.exec.Close()
if session.release != nil {
session.release()
}
})
}
func cloneExecFrame(frame *execstream.Frame) *execstream.Frame { func cloneExecFrame(frame *execstream.Frame) *execstream.Frame {
if frame == nil { if frame == nil {
return nil return nil

View File

@ -386,3 +386,31 @@ func TestExecSessionFinishKeepsReconnectableSubscriberOpen(t *testing.T) {
require.Equal(t, execstream.FrameTypeNoMoreHistory, noMoreHistory.Type) require.Equal(t, execstream.FrameTypeNoMoreHistory, noMoreHistory.Type)
require.EqualValues(t, 2, noMoreHistory.Watermark) require.EqualValues(t, 2, noMoreHistory.Watermark)
} }
func TestExecSessionFinishReleasesTransportWhileRetainingHistory(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
var releaseCalls atomic.Int32
session.release = func() {
releaseCalls.Add(1)
}
session.recordFrame(&execstream.Frame{Type: execstream.FrameTypeStdout, Data: []byte("out")})
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeExit,
Exit: &execstream.Exit{Code: 0},
})
session.markFinished()
require.False(t, session.closed)
require.EqualValues(t, 1, releaseCalls.Load())
require.EqualValues(t, 1, session.exec.(*fakeExec).closeCalls.Load())
subscriber, err := session.attach()
require.NoError(t, err)
session.sendHistory(subscriber, 0)
require.Equal(t, execstream.FrameTypeStdout, (<-subscriber.frames).Type)
require.Equal(t, execstream.FrameTypeExit, (<-subscriber.frames).Type)
require.Equal(t, execstream.FrameTypeNoMoreHistory, (<-subscriber.frames).Type)
}

View File

@ -0,0 +1,182 @@
package controller
import (
"context"
"sync"
"github.com/cirruslabs/orchard/internal/controller/sshexec"
)
type execSSHTransport interface {
NewExec(options sshexec.Options) (sshExecRunner, error)
Close() error
}
type execSSHClientTransport struct {
client *sshexec.Client
}
func (transport *execSSHClientTransport) NewExec(options sshexec.Options) (sshExecRunner, error) {
return transport.client.NewExec(options)
}
func (transport *execSSHClientTransport) Close() error {
return transport.client.Close()
}
type execSSHTransportKey struct {
workerName string
vmUID string
restartCount uint64
}
type execSSHTransportEntry struct {
key execSSHTransportKey
transport execSSHTransport
refs int
closed bool
}
type execSSHTransportCreation struct {
done chan struct{}
err error
}
type execSSHTransportPool struct {
mu sync.Mutex
entries map[execSSHTransportKey]*execSSHTransportEntry
creating map[execSSHTransportKey]*execSSHTransportCreation
}
func newExecSSHTransportPool() *execSSHTransportPool {
return &execSSHTransportPool{
entries: map[execSSHTransportKey]*execSSHTransportEntry{},
creating: map[execSSHTransportKey]*execSSHTransportCreation{},
}
}
func (pool *execSSHTransportPool) acquire(
ctx context.Context,
key execSSHTransportKey,
create func() (execSSHTransport, error),
) (*execSSHTransportLease, error) {
for {
pool.mu.Lock()
if entry, ok := pool.entries[key]; ok {
entry.refs++
pool.mu.Unlock()
return &execSSHTransportLease{
pool: pool,
entry: entry,
reused: true,
}, nil
}
if creation, ok := pool.creating[key]; ok {
pool.mu.Unlock()
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-creation.done:
if creation.err != nil {
return nil, creation.err
}
}
continue
}
creation := &execSSHTransportCreation{done: make(chan struct{})}
pool.creating[key] = creation
pool.mu.Unlock()
transport, err := create()
pool.mu.Lock()
delete(pool.creating, key)
creation.err = err
var entry *execSSHTransportEntry
if err == nil {
entry = &execSSHTransportEntry{
key: key,
transport: transport,
refs: 1,
}
pool.entries[key] = entry
}
close(creation.done)
pool.mu.Unlock()
if err != nil {
return nil, err
}
return &execSSHTransportLease{
pool: pool,
entry: entry,
}, nil
}
}
func (pool *execSSHTransportPool) release(entry *execSSHTransportEntry) {
var transport execSSHTransport
pool.mu.Lock()
if entry.closed {
pool.mu.Unlock()
return
}
entry.refs--
if entry.refs == 0 {
entry.closed = true
if pool.entries[entry.key] == entry {
delete(pool.entries, entry.key)
}
transport = entry.transport
}
pool.mu.Unlock()
if transport != nil {
_ = transport.Close()
}
}
func (pool *execSSHTransportPool) closeAll() {
pool.mu.Lock()
entries := make([]*execSSHTransportEntry, 0, len(pool.entries))
for key, entry := range pool.entries {
entry.closed = true
delete(pool.entries, key)
entries = append(entries, entry)
}
pool.mu.Unlock()
for _, entry := range entries {
_ = entry.transport.Close()
}
}
type execSSHTransportLease struct {
pool *execSSHTransportPool
entry *execSSHTransportEntry
reused bool
releaseOnce sync.Once
}
func (lease *execSSHTransportLease) transport() execSSHTransport {
return lease.entry.transport
}
func (lease *execSSHTransportLease) release() {
lease.releaseOnce.Do(func() {
lease.pool.release(lease.entry)
})
}

View File

@ -0,0 +1,186 @@
package controller
import (
"context"
"errors"
"sync"
"sync/atomic"
"testing"
"github.com/cirruslabs/orchard/internal/controller/sshexec"
"github.com/stretchr/testify/require"
)
type fakeExecSSHTransport struct {
newExec func(sshexec.Options) (sshExecRunner, error)
closeCalls atomic.Int32
}
func (transport *fakeExecSSHTransport) NewExec(options sshexec.Options) (sshExecRunner, error) {
if transport.newExec != nil {
return transport.newExec(options)
}
return &fakeExec{}, nil
}
func (transport *fakeExecSSHTransport) Close() error {
transport.closeCalls.Add(1)
return nil
}
func TestExecSSHTransportPoolConcurrentAcquireReusesOneTransport(t *testing.T) {
pool := newExecSSHTransportPool()
key := execSSHTransportKey{workerName: "worker", vmUID: "vm", restartCount: 1}
createStarted := make(chan struct{})
releaseCreate := make(chan struct{})
var createCalls atomic.Int32
transport := &fakeExecSSHTransport{}
create := func() (execSSHTransport, error) {
if createCalls.Add(1) == 1 {
close(createStarted)
}
<-releaseCreate
return transport, nil
}
const leasesCount = 16
leases := make([]*execSSHTransportLease, leasesCount)
errCh := make(chan error, leasesCount)
var wg sync.WaitGroup
wg.Add(leasesCount)
for i := range leasesCount {
go func() {
defer wg.Done()
lease, err := pool.acquire(context.Background(), key, create)
if err != nil {
errCh <- err
return
}
leases[i] = lease
}()
}
<-createStarted
close(releaseCreate)
wg.Wait()
close(errCh)
for err := range errCh {
require.NoError(t, err)
}
require.EqualValues(t, 1, createCalls.Load())
for _, lease := range leases {
require.NotNil(t, lease)
require.Same(t, transport, lease.transport())
lease.release()
}
}
func TestExecSSHTransportPoolClosesOnLastRelease(t *testing.T) {
pool := newExecSSHTransportPool()
key := execSSHTransportKey{workerName: "worker", vmUID: "vm", restartCount: 1}
transport := &fakeExecSSHTransport{}
create := func() (execSSHTransport, error) {
return transport, nil
}
firstLease, err := pool.acquire(context.Background(), key, create)
require.NoError(t, err)
secondLease, err := pool.acquire(context.Background(), key, create)
require.NoError(t, err)
firstLease.release()
require.EqualValues(t, 0, transport.closeCalls.Load())
secondLease.release()
require.EqualValues(t, 1, transport.closeCalls.Load())
require.Empty(t, pool.entries)
}
func TestExecSSHTransportPoolSeparatesVMIncarnations(t *testing.T) {
pool := newExecSSHTransportPool()
var createCalls atomic.Int32
create := func() (execSSHTransport, error) {
createCalls.Add(1)
return &fakeExecSSHTransport{}, nil
}
firstLease, err := pool.acquire(context.Background(),
execSSHTransportKey{workerName: "worker", vmUID: "vm", restartCount: 1}, create)
require.NoError(t, err)
secondLease, err := pool.acquire(context.Background(),
execSSHTransportKey{workerName: "worker", vmUID: "vm", restartCount: 2}, create)
require.NoError(t, err)
defer firstLease.release()
defer secondLease.release()
require.EqualValues(t, 2, createCalls.Load())
require.NotSame(t, firstLease.transport(), secondLease.transport())
}
func TestExecSSHTransportPoolKeepsSharedTransportAfterSessionCreationFailure(t *testing.T) {
pool := newExecSSHTransportPool()
key := execSSHTransportKey{workerName: "worker", vmUID: "vm", restartCount: 1}
var createCalls atomic.Int32
transport := &fakeExecSSHTransport{
newExec: func(sshexec.Options) (sshExecRunner, error) {
return nil, errors.New("failed to open channel")
},
}
create := func() (execSSHTransport, error) {
createCalls.Add(1)
return transport, nil
}
activeLease, err := pool.acquire(context.Background(), key, create)
require.NoError(t, err)
failedLease, err := pool.acquire(context.Background(), key, create)
require.NoError(t, err)
require.True(t, failedLease.reused)
_, err = failedLease.transport().NewExec(sshexec.Options{})
require.ErrorContains(t, err, "failed to open channel")
failedLease.release()
require.EqualValues(t, 1, createCalls.Load())
require.EqualValues(t, 0, transport.closeCalls.Load())
require.Len(t, pool.entries, 1)
activeLease.release()
require.EqualValues(t, 1, transport.closeCalls.Load())
}
func TestExecSSHTransportPoolCloseAllClosesActiveTransports(t *testing.T) {
pool := newExecSSHTransportPool()
transport := &fakeExecSSHTransport{}
lease, err := pool.acquire(context.Background(),
execSSHTransportKey{workerName: "worker", vmUID: "vm", restartCount: 1},
func() (execSSHTransport, error) {
return transport, nil
})
require.NoError(t, err)
pool.closeAll()
require.EqualValues(t, 1, transport.closeCalls.Load())
require.Empty(t, pool.entries)
lease.release()
require.EqualValues(t, 1, transport.closeCalls.Load())
}

View File

@ -10,6 +10,7 @@ import (
"slices" "slices"
"sort" "sort"
"strings" "strings"
"sync"
"github.com/cirruslabs/orchard/internal/execstream" "github.com/cirruslabs/orchard/internal/execstream"
"golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh"
@ -28,16 +29,24 @@ type Options struct {
} }
type Exec struct { type Exec struct {
sshClient *ssh.Client
sshSession *ssh.Session sshSession *ssh.Session
stdout io.Reader stdout io.Reader
stderr io.Reader stderr io.Reader
stdin io.WriteCloser stdin io.WriteCloser
stdinReader *io.PipeReader stdinReader *io.PipeReader
tty bool tty bool
closeOwner func() error
} }
func New(netConn net.Conn, user string, password string, options Options) (*Exec, error) { type Client struct {
netConn net.Conn
sshClient *ssh.Client
closeOnce sync.Once
closeErr error
}
func NewClient(netConn net.Conn, user string, password string) (*Client, error) {
// Establish an SSH connection // Establish an SSH connection
sshConn, sshChans, sshReqs, err := ssh.NewClientConn(netConn, "", &ssh.ClientConfig{ sshConn, sshChans, sshReqs, err := ssh.NewClientConn(netConn, "", &ssh.ClientConfig{
HostKeyCallback: func(hostname string, remote net.Addr, key ssh.PublicKey) error { HostKeyCallback: func(hostname string, remote net.Addr, key ssh.PublicKey) error {
@ -49,21 +58,29 @@ func New(netConn net.Conn, user string, password string, options Options) (*Exec
}, },
}) })
if err != nil { if err != nil {
_ = netConn.Close()
return nil, fmt.Errorf("failed to create an SSH connection: %w", err) return nil, fmt.Errorf("failed to create an SSH connection: %w", err)
} }
sshClient := ssh.NewClient(sshConn, sshChans, sshReqs) return &Client{
netConn: netConn,
sshClient: ssh.NewClient(sshConn, sshChans, sshReqs),
}, nil
}
func (client *Client) NewExec(options Options) (*Exec, error) {
if client == nil || client.sshClient == nil {
return nil, errors.New("SSH client is not initialized")
}
// Create a new SSH session // Create a new SSH session
sshSession, err := sshClient.NewSession() sshSession, err := client.sshClient.NewSession()
if err != nil { if err != nil {
_ = sshClient.Close()
return nil, fmt.Errorf("failed to create an SSH session: %w", err) return nil, fmt.Errorf("failed to create an SSH session: %w", err)
} }
exec := &Exec{ exec := &Exec{
sshClient: sshClient,
sshSession: sshSession, sshSession: sshSession,
tty: options.TTY, tty: options.TTY,
} }
@ -83,7 +100,6 @@ func New(netConn net.Conn, user string, password string, options Options) (*Exec
ssh.TerminalModes{}, ssh.TerminalModes{},
); err != nil { ); err != nil {
_ = sshSession.Close() _ = sshSession.Close()
_ = sshClient.Close()
return nil, fmt.Errorf("failed to request PTY for the SSH session: %w", err) return nil, fmt.Errorf("failed to request PTY for the SSH session: %w", err)
} }
@ -92,7 +108,6 @@ func New(netConn net.Conn, user string, password string, options Options) (*Exec
exec.stdout, err = sshSession.StdoutPipe() exec.stdout, err = sshSession.StdoutPipe()
if err != nil { if err != nil {
_ = sshSession.Close() _ = sshSession.Close()
_ = sshClient.Close()
return nil, fmt.Errorf("failed to create standard output pipe "+ return nil, fmt.Errorf("failed to create standard output pipe "+
"for the SSH session: %w", err) "for the SSH session: %w", err)
@ -101,7 +116,6 @@ func New(netConn net.Conn, user string, password string, options Options) (*Exec
exec.stderr, err = sshSession.StderrPipe() exec.stderr, err = sshSession.StderrPipe()
if err != nil { if err != nil {
_ = sshSession.Close() _ = sshSession.Close()
_ = sshClient.Close()
return nil, fmt.Errorf("failed to create standard error pipe "+ return nil, fmt.Errorf("failed to create standard error pipe "+
"for the SSH session: %w", err) "for the SSH session: %w", err)
@ -110,6 +124,45 @@ func New(netConn net.Conn, user string, password string, options Options) (*Exec
return exec, nil return exec, nil
} }
func (client *Client) Close() error {
if client == nil {
return nil
}
client.closeOnce.Do(func() {
if client.sshClient != nil {
client.closeErr = client.sshClient.Close()
if client.closeErr == nil {
return
}
}
if client.netConn != nil {
client.closeErr = errors.Join(client.closeErr, client.netConn.Close())
}
})
return client.closeErr
}
func New(netConn net.Conn, user string, password string, options Options) (*Exec, error) {
client, err := NewClient(netConn, user, password)
if err != nil {
return nil, err
}
exec, err := client.NewExec(options)
if err != nil {
_ = client.Close()
return nil, err
}
exec.closeOwner = client.Close
return exec, nil
}
func (exec *Exec) Stdin() io.WriteCloser { func (exec *Exec) Stdin() io.WriteCloser {
return exec.stdin return exec.stdin
} }
@ -285,11 +338,15 @@ func (exec *Exec) Close() error {
_ = exec.stdinReader.Close() _ = exec.stdinReader.Close()
} }
if err := exec.sshSession.Close(); err != nil { sessionErr := exec.sshSession.Close()
_ = exec.sshClient.Close() if exec.closeOwner != nil {
ownerErr := exec.closeOwner()
if sessionErr != nil {
return sessionErr
}
return err return ownerErr
} }
return exec.sshClient.Close() return sessionErr
} }

View File

@ -17,11 +17,23 @@ type execSSHServer struct {
config *ssh.ServerConfig config *ssh.ServerConfig
rejectFirstConnections atomic.Int32 rejectFirstConnections atomic.Int32
successfulConnections atomic.Int32
acceptedSessions atomic.Int32
releaseSessions <-chan struct{}
done chan struct{}
wg sync.WaitGroup wg sync.WaitGroup
} }
func startExecSSHServer(t *testing.T, rejectFirstConnections int32) *execSSHServer { func startExecSSHServer(t *testing.T, rejectFirstConnections int32) *execSSHServer {
return startExecSSHServerWithSessionGate(t, rejectFirstConnections, nil)
}
func startExecSSHServerWithSessionGate(
t *testing.T,
rejectFirstConnections int32,
releaseSessions <-chan struct{},
) *execSSHServer {
t.Helper() t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0") listener, err := net.Listen("tcp", "127.0.0.1:0")
@ -34,7 +46,9 @@ func startExecSSHServer(t *testing.T, rejectFirstConnections int32) *execSSHServ
require.NoError(t, err) require.NoError(t, err)
server := &execSSHServer{ server := &execSSHServer{
listener: listener, listener: listener,
releaseSessions: releaseSessions,
done: make(chan struct{}),
config: &ssh.ServerConfig{ config: &ssh.ServerConfig{
PasswordCallback: func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) { PasswordCallback: func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) {
if conn.User() != "admin" || string(password) != "admin" { if conn.User() != "admin" || string(password) != "admin" {
@ -52,6 +66,7 @@ func startExecSSHServer(t *testing.T, rejectFirstConnections int32) *execSSHServ
go server.run() go server.run()
t.Cleanup(func() { t.Cleanup(func() {
close(server.done)
require.NoError(t, server.listener.Close()) require.NoError(t, server.listener.Close())
server.wg.Wait() server.wg.Wait()
}) })
@ -94,6 +109,7 @@ func (server *execSSHServer) serve(conn net.Conn) {
if err != nil { if err != nil {
return return
} }
server.successfulConnections.Add(1)
defer serverConn.Close() defer serverConn.Close()
go ssh.DiscardRequests(requests) go ssh.DiscardRequests(requests)
@ -109,17 +125,18 @@ func (server *execSSHServer) serve(conn net.Conn) {
if err != nil { if err != nil {
continue continue
} }
server.acceptedSessions.Add(1)
server.wg.Add(1) server.wg.Add(1)
go func() { go func() {
defer server.wg.Done() defer server.wg.Done()
serveExecSSHSession(channel, requests) server.serveExecSSHSession(channel, requests)
}() }()
} }
} }
func serveExecSSHSession(channel ssh.Channel, requests <-chan *ssh.Request) { func (server *execSSHServer) serveExecSSHSession(channel ssh.Channel, requests <-chan *ssh.Request) {
defer channel.Close() defer channel.Close()
for request := range requests { for request := range requests {
@ -127,6 +144,13 @@ func serveExecSSHSession(channel ssh.Channel, requests <-chan *ssh.Request) {
case "exec": case "exec":
_ = request.Reply(true, nil) _ = request.Reply(true, nil)
_, _ = io.WriteString(channel, "ok") _, _ = io.WriteString(channel, "ok")
if server.releaseSessions != nil {
select {
case <-server.releaseSessions:
case <-server.done:
return
}
}
_, _ = channel.SendRequest("exit-status", false, ssh.Marshal(struct { _, _ = channel.SendRequest("exit-status", false, ssh.Marshal(struct {
Status uint32 Status uint32
}{Status: 0})) }{Status: 0}))

View File

@ -7,6 +7,7 @@ import (
"fmt" "fmt"
"net" "net"
"sync" "sync"
"sync/atomic"
"testing" "testing"
"time" "time"
@ -228,10 +229,14 @@ func TestVMExecScript(t *testing.T) {
} }
func TestVMExecManyConcurrentSessions(t *testing.T) { func TestVMExecManyConcurrentSessions(t *testing.T) {
sshServer := startExecSSHServer(t, 24) releaseSessions := make(chan struct{})
sshServer := startExecSSHServerWithSessionGate(t, 0, releaseSessions)
var vmDialCalls atomic.Int32
devClient, vmName := prepareForSyntheticExec(t, dialer.DialFunc( devClient, vmName := prepareForSyntheticExec(t, dialer.DialFunc(
func(ctx context.Context, network string, addr string) (net.Conn, error) { func(ctx context.Context, network string, addr string) (net.Conn, error) {
vmDialCalls.Add(1)
var netDialer net.Dialer var netDialer net.Dialer
return netDialer.DialContext(ctx, network, sshServer.Addr()) return netDialer.DialContext(ctx, network, sshServer.Addr())
@ -301,12 +306,21 @@ func TestVMExecManyConcurrentSessions(t *testing.T) {
} }
close(start) close(start)
require.Eventually(t, func() bool {
return sshServer.acceptedSessions.Load() == concurrentExecs
}, 30*time.Second, 10*time.Millisecond)
require.EqualValues(t, 1, vmDialCalls.Load())
require.EqualValues(t, 1, sshServer.successfulConnections.Load())
close(releaseSessions)
wg.Wait() wg.Wait()
close(errCh) close(errCh)
for err := range errCh { for err := range errCh {
require.NoError(t, err) require.NoError(t, err)
} }
require.EqualValues(t, concurrentExecs, sshServer.acceptedSessions.Load())
} }
func TestVMExecSessionReconnectHistory(t *testing.T) { func TestVMExecSessionReconnectHistory(t *testing.T) {