Retry VM exec SSH setup under concurrent load (#437)
This commit is contained in:
parent
2667a01cf8
commit
3083b541df
|
|
@ -189,19 +189,27 @@ func (controller *Controller) newSSHExecSession(
|
||||||
) (*execSession, error) {
|
) (*execSession, error) {
|
||||||
sessionContext, sessionContextCancel := context.WithCancel(context.Background())
|
sessionContext, sessionContextCancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
portForwardConn, err := retry.NewWithData[net.Conn](
|
type sshExecAttempt struct {
|
||||||
|
portForwardConn net.Conn
|
||||||
|
exec *sshexec.Exec
|
||||||
|
}
|
||||||
|
|
||||||
|
attempt, err := retry.NewWithData[sshExecAttempt](
|
||||||
retry.Context(waitContext),
|
retry.Context(waitContext),
|
||||||
retry.DelayType(retry.FixedDelay),
|
retry.DelayType(retry.FixedDelay),
|
||||||
retry.Delay(time.Second),
|
retry.Delay(time.Second),
|
||||||
retry.Attempts(0),
|
retry.Attempts(0),
|
||||||
retry.LastErrorOnly(true),
|
retry.LastErrorOnly(true),
|
||||||
).Do(func() (net.Conn, error) {
|
).Do(func() (sshExecAttempt, error) {
|
||||||
return controller.portForwardConnection(sessionContext, waitContext, vm.Worker, vm.UID, 22)
|
portForwardConn, err := controller.portForwardConnection(
|
||||||
})
|
sessionContext,
|
||||||
|
waitContext,
|
||||||
|
vm.Worker,
|
||||||
|
vm.UID,
|
||||||
|
22,
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
sessionContextCancel()
|
return sshExecAttempt{}, err
|
||||||
|
|
||||||
return nil, err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
exec, err := sshexec.New(portForwardConn, vm.SSHUsername(), vm.SSHPassword(), sshexec.Options{
|
exec, err := sshexec.New(portForwardConn, vm.SSHUsername(), vm.SSHPassword(), sshexec.Options{
|
||||||
|
|
@ -211,10 +219,20 @@ func (controller *Controller) newSSHExecSession(
|
||||||
Cols: spec.cols,
|
Cols: spec.cols,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
sessionContextCancel()
|
|
||||||
_ = portForwardConn.Close()
|
_ = portForwardConn.Close()
|
||||||
|
|
||||||
return nil, fmt.Errorf("failed to establish SSH connection to a VM: %w", err)
|
return sshExecAttempt{}, fmt.Errorf("failed to establish SSH connection to a VM: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return sshExecAttempt{
|
||||||
|
portForwardConn: portForwardConn,
|
||||||
|
exec: exec,
|
||||||
|
}, nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
sessionContextCancel()
|
||||||
|
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return newExecSessionWithContextAndSpec(
|
return newExecSessionWithContextAndSpec(
|
||||||
|
|
@ -223,8 +241,8 @@ func (controller *Controller) newSSHExecSession(
|
||||||
key,
|
key,
|
||||||
spec,
|
spec,
|
||||||
runCommand,
|
runCommand,
|
||||||
exec,
|
attempt.exec,
|
||||||
portForwardConn,
|
attempt.portForwardConn,
|
||||||
registry,
|
registry,
|
||||||
controller.execSessionExitTTL,
|
controller.execSessionExitTTL,
|
||||||
policy,
|
policy,
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,139 @@
|
||||||
|
package tests_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/ed25519"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/ssh"
|
||||||
|
)
|
||||||
|
|
||||||
|
type execSSHServer struct {
|
||||||
|
listener net.Listener
|
||||||
|
config *ssh.ServerConfig
|
||||||
|
|
||||||
|
rejectFirstConnections atomic.Int32
|
||||||
|
|
||||||
|
wg sync.WaitGroup
|
||||||
|
}
|
||||||
|
|
||||||
|
func startExecSSHServer(t *testing.T, rejectFirstConnections int32) *execSSHServer {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
_, privateKey, err := ed25519.GenerateKey(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
signer, err := ssh.NewSignerFromKey(privateKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
server := &execSSHServer{
|
||||||
|
listener: listener,
|
||||||
|
config: &ssh.ServerConfig{
|
||||||
|
PasswordCallback: func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) {
|
||||||
|
if conn.User() != "admin" || string(password) != "admin" {
|
||||||
|
return nil, ssh.ErrNoAuth
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ssh.Permissions{}, nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
server.rejectFirstConnections.Store(rejectFirstConnections)
|
||||||
|
server.config.AddHostKey(signer)
|
||||||
|
|
||||||
|
server.wg.Add(1)
|
||||||
|
go server.run()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
require.NoError(t, server.listener.Close())
|
||||||
|
server.wg.Wait()
|
||||||
|
})
|
||||||
|
|
||||||
|
return server
|
||||||
|
}
|
||||||
|
|
||||||
|
func (server *execSSHServer) Addr() string {
|
||||||
|
return server.listener.Addr().String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (server *execSSHServer) run() {
|
||||||
|
defer server.wg.Done()
|
||||||
|
|
||||||
|
for {
|
||||||
|
conn, err := server.listener.Accept()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if server.rejectFirstConnections.Add(-1) >= 0 {
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
server.wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer server.wg.Done()
|
||||||
|
|
||||||
|
server.serve(conn)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (server *execSSHServer) serve(conn net.Conn) {
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
serverConn, newChannels, requests, err := ssh.NewServerConn(conn, server.config)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
go ssh.DiscardRequests(requests)
|
||||||
|
|
||||||
|
for newChannel := range newChannels {
|
||||||
|
if newChannel.ChannelType() != "session" {
|
||||||
|
_ = newChannel.Reject(ssh.UnknownChannelType, "unsupported channel type")
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
channel, requests, err := newChannel.Accept()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
server.wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer server.wg.Done()
|
||||||
|
|
||||||
|
serveExecSSHSession(channel, requests)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func serveExecSSHSession(channel ssh.Channel, requests <-chan *ssh.Request) {
|
||||||
|
defer channel.Close()
|
||||||
|
|
||||||
|
for request := range requests {
|
||||||
|
switch request.Type {
|
||||||
|
case "exec":
|
||||||
|
_ = request.Reply(true, nil)
|
||||||
|
_, _ = io.WriteString(channel, "ok")
|
||||||
|
_, _ = channel.SendRequest("exit-status", false, ssh.Marshal(struct {
|
||||||
|
Status uint32
|
||||||
|
}{Status: 0}))
|
||||||
|
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
_ = request.Reply(false, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -4,13 +4,19 @@ import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/cirruslabs/orchard/internal/controller"
|
||||||
|
"github.com/cirruslabs/orchard/internal/dialer"
|
||||||
"github.com/cirruslabs/orchard/internal/execstream"
|
"github.com/cirruslabs/orchard/internal/execstream"
|
||||||
"github.com/cirruslabs/orchard/internal/tests/devcontroller"
|
"github.com/cirruslabs/orchard/internal/tests/devcontroller"
|
||||||
"github.com/cirruslabs/orchard/internal/tests/platformdependent"
|
"github.com/cirruslabs/orchard/internal/tests/platformdependent"
|
||||||
"github.com/cirruslabs/orchard/internal/tests/wait"
|
"github.com/cirruslabs/orchard/internal/tests/wait"
|
||||||
|
"github.com/cirruslabs/orchard/internal/worker"
|
||||||
"github.com/cirruslabs/orchard/pkg/client"
|
"github.com/cirruslabs/orchard/pkg/client"
|
||||||
v1 "github.com/cirruslabs/orchard/pkg/resource/v1"
|
v1 "github.com/cirruslabs/orchard/pkg/resource/v1"
|
||||||
"github.com/coder/websocket"
|
"github.com/coder/websocket"
|
||||||
|
|
@ -221,6 +227,88 @@ func TestVMExecScript(t *testing.T) {
|
||||||
require.Equal(t, websocket.StatusNormalClosure, closeError.Code)
|
require.Equal(t, websocket.StatusNormalClosure, closeError.Code)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestVMExecManyConcurrentSessions(t *testing.T) {
|
||||||
|
sshServer := startExecSSHServer(t, 24)
|
||||||
|
|
||||||
|
devClient, vmName := prepareForSyntheticExec(t, dialer.DialFunc(
|
||||||
|
func(ctx context.Context, network string, addr string) (net.Conn, error) {
|
||||||
|
var netDialer net.Dialer
|
||||||
|
|
||||||
|
return netDialer.DialContext(ctx, network, sshServer.Addr())
|
||||||
|
},
|
||||||
|
))
|
||||||
|
|
||||||
|
const concurrentExecs = 32
|
||||||
|
|
||||||
|
start := make(chan struct{})
|
||||||
|
errCh := make(chan error, concurrentExecs)
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(concurrentExecs)
|
||||||
|
|
||||||
|
for i := range concurrentExecs {
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
|
||||||
|
<-start
|
||||||
|
|
||||||
|
wsConn, err := devClient.VMs().Exec(
|
||||||
|
t.Context(),
|
||||||
|
vmName,
|
||||||
|
fmt.Sprintf("sh -c 'sleep 2; printf exec-%02d'", i),
|
||||||
|
false,
|
||||||
|
30,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
errCh <- fmt.Errorf("exec %d failed to start: %w", i, err)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer wsConn.CloseNow()
|
||||||
|
|
||||||
|
frame, err := readFrameErr(t.Context(), wsConn)
|
||||||
|
if err != nil {
|
||||||
|
errCh <- fmt.Errorf("exec %d failed to read stdout frame: %w", i, err)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if frame.Type != execstream.FrameTypeStdout {
|
||||||
|
errCh <- fmt.Errorf("exec %d produced first frame %q", i, frame.Type)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if got, want := string(frame.Data), "ok"; got != want {
|
||||||
|
errCh <- fmt.Errorf("exec %d produced stdout %q, want %q", i, got, want)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
frame, err = readFrameErr(t.Context(), wsConn)
|
||||||
|
if err != nil {
|
||||||
|
errCh <- fmt.Errorf("exec %d failed to read exit frame: %w", i, err)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if frame.Type != execstream.FrameTypeExit {
|
||||||
|
errCh <- fmt.Errorf("exec %d produced second frame %q", i, frame.Type)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if frame.Exit.Code != 0 {
|
||||||
|
errCh <- fmt.Errorf("exec %d exited with code %d", i, frame.Exit.Code)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
close(start)
|
||||||
|
wg.Wait()
|
||||||
|
close(errCh)
|
||||||
|
|
||||||
|
for err := range errCh {
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestVMExecSessionReconnectHistory(t *testing.T) {
|
func TestVMExecSessionReconnectHistory(t *testing.T) {
|
||||||
devClient, vmName := prepareForExec(t)
|
devClient, vmName := prepareForExec(t)
|
||||||
sessionID := uuid.NewString()
|
sessionID := uuid.NewString()
|
||||||
|
|
@ -411,25 +499,63 @@ func prepareForExec(t *testing.T) (*client.Client, string) {
|
||||||
return devClient, vmName
|
return devClient, vmName
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func prepareForSyntheticExec(t *testing.T, vmDialer dialer.Dialer) (*client.Client, string) {
|
||||||
|
devClient, _, _ := devcontroller.StartIntegrationTestEnvironmentWithAdditionalOpts(t,
|
||||||
|
false, []controller.Option{controller.WithSynthetic()},
|
||||||
|
false, []worker.Option{
|
||||||
|
worker.WithSynthetic(),
|
||||||
|
worker.WithDialer(vmDialer),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
vmName := "test-vm-exec-" + uuid.NewString()
|
||||||
|
|
||||||
|
err := devClient.VMs().Create(t.Context(), platformdependent.VM(vmName))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.True(t, wait.Wait(30*time.Second, func() bool {
|
||||||
|
vm, err := devClient.VMs().Get(t.Context(), vmName)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Logf("Waiting for the synthetic VM to start. Current status: %s", vm.Status)
|
||||||
|
|
||||||
|
return vm.Status == v1.VMStatusRunning
|
||||||
|
}), "failed to start a synthetic VM")
|
||||||
|
|
||||||
|
return devClient, vmName
|
||||||
|
}
|
||||||
|
|
||||||
func readFrame(t *testing.T, wsConn *websocket.Conn) *execstream.Frame {
|
func readFrame(t *testing.T, wsConn *websocket.Conn) *execstream.Frame {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
|
frame, err := readFrameErr(t.Context(), wsConn)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
return frame
|
||||||
|
}
|
||||||
|
|
||||||
|
func readFrameErr(ctx context.Context, wsConn *websocket.Conn) (*execstream.Frame, error) {
|
||||||
var frame execstream.Frame
|
var frame execstream.Frame
|
||||||
|
|
||||||
readCtx, readCtxCancel := context.WithTimeout(t.Context(), 30*time.Second)
|
readCtx, readCtxCancel := context.WithTimeout(ctx, 30*time.Second)
|
||||||
defer readCtxCancel()
|
defer readCtxCancel()
|
||||||
|
|
||||||
messageType, payloadBytes, err := wsConn.Read(readCtx)
|
messageType, payloadBytes, err := wsConn.Read(readCtx)
|
||||||
require.NoError(t, err)
|
if err != nil {
|
||||||
require.Equal(t, websocket.MessageText, messageType)
|
return nil, err
|
||||||
|
}
|
||||||
err = json.Unmarshal(payloadBytes, &frame)
|
if messageType != websocket.MessageText {
|
||||||
require.NoError(t, err)
|
return nil, fmt.Errorf("unexpected websocket message type %q", messageType)
|
||||||
if frame.Type == execstream.FrameTypeError {
|
|
||||||
require.FailNowf(t, "exec stream error", "%s", frame.Error)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return &frame
|
if err := json.Unmarshal(payloadBytes, &frame); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if frame.Type == execstream.FrameTypeError {
|
||||||
|
return nil, fmt.Errorf("exec stream error: %s", frame.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &frame, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func readFramesUntilExit(t *testing.T, wsConn *websocket.Conn) []*execstream.Frame {
|
func readFramesUntilExit(t *testing.T, wsConn *websocket.Conn) []*execstream.Frame {
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue