Reuse SSH transports for concurrent exec sessions
This commit is contained in:
parent
3083b541df
commit
2e8a6d7a58
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
@ -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())
|
||||||
|
}
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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}))
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue