Retry VM exec SSH setup under concurrent load

This commit is contained in:
Fedor Korotkov 2026-05-05 14:31:07 -04:00
parent 2667a01cf8
commit 85ece2e859
3 changed files with 310 additions and 27 deletions

View File

@ -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,

View File

@ -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)
}
}
}

View File

@ -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 {