diff --git a/internal/controller/api_vms_exec.go b/internal/controller/api_vms_exec.go index cf684d7..f4cd663 100644 --- a/internal/controller/api_vms_exec.go +++ b/internal/controller/api_vms_exec.go @@ -5,7 +5,6 @@ import ( "encoding/json" "errors" "fmt" - "net" "net/http" "strconv" "time" @@ -190,8 +189,14 @@ func (controller *Controller) newSSHExecSession( sessionContext, sessionContextCancel := context.WithCancel(context.Background()) type sshExecAttempt struct { - portForwardConn net.Conn - exec *sshexec.Exec + lease *execSSHTransportLease + exec sshExecRunner + } + + transportKey := execSSHTransportKey{ + workerName: vm.Worker, + vmUID: vm.UID, + restartCount: vm.RestartCount, } attempt, err := retry.NewWithData[sshExecAttempt]( @@ -201,32 +206,49 @@ func (controller *Controller) newSSHExecSession( retry.Attempts(0), retry.LastErrorOnly(true), ).Do(func() (sshExecAttempt, error) { - portForwardConn, err := controller.portForwardConnection( - sessionContext, - waitContext, - vm.Worker, - vm.UID, - 22, - ) + lease, err := controller.execSSHPool.acquire(waitContext, transportKey, func() (execSSHTransport, error) { + portForwardConn, err := controller.portForwardConnection( + context.Background(), + waitContext, + vm.Worker, + 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 { return sshExecAttempt{}, err } - exec, err := sshexec.New(portForwardConn, vm.SSHUsername(), vm.SSHPassword(), sshexec.Options{ + exec, err := lease.transport().NewExec(sshexec.Options{ Interactive: spec.interactive, TTY: spec.tty, Rows: spec.rows, Cols: spec.cols, }) 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{ - portForwardConn: portForwardConn, - exec: exec, + lease: lease, + exec: exec, }, nil }) if err != nil { @@ -242,7 +264,7 @@ func (controller *Controller) newSSHExecSession( spec, runCommand, attempt.exec, - attempt.portForwardConn, + attempt.lease.release, registry, controller.execSessionExitTTL, policy, diff --git a/internal/controller/controller.go b/internal/controller/controller.go index 1521e0d..0cdf04e 100644 --- a/internal/controller/controller.go +++ b/internal/controller/controller.go @@ -66,6 +66,7 @@ type Controller struct { sshNoClientAuth bool sshServer *sshserver.SSHServer execSessions *execSessionRegistry + execSSHPool *execSSHTransportPool single singleflight.Group @@ -80,6 +81,7 @@ func New(opts ...Option) (*Controller, error) { execSessionExitTTL: 10 * time.Minute, pingInterval: 30 * time.Second, execSessions: newExecSessionRegistry(), + execSSHPool: newExecSSHTransportPool(), single: singleflight.Group{}, } @@ -313,6 +315,7 @@ func (controller *Controller) Run(ctx context.Context) error { <-ctx.Done() controller.execSessions.closeAll() + controller.execSSHPool.closeAll() if err := controller.httpServer.Shutdown(ctx); err != nil { controller.logger.Errorf("failed to cleanly shutdown the HTTP server: %v", err) diff --git a/internal/controller/exec_sessions.go b/internal/controller/exec_sessions.go index 13a6212..f55aff8 100644 --- a/internal/controller/exec_sessions.go +++ b/internal/controller/exec_sessions.go @@ -5,7 +5,6 @@ import ( "errors" "io" "maps" - "net" "sync" "time" @@ -325,14 +324,14 @@ func (subscriber *execSessionSubscriber) close() { } type execSession struct { - key execSessionKey - spec execSessionSpec - command string - exec sshExecRunner - transport net.Conn - registry *execSessionRegistry - exitTTL time.Duration - policy execSessionPolicy + key execSessionKey + spec execSessionSpec + command string + exec sshExecRunner + release func() + registry *execSessionRegistry + exitTTL time.Duration + policy execSessionPolicy ctx context.Context cancel context.CancelFunc @@ -348,6 +347,7 @@ type execSession struct { expiryTimer *time.Timer startOnce sync.Once + closeOnce sync.Once done chan struct{} doneOnce sync.Once } @@ -356,7 +356,7 @@ func newExecSession( key execSessionKey, command string, exec sshExecRunner, - transport net.Conn, + release func(), registry *execSessionRegistry, exitTTL time.Duration, policy execSessionPolicy, @@ -370,7 +370,7 @@ func newExecSession( execSessionSpec{command: command}, command, exec, - transport, + release, registry, exitTTL, policy, @@ -384,7 +384,7 @@ func newExecSessionWithContextAndSpec( spec execSessionSpec, command string, exec sshExecRunner, - transport net.Conn, + release func(), registry *execSessionRegistry, exitTTL time.Duration, policy execSessionPolicy, @@ -398,7 +398,7 @@ func newExecSessionWithContextAndSpec( spec: spec.clone(), command: command, exec: exec, - transport: transport, + release: release, registry: registry, exitTTL: exitTTL, policy: policy, @@ -574,10 +574,7 @@ func (session *execSession) close() { closeSubscribers(subscribers) session.cancel() - _ = session.exec.Close() - if session.transport != nil { - _ = session.transport.Close() - } + session.closeCommandResources() if session.registry != nil { session.registry.remove(session.key, session) } @@ -661,6 +658,8 @@ func (session *execSession) markFinished() { close(session.done) }) + session.closeCommandResources() + if shouldClose { session.close() } @@ -693,6 +692,15 @@ func (session *execSession) dropSubscriber(subscriber *execSessionSubscriber) { 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 { if frame == nil { return nil diff --git a/internal/controller/exec_sessions_test.go b/internal/controller/exec_sessions_test.go index 7164635..78aded6 100644 --- a/internal/controller/exec_sessions_test.go +++ b/internal/controller/exec_sessions_test.go @@ -386,3 +386,31 @@ func TestExecSessionFinishKeepsReconnectableSubscriberOpen(t *testing.T) { require.Equal(t, execstream.FrameTypeNoMoreHistory, noMoreHistory.Type) 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) +} diff --git a/internal/controller/exec_ssh_pool.go b/internal/controller/exec_ssh_pool.go new file mode 100644 index 0000000..777e5bb --- /dev/null +++ b/internal/controller/exec_ssh_pool.go @@ -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) + }) +} diff --git a/internal/controller/exec_ssh_pool_test.go b/internal/controller/exec_ssh_pool_test.go new file mode 100644 index 0000000..47cef9c --- /dev/null +++ b/internal/controller/exec_ssh_pool_test.go @@ -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()) +} diff --git a/internal/controller/sshexec/sshexec.go b/internal/controller/sshexec/sshexec.go index 91cc8ab..269fedf 100644 --- a/internal/controller/sshexec/sshexec.go +++ b/internal/controller/sshexec/sshexec.go @@ -10,6 +10,7 @@ import ( "slices" "sort" "strings" + "sync" "github.com/cirruslabs/orchard/internal/execstream" "golang.org/x/crypto/ssh" @@ -28,16 +29,24 @@ type Options struct { } type Exec struct { - sshClient *ssh.Client sshSession *ssh.Session stdout io.Reader stderr io.Reader stdin io.WriteCloser stdinReader *io.PipeReader 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 sshConn, sshChans, sshReqs, err := ssh.NewClientConn(netConn, "", &ssh.ClientConfig{ 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 { + _ = netConn.Close() + 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 - sshSession, err := sshClient.NewSession() + sshSession, err := client.sshClient.NewSession() if err != nil { - _ = sshClient.Close() - return nil, fmt.Errorf("failed to create an SSH session: %w", err) } exec := &Exec{ - sshClient: sshClient, sshSession: sshSession, tty: options.TTY, } @@ -83,7 +100,6 @@ func New(netConn net.Conn, user string, password string, options Options) (*Exec ssh.TerminalModes{}, ); err != nil { _ = sshSession.Close() - _ = sshClient.Close() 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() if err != nil { _ = sshSession.Close() - _ = sshClient.Close() return nil, fmt.Errorf("failed to create standard output pipe "+ "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() if err != nil { _ = sshSession.Close() - _ = sshClient.Close() return nil, fmt.Errorf("failed to create standard error pipe "+ "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 } +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 { return exec.stdin } @@ -285,11 +338,15 @@ func (exec *Exec) Close() error { _ = exec.stdinReader.Close() } - if err := exec.sshSession.Close(); err != nil { - _ = exec.sshClient.Close() + sessionErr := exec.sshSession.Close() + if exec.closeOwner != nil { + ownerErr := exec.closeOwner() + if sessionErr != nil { + return sessionErr + } - return err + return ownerErr } - return exec.sshClient.Close() + return sessionErr } diff --git a/internal/tests/exec_ssh_server_test.go b/internal/tests/exec_ssh_server_test.go index 3ea9782..586713b 100644 --- a/internal/tests/exec_ssh_server_test.go +++ b/internal/tests/exec_ssh_server_test.go @@ -17,11 +17,23 @@ type execSSHServer struct { config *ssh.ServerConfig rejectFirstConnections atomic.Int32 + successfulConnections atomic.Int32 + acceptedSessions atomic.Int32 + releaseSessions <-chan struct{} + done chan struct{} wg sync.WaitGroup } 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() 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) server := &execSSHServer{ - listener: listener, + listener: listener, + releaseSessions: releaseSessions, + done: make(chan struct{}), config: &ssh.ServerConfig{ PasswordCallback: func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) { if conn.User() != "admin" || string(password) != "admin" { @@ -52,6 +66,7 @@ func startExecSSHServer(t *testing.T, rejectFirstConnections int32) *execSSHServ go server.run() t.Cleanup(func() { + close(server.done) require.NoError(t, server.listener.Close()) server.wg.Wait() }) @@ -94,6 +109,7 @@ func (server *execSSHServer) serve(conn net.Conn) { if err != nil { return } + server.successfulConnections.Add(1) defer serverConn.Close() go ssh.DiscardRequests(requests) @@ -109,17 +125,18 @@ func (server *execSSHServer) serve(conn net.Conn) { if err != nil { continue } + server.acceptedSessions.Add(1) server.wg.Add(1) go func() { 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() for request := range requests { @@ -127,6 +144,13 @@ func serveExecSSHSession(channel ssh.Channel, requests <-chan *ssh.Request) { case "exec": _ = request.Reply(true, nil) _, _ = io.WriteString(channel, "ok") + if server.releaseSessions != nil { + select { + case <-server.releaseSessions: + case <-server.done: + return + } + } _, _ = channel.SendRequest("exit-status", false, ssh.Marshal(struct { Status uint32 }{Status: 0})) diff --git a/internal/tests/exec_test.go b/internal/tests/exec_test.go index 31c44a7..5e32052 100644 --- a/internal/tests/exec_test.go +++ b/internal/tests/exec_test.go @@ -7,6 +7,7 @@ import ( "fmt" "net" "sync" + "sync/atomic" "testing" "time" @@ -228,10 +229,14 @@ func TestVMExecScript(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( func(ctx context.Context, network string, addr string) (net.Conn, error) { + vmDialCalls.Add(1) + var netDialer net.Dialer return netDialer.DialContext(ctx, network, sshServer.Addr()) @@ -301,12 +306,21 @@ func TestVMExecManyConcurrentSessions(t *testing.T) { } 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() close(errCh) for err := range errCh { require.NoError(t, err) } + + require.EqualValues(t, concurrentExecs, sshServer.acceptedSessions.Load()) } func TestVMExecSessionReconnectHistory(t *testing.T) {