Retry VM exec SSH setup under concurrent load
This commit is contained in:
parent
2667a01cf8
commit
85ece2e859
|
|
@ -189,19 +189,27 @@ func (controller *Controller) newSSHExecSession(
|
|||
) (*execSession, error) {
|
||||
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.DelayType(retry.FixedDelay),
|
||||
retry.Delay(time.Second),
|
||||
retry.Attempts(0),
|
||||
retry.LastErrorOnly(true),
|
||||
).Do(func() (net.Conn, error) {
|
||||
return controller.portForwardConnection(sessionContext, waitContext, vm.Worker, vm.UID, 22)
|
||||
})
|
||||
).Do(func() (sshExecAttempt, error) {
|
||||
portForwardConn, err := controller.portForwardConnection(
|
||||
sessionContext,
|
||||
waitContext,
|
||||
vm.Worker,
|
||||
vm.UID,
|
||||
22,
|
||||
)
|
||||
if err != nil {
|
||||
sessionContextCancel()
|
||||
|
||||
return nil, err
|
||||
return sshExecAttempt{}, err
|
||||
}
|
||||
|
||||
exec, err := sshexec.New(portForwardConn, vm.SSHUsername(), vm.SSHPassword(), sshexec.Options{
|
||||
|
|
@ -211,10 +219,20 @@ func (controller *Controller) newSSHExecSession(
|
|||
Cols: spec.cols,
|
||||
})
|
||||
if err != nil {
|
||||
sessionContextCancel()
|
||||
_ = 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(
|
||||
|
|
@ -223,8 +241,8 @@ func (controller *Controller) newSSHExecSession(
|
|||
key,
|
||||
spec,
|
||||
runCommand,
|
||||
exec,
|
||||
portForwardConn,
|
||||
attempt.exec,
|
||||
attempt.portForwardConn,
|
||||
registry,
|
||||
controller.execSessionExitTTL,
|
||||
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"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/cirruslabs/orchard/internal/controller"
|
||||
"github.com/cirruslabs/orchard/internal/dialer"
|
||||
"github.com/cirruslabs/orchard/internal/execstream"
|
||||
"github.com/cirruslabs/orchard/internal/tests/devcontroller"
|
||||
"github.com/cirruslabs/orchard/internal/tests/platformdependent"
|
||||
"github.com/cirruslabs/orchard/internal/tests/wait"
|
||||
"github.com/cirruslabs/orchard/internal/worker"
|
||||
"github.com/cirruslabs/orchard/pkg/client"
|
||||
v1 "github.com/cirruslabs/orchard/pkg/resource/v1"
|
||||
"github.com/coder/websocket"
|
||||
|
|
@ -221,6 +227,88 @@ func TestVMExecScript(t *testing.T) {
|
|||
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) {
|
||||
devClient, vmName := prepareForExec(t)
|
||||
sessionID := uuid.NewString()
|
||||
|
|
@ -411,25 +499,63 @@ func prepareForExec(t *testing.T) (*client.Client, string) {
|
|||
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 {
|
||||
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
|
||||
|
||||
readCtx, readCtxCancel := context.WithTimeout(t.Context(), 30*time.Second)
|
||||
readCtx, readCtxCancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer readCtxCancel()
|
||||
|
||||
messageType, payloadBytes, err := wsConn.Read(readCtx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, websocket.MessageText, messageType)
|
||||
|
||||
err = json.Unmarshal(payloadBytes, &frame)
|
||||
require.NoError(t, err)
|
||||
if frame.Type == execstream.FrameTypeError {
|
||||
require.FailNowf(t, "exec stream error", "%s", frame.Error)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if messageType != websocket.MessageText {
|
||||
return nil, fmt.Errorf("unexpected websocket message type %q", messageType)
|
||||
}
|
||||
|
||||
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 {
|
||||
|
|
|
|||
Loading…
Reference in New Issue