/exec API: drop session reconnection support (#487)

This commit is contained in:
edi-oai
2026-09-09 08:05:08 +01:00
committed by GitHub
parent 9ac6ca6326
commit 4cd43d2b5a
14 changed files with 106 additions and 1880 deletions
+72 -223
View File
@@ -26,18 +26,13 @@ func (controller *Controller) execVM(ctx *gin.Context) responder.Responder {
// Retrieve and parse path and query parameters
name := ctx.Param("name")
sessionID := ctx.Query("session")
if sessionID == "" {
sessionID = ctx.Query("cmux_session_id")
}
command := ctx.Query("command")
if sessionID == "" && command == "" {
if command == "" {
return responder.JSON(http.StatusBadRequest,
NewErrorResponse("\"command\" parameter cannot be empty"))
}
spec, runCommand, err := parseExecSessionSpec(ctx, command)
options, runCommand, err := parseExecOptions(ctx, command)
if err != nil {
return responder.JSON(http.StatusBadRequest, NewErrorResponse("%v", err))
}
@@ -48,20 +43,6 @@ func (controller *Controller) execVM(ctx *gin.Context) responder.Responder {
return responder.Code(http.StatusBadRequest)
}
if sessionID != "" {
return controller.execVMReconnectable(ctx, name, sessionID, spec, runCommand, wait)
}
return controller.execVMLegacy(ctx, name, spec, runCommand, wait)
}
func (controller *Controller) execVMLegacy(
ctx *gin.Context,
name string,
spec execSessionSpec,
runCommand string,
wait uint64,
) responder.Responder {
// Look-up the VM
waitContext, waitContextCancel := context.WithTimeout(ctx, time.Duration(wait)*time.Second)
defer waitContextCancel()
@@ -71,16 +52,7 @@ func (controller *Controller) execVMLegacy(
return responderImpl
}
session, err := controller.newSSHExecSession(
ctx,
waitContext,
vm,
execSessionKey{vmName: name},
spec,
runCommand,
nil,
legacyExecSessionPolicy,
)
exec, err := controller.newSSHExec(waitContext, vm, options)
if err != nil {
return responder.JSON(http.StatusServiceUnavailable, NewErrorResponse("%v", err))
}
@@ -90,7 +62,7 @@ func (controller *Controller) execVMLegacy(
OriginPatterns: []string{"*"},
})
if err != nil {
session.closeIfUnused()
_ = exec.Close()
return responder.Error(err)
}
@@ -102,109 +74,22 @@ func (controller *Controller) execVMLegacy(
_ = wsConn.CloseNow()
}()
return controller.serveExecSession(ctx, wsConn, session)
return controller.serveExec(ctx, wsConn, exec, runCommand)
}
func (controller *Controller) execVMReconnectable(
ctx *gin.Context,
name string,
sessionID string,
spec execSessionSpec,
runCommand string,
wait uint64,
) responder.Responder {
key := execSessionKey{
vmName: name,
sessionID: sessionID,
}
session, ok := controller.execSessions.get(key)
if ok {
if !session.specMatches(spec) {
return responder.JSON(http.StatusConflict,
NewErrorResponse("exec session %q is already running with different options", sessionID))
}
} else {
if spec.command == "" {
return responder.JSON(http.StatusNotFound,
NewErrorResponse("exec session %q does not exist", sessionID))
}
waitContext, waitContextCancel := context.WithTimeout(ctx, time.Duration(wait)*time.Second)
defer waitContextCancel()
vm, responderImpl := controller.waitForVM(waitContext, name)
if responderImpl != nil {
return responderImpl
}
var err error
session, _, err = controller.execSessions.getOrCreate(waitContext, key, func() (*execSession, error) {
return controller.newSSHExecSession(
ctx,
waitContext,
vm,
key,
spec,
runCommand,
controller.execSessions,
reconnectableExecSessionPolicy,
)
})
if err != nil {
return responder.JSON(http.StatusServiceUnavailable, NewErrorResponse("%v", err))
}
if !session.specMatches(spec) {
return responder.JSON(http.StatusConflict,
NewErrorResponse("exec session %q is already running with different options", sessionID))
}
}
wsConn, err := websocket.Accept(ctx.Writer, ctx.Request, &websocket.AcceptOptions{
OriginPatterns: []string{"*"},
})
if err != nil {
session.closeIfUnused()
return responder.Error(err)
}
defer func() {
_ = wsConn.CloseNow()
}()
return controller.serveExecSession(ctx, wsConn, session)
}
func (controller *Controller) newSSHExecSession(
_ *gin.Context,
func (controller *Controller) newSSHExec(
waitContext context.Context,
vm *v1.VM,
key execSessionKey,
spec execSessionSpec,
runCommand string,
registry *execSessionRegistry,
policy execSessionPolicy,
) (*execSession, error) {
sessionContext, sessionContextCancel := context.WithCancel(context.Background())
type sshExecAttempt struct {
exec *sshexec.Exec
}
attempt, err := retry.NewWithData[sshExecAttempt](
options sshexec.Options,
) (*sshexec.Exec, error) {
return retry.NewWithData[*sshexec.Exec](
retry.Context(waitContext),
retry.DelayType(retry.FixedDelay),
retry.Delay(time.Second),
retry.Attempts(0),
retry.LastErrorOnly(true),
).Do(func() (sshExecAttempt, error) {
exec, err := controller.execSSHClients.newExec(vm.UID, sshexec.Options{
Interactive: spec.interactive,
TTY: spec.tty,
Rows: spec.rows,
Cols: spec.cols,
}, func() (sshExecClient, error) {
).Do(func() (*sshexec.Exec, error) {
exec, err := controller.execSSHClients.newExec(vm.UID, options, func() (sshExecClient, error) {
portForwardConn, err := controller.portForwardConnection(
context.Background(),
waitContext,
@@ -228,64 +113,54 @@ func (controller *Controller) newSSHExecSession(
return client, nil
})
if err != nil {
return sshExecAttempt{}, fmt.Errorf("failed to establish SSH connection to a VM: %w", err)
return nil, fmt.Errorf("failed to establish SSH connection to a VM: %w", err)
}
return sshExecAttempt{
exec: exec,
}, nil
return exec, nil
})
if err != nil {
sessionContextCancel()
return nil, err
}
return newExecSessionWithContextAndSpec(
sessionContext,
sessionContextCancel,
key,
spec,
runCommand,
attempt.exec,
nil,
registry,
controller.execSessionRetentionTTL,
policy,
), nil
}
func (controller *Controller) serveExecSession(
func (controller *Controller) serveExec(
ctx *gin.Context,
wsConn *websocket.Conn,
session *execSession,
) responder.Responder {
subscriber, err := session.attach()
if err != nil {
_ = wsConn.Close(websocket.StatusNormalClosure, err.Error())
exec *sshexec.Exec,
command string,
) *responder.EmptyResponder {
execContext, cancel := context.WithCancel(context.Background())
defer func() {
cancel()
_ = exec.Close()
}()
return responder.Empty()
}
defer session.detach(subscriber)
session.start()
// A bounded channel applies backpressure without dropping output. The SSH
// runner drains stdout/stderr before sending exit, then we drain this channel.
outgoingFrames := make(chan *execstream.Frame, 128)
go func() {
defer close(outgoingFrames)
if err := exec.Run(execContext, command, outgoingFrames); err != nil && !errors.Is(err, context.Canceled) {
select {
case outgoingFrames <- &execstream.Frame{Type: execstream.FrameTypeError, Error: err.Error()}:
case <-execContext.Done():
}
}
}()
readFramesErrCh := make(chan error, 1)
go func() {
readFramesErrCh <- controller.readExecSessionFrames(ctx, wsConn, session, subscriber)
readFramesErrCh <- readExecFrames(ctx, wsConn, exec)
}()
for {
select {
case readFramesErr := <-readFramesErrCh:
if readFramesErr != nil &&
!errors.Is(readFramesErr, errExecSessionDetached) &&
!errors.Is(readFramesErr, errExecSessionClosed) {
if readFramesErr != nil {
controller.logger.Warnf("failed to read and process exec frames from WebSocket: %v",
readFramesErr)
}
return responder.Empty()
case outgoingFrame, ok := <-subscriber.frames:
case outgoingFrame, ok := <-outgoingFrames:
if !ok {
if err := wsConn.Close(websocket.StatusNormalClosure, "Command finished"); err != nil {
controller.logger.Warnf("exec: failed to close WebSocket cleanly: %v", err)
@@ -316,20 +191,15 @@ func (controller *Controller) serveExecSession(
}
}
var (
errExecSessionDetached = errors.New("exec session detached")
errExecSessionClosed = errors.New("exec session closed")
)
func parseExecSessionSpec(ctx *gin.Context, command string) (execSessionSpec, string, error) {
func parseExecOptions(ctx *gin.Context, command string) (sshexec.Options, string, error) {
interactive, err := parseExecInteractive(ctx)
if err != nil {
return execSessionSpec{}, "", err
return sshexec.Options{}, "", err
}
tty, err := parseExecBool(ctx, "tty")
if err != nil {
return execSessionSpec{}, "", err
return sshexec.Options{}, "", err
}
if tty {
interactive = true
@@ -337,35 +207,31 @@ func parseExecSessionSpec(ctx *gin.Context, command string) (execSessionSpec, st
rows, err := parseExecUint32(ctx.Query("rows"), "rows")
if err != nil {
return execSessionSpec{}, "", err
return sshexec.Options{}, "", err
}
cols, err := parseExecUint32(ctx.Query("cols"), "cols")
if err != nil {
return execSessionSpec{}, "", err
return sshexec.Options{}, "", err
}
if (rows == 0) != (cols == 0) {
return execSessionSpec{}, "", errors.New("\"rows\" and \"cols\" must be provided together")
return sshexec.Options{}, "", errors.New("\"rows\" and \"cols\" must be provided together")
}
spec := execSessionSpec{
command: command,
interactive: interactive,
tty: tty,
rows: rows,
cols: cols,
env: ctx.QueryMap("env"),
workdir: ctx.Query("workdir"),
options := sshexec.Options{
Interactive: interactive,
TTY: tty,
Rows: rows,
Cols: cols,
Env: ctx.QueryMap("env"),
Workdir: ctx.Query("workdir"),
}
runCommand, err := sshexec.CommandWithOptions(command, sshexec.Options{
Env: spec.env,
Workdir: spec.workdir,
})
runCommand, err := sshexec.CommandWithOptions(command, options)
if err != nil {
return execSessionSpec{}, "", err
return sshexec.Options{}, "", err
}
return spec, runCommand, nil
return options, runCommand, nil
}
func parseExecInteractive(ctx *gin.Context) (bool, error) {
@@ -426,12 +292,10 @@ func parseExecUint32(raw string, name string) (uint32, error) {
return uint32(value), nil
}
func (controller *Controller) readExecSessionFrames(
ctx context.Context,
wsConn *websocket.Conn,
session *execSession,
subscriber *execSessionSubscriber,
) error {
func readExecFrames(ctx context.Context, wsConn *websocket.Conn, exec *sshexec.Exec) error {
stdin := exec.Stdin()
stdinClosed := false
for {
var frame execstream.Frame
@@ -439,7 +303,7 @@ func (controller *Controller) readExecSessionFrames(
if err != nil {
var closeErr websocket.CloseError
if errors.As(err, &closeErr) && closeErr.Code == websocket.StatusNormalClosure {
return errExecSessionDetached
return nil
}
return fmt.Errorf("failed to read next frame from WebSocket: %w", err)
@@ -455,7 +319,18 @@ func (controller *Controller) readExecSessionFrames(
switch frame.Type {
case execstream.FrameTypeStdin:
if err := session.writeStdin(frame.Data); err != nil {
if stdin == nil || stdinClosed {
return fmt.Errorf("failed to handle %q frame: this exec session has no stdin enabled or it is already closed",
frame.Type)
}
if len(frame.Data) == 0 {
err = stdin.Close()
stdinClosed = true
} else {
_, err = stdin.Write(frame.Data)
}
if err != nil {
return fmt.Errorf("failed to handle %q frame: %w", frame.Type, err)
}
case execstream.FrameTypeResize:
@@ -463,35 +338,9 @@ func (controller *Controller) readExecSessionFrames(
return fmt.Errorf("failed to handle %q frame: terminal size is required", frame.Type)
}
if err := session.resize(frame.Terminal.Rows, frame.Terminal.Cols); err != nil {
if err := exec.Resize(frame.Terminal.Rows, frame.Terminal.Cols); err != nil {
return fmt.Errorf("failed to handle %q frame: %w", frame.Type, err)
}
case execstream.FrameTypeHistory:
if !session.policy.replayEnabled {
return fmt.Errorf("unexpected frame type received: %q", frame.Type)
}
session.sendHistory(subscriber, frame.Watermark)
case execstream.FrameTypeAck:
if !session.policy.replayEnabled {
return fmt.Errorf("unexpected frame type received: %q", frame.Type)
}
session.ack(frame.Watermark)
case execstream.FrameTypeDetach:
if !session.policy.replayEnabled {
return fmt.Errorf("unexpected frame type received: %q", frame.Type)
}
return errExecSessionDetached
case execstream.FrameTypeClose:
if !session.policy.replayEnabled {
return fmt.Errorf("unexpected frame type received: %q", frame.Type)
}
session.close()
return errExecSessionClosed
default:
return fmt.Errorf("unexpected frame type received: %q", frame.Type)
}
+13 -13
View File
@@ -6,6 +6,7 @@ import (
"net/http/httptest"
"testing"
"github.com/cirruslabs/orchard/internal/controller/sshexec"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
@@ -66,26 +67,25 @@ func TestParseExecInteractive(t *testing.T) {
}
}
func TestParseExecSessionSpec(t *testing.T) {
spec, runCommand, err := parseExecSessionSpec(
func TestParseExecOptions(t *testing.T) {
options, runCommand, err := parseExecOptions(
execQueryContext("interactive=true&tty=true&rows=24&cols=80&env[GREETING]=hello&workdir=/tmp"),
"printf '%s' \"$GREETING\"",
)
require.NoError(t, err)
require.Equal(t, execSessionSpec{
command: "printf '%s' \"$GREETING\"",
interactive: true,
tty: true,
rows: 24,
cols: 80,
env: map[string]string{"GREETING": "hello"},
workdir: "/tmp",
}, spec)
require.Equal(t, sshexec.Options{
Interactive: true,
TTY: true,
Rows: 24,
Cols: 80,
Env: map[string]string{"GREETING": "hello"},
Workdir: "/tmp",
}, options)
require.Equal(t, "cd '/tmp' || exit $?\nexport GREETING='hello'\nprintf '%s' \"$GREETING\"", runCommand)
}
func TestParseExecSessionSpecRejectsPartialTTYSize(t *testing.T) {
_, _, err := parseExecSessionSpec(execQueryContext("tty=true&rows=24"), "echo hello")
func TestParseExecOptionsRejectsPartialTTYSize(t *testing.T) {
_, _, err := parseExecOptions(execQueryContext("tty=true&rows=24"), "echo hello")
require.ErrorContains(t, err, "provided together")
}
-5
View File
@@ -56,7 +56,6 @@ type Controller struct {
ipRendezvous *rendezvous.Rendezvous[rendezvous.ResultWithErrorMessage[string]]
enableSwaggerDocs bool
workerOfflineTimeout time.Duration
execSessionRetentionTTL time.Duration
execSSHConnectionKeepaliveInterval time.Duration
experimentalRPCV2 bool
disableDBCompression bool
@@ -67,7 +66,6 @@ type Controller struct {
sshSigner ssh.Signer
sshNoClientAuth bool
sshServer *sshserver.SSHServer
execSessions *execSessionRegistry
execSSHClients *execSSHClientPool
single singleflight.Group
@@ -80,10 +78,8 @@ func New(opts ...Option) (*Controller, error) {
connRendezvous: rendezvous.New[rendezvous.ResultWithErrorMessage[net.Conn]](),
ipRendezvous: rendezvous.New[rendezvous.ResultWithErrorMessage[string]](),
workerOfflineTimeout: 3 * time.Minute,
execSessionRetentionTTL: 10 * time.Minute,
execSSHConnectionKeepaliveInterval: 30 * time.Second,
pingInterval: 30 * time.Second,
execSessions: newExecSessionRegistry(),
single: singleflight.Group{},
}
@@ -357,7 +353,6 @@ func (controller *Controller) Run(ctx context.Context) error {
go func() {
<-ctx.Done()
controller.execSessions.closeAll()
controller.execSSHClients.closeAll()
if err := controller.httpServer.Shutdown(ctx); err != nil {
-745
View File
@@ -1,745 +0,0 @@
package controller
import (
"context"
"errors"
"io"
"maps"
"net"
"sync"
"time"
"github.com/cirruslabs/orchard/internal/execstream"
)
const execSessionReplayBufferBytes = 4 * 1024 * 1024
type execSessionPolicy struct {
closeOnDetach bool
retainAfterExit bool
replayEnabled bool
blockOnSubscriberBackpressure bool
}
var (
legacyExecSessionPolicy = execSessionPolicy{
closeOnDetach: true,
blockOnSubscriberBackpressure: true,
}
reconnectableExecSessionPolicy = execSessionPolicy{
retainAfterExit: true,
replayEnabled: true,
}
)
type sshExecRunner interface {
Stdin() io.WriteCloser
Resize(rows uint32, cols uint32) error
Run(ctx context.Context, command string, outgoingFrames chan<- *execstream.Frame) error
Close() error
}
type execSessionSpec struct {
command string
interactive bool
tty bool
rows uint32
cols uint32
env map[string]string
workdir string
}
func (spec execSessionSpec) clone() execSessionSpec {
spec.env = maps.Clone(spec.env)
return spec
}
func (spec execSessionSpec) equal(other execSessionSpec) bool {
return spec.command == other.command &&
spec.interactive == other.interactive &&
spec.tty == other.tty &&
spec.rows == other.rows &&
spec.cols == other.cols &&
spec.workdir == other.workdir &&
maps.Equal(spec.env, other.env)
}
type execSessionKey struct {
vmName string
sessionID string
}
type execSessionCreation struct {
done chan struct{}
session *execSession
err error
}
type execSessionRegistry struct {
mu sync.Mutex
sessions map[execSessionKey]*execSession
creating map[execSessionKey]*execSessionCreation
}
func newExecSessionRegistry() *execSessionRegistry {
return &execSessionRegistry{
sessions: map[execSessionKey]*execSession{},
creating: map[execSessionKey]*execSessionCreation{},
}
}
func (registry *execSessionRegistry) get(key execSessionKey) (*execSession, bool) {
registry.mu.Lock()
defer registry.mu.Unlock()
session, ok := registry.sessions[key]
return session, ok
}
func (registry *execSessionRegistry) getOrCreate(
ctx context.Context,
key execSessionKey,
create func() (*execSession, error),
) (*execSession, bool, error) {
registry.mu.Lock()
if session, ok := registry.sessions[key]; ok {
registry.mu.Unlock()
return session, false, nil
}
if creation, ok := registry.creating[key]; ok {
registry.mu.Unlock()
select {
case <-ctx.Done():
return nil, false, ctx.Err()
case <-creation.done:
if creation.err != nil {
return nil, false, creation.err
}
return creation.session, false, nil
}
}
creation := &execSessionCreation{done: make(chan struct{})}
registry.creating[key] = creation
registry.mu.Unlock()
session, err := create()
registry.mu.Lock()
delete(registry.creating, key)
if err == nil {
registry.sessions[key] = session
}
creation.session = session
creation.err = err
close(creation.done)
registry.mu.Unlock()
return session, true, err
}
func (registry *execSessionRegistry) remove(key execSessionKey, expected *execSession) {
registry.mu.Lock()
defer registry.mu.Unlock()
if registry.sessions[key] == expected {
delete(registry.sessions, key)
}
}
func (registry *execSessionRegistry) closeAll() {
registry.mu.Lock()
sessions := make([]*execSession, 0, len(registry.sessions))
for _, session := range registry.sessions {
sessions = append(sessions, session)
}
registry.mu.Unlock()
for _, session := range sessions {
session.close()
}
}
type execReplayFrame struct {
frame *execstream.Frame
size int
}
type execReplayBuffer struct {
frames []execReplayFrame
bufferBytes int
nextWatermark uint64
ackedWatermark uint64
}
func (buffer *execReplayBuffer) append(frame *execstream.Frame) *execstream.Frame {
frame = cloneExecFrame(frame)
buffer.nextWatermark++
frame.Watermark = buffer.nextWatermark
frameSize := execFrameSize(frame)
buffer.frames = append(buffer.frames, execReplayFrame{
frame: frame,
size: frameSize,
})
buffer.bufferBytes += frameSize
buffer.trimAcknowledged()
buffer.trimToLimit()
return frame
}
func (buffer *execReplayBuffer) ack(watermark uint64) {
if watermark <= buffer.ackedWatermark {
return
}
buffer.ackedWatermark = watermark
buffer.trimAcknowledged()
}
func (buffer *execReplayBuffer) replayAfter(
watermark uint64,
frames []*execstream.Frame,
) []*execstream.Frame {
for _, record := range buffer.frames {
if record.frame.Watermark <= watermark {
continue
}
frames = append(frames, record.frame)
}
return frames
}
func (buffer *execReplayBuffer) trimAcknowledged() {
for len(buffer.frames) > 0 && buffer.frames[0].frame.Watermark <= buffer.ackedWatermark {
buffer.bufferBytes -= buffer.frames[0].size
buffer.frames = buffer.frames[1:]
}
}
func (buffer *execReplayBuffer) trimToLimit() {
for buffer.bufferBytes > execSessionReplayBufferBytes && len(buffer.frames) > 0 {
buffer.bufferBytes -= buffer.frames[0].size
buffer.frames = buffer.frames[1:]
}
}
type execSessionSubscriber struct {
frames chan *execstream.Frame
closed chan struct{}
closeOnce sync.Once
sendMu sync.Mutex
sentWatermark uint64
}
func newExecSessionSubscriber() *execSessionSubscriber {
return &execSessionSubscriber{
frames: make(chan *execstream.Frame, 128),
closed: make(chan struct{}),
}
}
func (subscriber *execSessionSubscriber) enqueue(frame *execstream.Frame, block bool) bool {
subscriber.sendMu.Lock()
defer subscriber.sendMu.Unlock()
if block {
return subscriber.sendLocked(frame)
}
if subscriber.alreadySentLocked(frame) {
return true
}
select {
case <-subscriber.closed:
return false
default:
}
select {
case subscriber.frames <- subscriber.markSentLocked(frame):
return true
case <-subscriber.closed:
return false
default:
return false
}
}
func (subscriber *execSessionSubscriber) sendHistory(frames []*execstream.Frame) bool {
for _, frame := range frames {
if !subscriber.sendLocked(frame) {
return false
}
}
return true
}
func (subscriber *execSessionSubscriber) sendLocked(frame *execstream.Frame) bool {
if subscriber.alreadySentLocked(frame) {
return true
}
select {
case <-subscriber.closed:
return false
default:
}
select {
case subscriber.frames <- subscriber.markSentLocked(frame):
return true
case <-subscriber.closed:
return false
}
}
func (subscriber *execSessionSubscriber) alreadySentLocked(frame *execstream.Frame) bool {
return isReplayOutputFrame(frame) &&
frame.Watermark != 0 &&
frame.Watermark <= subscriber.sentWatermark
}
func (subscriber *execSessionSubscriber) markSentLocked(frame *execstream.Frame) *execstream.Frame {
frame = cloneExecFrame(frame)
if isReplayOutputFrame(frame) && frame.Watermark > subscriber.sentWatermark {
subscriber.sentWatermark = frame.Watermark
}
return frame
}
func (subscriber *execSessionSubscriber) close() {
subscriber.closeOnce.Do(func() {
close(subscriber.closed)
subscriber.sendMu.Lock()
close(subscriber.frames)
subscriber.sendMu.Unlock()
})
}
type execSession struct {
key execSessionKey
spec execSessionSpec
command string
exec sshExecRunner
transport net.Conn
registry *execSessionRegistry
retentionTTL time.Duration
policy execSessionPolicy
ctx context.Context
cancel context.CancelFunc
mu sync.Mutex
stdin io.WriteCloser
stdinClosed bool
subscribers map[*execSessionSubscriber]struct{}
replay execReplayBuffer
started bool
finished bool
closed bool
expiryTimer *time.Timer
startOnce sync.Once
done chan struct{}
doneOnce sync.Once
}
func newExecSession(
key execSessionKey,
command string,
exec sshExecRunner,
transport net.Conn,
registry *execSessionRegistry,
retentionTTL time.Duration,
policy execSessionPolicy,
) *execSession {
ctx, cancel := context.WithCancel(context.Background())
return newExecSessionWithContextAndSpec(
ctx,
cancel,
key,
execSessionSpec{command: command},
command,
exec,
transport,
registry,
retentionTTL,
policy,
)
}
func newExecSessionWithContextAndSpec(
ctx context.Context,
cancel context.CancelFunc,
key execSessionKey,
spec execSessionSpec,
command string,
exec sshExecRunner,
transport net.Conn,
registry *execSessionRegistry,
retentionTTL time.Duration,
policy execSessionPolicy,
) *execSession {
if ctx == nil || cancel == nil {
ctx, cancel = context.WithCancel(context.Background())
}
session := &execSession{
key: key,
spec: spec.clone(),
command: command,
exec: exec,
transport: transport,
registry: registry,
retentionTTL: retentionTTL,
policy: policy,
ctx: ctx,
cancel: cancel,
stdin: exec.Stdin(),
subscribers: map[*execSessionSubscriber]struct{}{},
done: make(chan struct{}),
}
return session
}
func (session *execSession) specMatches(spec execSessionSpec) bool {
return spec.command == "" || session.spec.equal(spec)
}
func (session *execSession) start() {
session.startOnce.Do(func() {
session.mu.Lock()
if session.closed {
session.mu.Unlock()
return
}
session.started = true
session.mu.Unlock()
go session.run()
})
}
func (session *execSession) closeIfUnused() {
session.mu.Lock()
unused := !session.started && len(session.subscribers) == 0
session.mu.Unlock()
if unused {
session.close()
}
}
func (session *execSession) attach() (*execSessionSubscriber, error) {
session.mu.Lock()
defer session.mu.Unlock()
if session.closed {
return nil, errors.New("exec session is closed")
}
subscriber := newExecSessionSubscriber()
session.subscribers[subscriber] = struct{}{}
return subscriber, nil
}
func (session *execSession) detach(subscriber *execSessionSubscriber) {
if session.policy.closeOnDetach {
session.close()
return
}
session.mu.Lock()
defer session.mu.Unlock()
session.detachLocked(subscriber)
}
func (session *execSession) detachLocked(subscriber *execSessionSubscriber) {
if _, ok := session.subscribers[subscriber]; !ok {
return
}
delete(session.subscribers, subscriber)
subscriber.close()
}
func (session *execSession) writeStdin(data []byte) error {
session.mu.Lock()
defer session.mu.Unlock()
if session.stdin == nil || session.stdinClosed {
return errors.New("this exec session has no stdin enabled or it is already closed")
}
if len(data) == 0 {
if err := session.stdin.Close(); err != nil {
return err
}
session.stdinClosed = true
return nil
}
_, err := session.stdin.Write(data)
return err
}
func (session *execSession) resize(rows uint32, cols uint32) error {
session.mu.Lock()
defer session.mu.Unlock()
if !session.spec.tty {
return errors.New("this exec session has no TTY")
}
return session.exec.Resize(rows, cols)
}
func (session *execSession) ack(watermark uint64) {
if !session.policy.replayEnabled {
return
}
session.mu.Lock()
defer session.mu.Unlock()
session.replay.ack(watermark)
}
func (session *execSession) sendHistory(
subscriber *execSessionSubscriber,
watermark uint64,
) {
if !session.policy.replayEnabled {
return
}
session.mu.Lock()
if _, ok := session.subscribers[subscriber]; !ok {
session.mu.Unlock()
return
}
subscriber.sendMu.Lock()
frames := session.replay.replayAfter(watermark, nil)
frames = append(frames, &execstream.Frame{
Type: execstream.FrameTypeNoMoreHistory,
Watermark: session.replay.nextWatermark,
})
session.mu.Unlock()
ok := subscriber.sendHistory(frames)
subscriber.sendMu.Unlock()
if !ok {
session.dropSubscriber(subscriber)
}
}
func (session *execSession) close() {
session.mu.Lock()
if session.closed {
session.mu.Unlock()
return
}
session.closed = true
if session.expiryTimer != nil {
session.expiryTimer.Stop()
session.expiryTimer = nil
}
subscribers := session.takeSubscribersLocked()
session.mu.Unlock()
closeSubscribers(subscribers)
session.cancel()
_ = session.exec.Close()
if session.transport != nil {
_ = session.transport.Close()
}
if session.registry != nil {
session.registry.remove(session.key, session)
}
}
func (session *execSession) run() {
outgoingFrames := make(chan *execstream.Frame)
runErrCh := make(chan error, 1)
go func() {
runErrCh <- session.exec.Run(session.ctx, session.command, outgoingFrames)
close(outgoingFrames)
}()
for frame := range outgoingFrames {
session.recordFrame(frame)
}
runErr := <-runErrCh
if runErr != nil && !errors.Is(runErr, context.Canceled) {
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeError,
Error: runErr.Error(),
})
}
session.markFinished()
}
func (session *execSession) recordFrame(frame *execstream.Frame) {
session.mu.Lock()
if session.closed {
session.mu.Unlock()
return
}
if session.policy.replayEnabled {
frame = session.replay.append(frame)
} else {
frame = cloneExecFrame(frame)
}
subscribers := make([]*execSessionSubscriber, 0, len(session.subscribers))
for subscriber := range session.subscribers {
subscribers = append(subscribers, subscriber)
}
session.mu.Unlock()
for _, subscriber := range subscribers {
if !subscriber.enqueue(frame, session.policy.blockOnSubscriberBackpressure) {
session.dropSubscriber(subscriber)
}
}
}
func (session *execSession) markFinished() {
session.mu.Lock()
if session.finished {
session.mu.Unlock()
return
}
session.finished = true
shouldClose := !session.policy.retainAfterExit
if !session.closed && session.policy.retainAfterExit {
session.expiryTimer = time.AfterFunc(session.retentionTTL, session.expire)
}
var subscribers []*execSessionSubscriber
if shouldClose {
subscribers = session.takeSubscribersLocked()
}
session.mu.Unlock()
closeSubscribers(subscribers)
session.doneOnce.Do(func() {
close(session.done)
})
if shouldClose {
session.close()
}
}
func (session *execSession) expire() {
session.close()
}
func (session *execSession) takeSubscribersLocked() []*execSessionSubscriber {
subscribers := make([]*execSessionSubscriber, 0, len(session.subscribers))
for subscriber := range session.subscribers {
subscribers = append(subscribers, subscriber)
}
session.subscribers = map[*execSessionSubscriber]struct{}{}
return subscribers
}
func closeSubscribers(subscribers []*execSessionSubscriber) {
for _, subscriber := range subscribers {
subscriber.close()
}
}
func (session *execSession) dropSubscriber(subscriber *execSessionSubscriber) {
session.mu.Lock()
defer session.mu.Unlock()
session.detachLocked(subscriber)
}
func cloneExecFrame(frame *execstream.Frame) *execstream.Frame {
if frame == nil {
return nil
}
clone := *frame
if frame.Data != nil {
clone.Data = append([]byte(nil), frame.Data...)
}
if frame.Exit != nil {
exit := *frame.Exit
clone.Exit = &exit
}
if frame.Terminal != nil {
terminal := *frame.Terminal
clone.Terminal = &terminal
}
return &clone
}
func execFrameSize(frame *execstream.Frame) int {
if frame == nil {
return 0
}
return len(frame.Data) + len(frame.Error) + 16
}
func isReplayOutputFrame(frame *execstream.Frame) bool {
if frame == nil {
return false
}
switch frame.Type {
case execstream.FrameTypeStdout,
execstream.FrameTypeStderr,
execstream.FrameTypeExit,
execstream.FrameTypeError:
return true
default:
return false
}
}
-517
View File
@@ -1,517 +0,0 @@
package controller
import (
"context"
"io"
"sync/atomic"
"testing"
"time"
"github.com/cirruslabs/orchard/internal/execstream"
"github.com/stretchr/testify/require"
)
type fakeExec struct {
stdin io.WriteCloser
run func(context.Context, string, chan<- *execstream.Frame) error
resize func(uint32, uint32) error
closeCalls atomic.Int32
}
func (exec *fakeExec) Stdin() io.WriteCloser {
return exec.stdin
}
func (exec *fakeExec) Run(
ctx context.Context,
command string,
outgoingFrames chan<- *execstream.Frame,
) error {
if exec.run != nil {
return exec.run(ctx, command, outgoingFrames)
}
return nil
}
func (exec *fakeExec) Resize(rows uint32, cols uint32) error {
if exec.resize != nil {
return exec.resize(rows, cols)
}
return nil
}
func (exec *fakeExec) Close() error {
exec.closeCalls.Add(1)
return nil
}
func newManualExecSessionForTest(
key execSessionKey,
registry *execSessionRegistry,
) *execSession {
ctx, cancel := context.WithCancel(context.Background())
return &execSession{
key: key,
spec: execSessionSpec{command: "echo test"},
command: "echo test",
exec: &fakeExec{},
registry: registry,
retentionTTL: time.Minute,
policy: reconnectableExecSessionPolicy,
ctx: ctx,
cancel: cancel,
subscribers: map[*execSessionSubscriber]struct{}{},
done: make(chan struct{}),
}
}
func TestExecSessionRegistryGetOrCreateReusesInflightCreation(t *testing.T) {
registry := newExecSessionRegistry()
key := execSessionKey{vmName: "vm", sessionID: "session"}
createStarted := make(chan struct{})
releaseCreate := make(chan struct{})
var createCalls atomic.Int32
create := func() (*execSession, error) {
createCalls.Add(1)
close(createStarted)
<-releaseCreate
return newManualExecSessionForTest(key, registry), nil
}
firstDone := make(chan struct{})
go func() {
defer close(firstDone)
_, _, err := registry.getOrCreate(context.Background(), key, create)
require.NoError(t, err)
}()
<-createStarted
secondDone := make(chan struct{})
go func() {
defer close(secondDone)
_, created, err := registry.getOrCreate(context.Background(), key, create)
require.NoError(t, err)
require.False(t, created)
}()
close(releaseCreate)
<-firstDone
<-secondDone
require.EqualValues(t, 1, createCalls.Load())
}
func TestExecSessionStartRunsCommandOnlyOnce(t *testing.T) {
var runCalls atomic.Int32
runStarted := make(chan struct{})
session := newExecSession(
execSessionKey{vmName: "vm", sessionID: "session"},
"echo test",
&fakeExec{
run: func(ctx context.Context, _ string, _ chan<- *execstream.Frame) error {
runCalls.Add(1)
close(runStarted)
<-ctx.Done()
return ctx.Err()
},
},
nil,
nil,
time.Minute,
reconnectableExecSessionPolicy,
)
defer session.close()
session.start()
session.start()
<-runStarted
require.EqualValues(t, 1, runCalls.Load())
}
func TestExecSessionSpecMatchesOptions(t *testing.T) {
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, nil)
session.spec = execSessionSpec{
command: "echo test",
interactive: true,
tty: true,
rows: 24,
cols: 80,
env: map[string]string{"GREETING": "hello"},
workdir: "/tmp",
}
require.True(t, session.specMatches(execSessionSpec{}))
require.True(t, session.specMatches(execSessionSpec{
command: "echo test",
interactive: true,
tty: true,
rows: 24,
cols: 80,
env: map[string]string{"GREETING": "hello"},
workdir: "/tmp",
}))
require.False(t, session.specMatches(execSessionSpec{
command: "echo test",
interactive: true,
tty: true,
rows: 24,
cols: 80,
env: map[string]string{"GREETING": "goodbye"},
workdir: "/tmp",
}))
}
func TestExecSessionResizeRequiresTTY(t *testing.T) {
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, nil)
err := session.resize(24, 80)
require.ErrorContains(t, err, "no TTY")
}
func TestExecSessionResizeDelegatesToRunner(t *testing.T) {
var resizedRows, resizedCols uint32
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, nil)
session.spec.tty = true
session.exec = &fakeExec{
resize: func(rows uint32, cols uint32) error {
resizedRows = rows
resizedCols = cols
return nil
},
}
require.NoError(t, session.resize(24, 80))
require.EqualValues(t, 24, resizedRows)
require.EqualValues(t, 80, resizedCols)
}
func TestExecSessionHistoryReplayAndAck(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
session.recordFrame(&execstream.Frame{Type: execstream.FrameTypeStdout, Data: []byte("out")})
session.recordFrame(&execstream.Frame{Type: execstream.FrameTypeStderr, Data: []byte("err")})
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeExit,
Exit: &execstream.Exit{Code: 7},
})
subscriber, err := session.attach()
require.NoError(t, err)
session.sendHistory(subscriber, 0)
require.Equal(t, execstream.FrameTypeStdout, (<-subscriber.frames).Type)
require.Equal(t, execstream.FrameTypeStderr, (<-subscriber.frames).Type)
require.Equal(t, execstream.FrameTypeExit, (<-subscriber.frames).Type)
noMoreHistory := <-subscriber.frames
require.Equal(t, execstream.FrameTypeNoMoreHistory, noMoreHistory.Type)
require.EqualValues(t, 3, noMoreHistory.Watermark)
session.ack(2)
require.Len(t, session.replay.frames, 1)
require.EqualValues(t, 3, session.replay.frames[0].frame.Watermark)
}
func TestExecSessionHistoryReplayStreamsPastSubscriberBuffer(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
const frameCount = 256
for i := 0; i < frameCount; i++ {
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeStdout,
Data: []byte{byte(i)},
})
}
subscriber, err := session.attach()
require.NoError(t, err)
done := make(chan struct{})
go func() {
defer close(done)
session.sendHistory(subscriber, 0)
}()
for i := 1; i <= frameCount; i++ {
frame := <-subscriber.frames
require.Equal(t, execstream.FrameTypeStdout, frame.Type)
require.EqualValues(t, i, frame.Watermark)
}
noMoreHistory := <-subscriber.frames
require.Equal(t, execstream.FrameTypeNoMoreHistory, noMoreHistory.Type)
require.EqualValues(t, frameCount, noMoreHistory.Watermark)
require.Eventually(t, func() bool {
select {
case <-done:
return true
default:
return false
}
}, time.Second, 10*time.Millisecond)
}
func TestExecSessionLiveOutputAppliesBackpressure(t *testing.T) {
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, nil)
session.policy = legacyExecSessionPolicy
t.Cleanup(session.close)
subscriber, err := session.attach()
require.NoError(t, err)
const frameCount = 256
done := make(chan struct{})
go func() {
defer close(done)
for i := range frameCount {
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeStdout,
Data: []byte{byte(i)},
Terminal: nil,
Exit: nil,
Error: "",
Watermark: 0,
})
}
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeExit,
Data: nil,
Terminal: nil,
Exit: &execstream.Exit{Code: 0},
Error: "",
Watermark: 0,
})
}()
require.Eventually(t, func() bool {
return len(subscriber.frames) == cap(subscriber.frames)
}, time.Second, time.Millisecond)
for i := range frameCount {
frame, ok := <-subscriber.frames
require.True(t, ok, "subscriber closed before output frame %d", i)
require.Equal(t, execstream.FrameTypeStdout, frame.Type)
require.Equal(t, []byte{byte(i)}, frame.Data)
}
exitFrame, ok := <-subscriber.frames
require.True(t, ok, "subscriber closed before the exit frame")
require.Equal(t, execstream.FrameTypeExit, exitFrame.Type)
require.EqualValues(t, 0, exitFrame.Exit.Code)
require.Eventually(t, func() bool {
select {
case <-done:
return true
default:
return false
}
}, time.Second, time.Millisecond)
}
func TestReconnectableExecSessionDropsStalledSubscriber(t *testing.T) {
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, nil)
t.Cleanup(session.close)
stalledSubscriber, err := session.attach()
require.NoError(t, err)
for i := range cap(stalledSubscriber.frames) {
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeStdout,
Data: []byte{byte(i)},
Terminal: nil,
Exit: nil,
Error: "",
Watermark: 0,
})
}
healthySubscriber, err := session.attach()
require.NoError(t, err)
done := make(chan struct{})
go func() {
defer close(done)
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeStdout,
Data: []byte("still running"),
Terminal: nil,
Exit: nil,
Error: "",
Watermark: 0,
})
session.recordFrame(&execstream.Frame{
Type: execstream.FrameTypeExit,
Data: nil,
Terminal: nil,
Exit: &execstream.Exit{Code: 0},
Error: "",
Watermark: 0,
})
}()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("a stalled reconnectable subscriber blocked live output")
}
outputFrame := <-healthySubscriber.frames
require.Equal(t, execstream.FrameTypeStdout, outputFrame.Type)
require.Equal(t, []byte("still running"), outputFrame.Data)
exitFrame := <-healthySubscriber.frames
require.Equal(t, execstream.FrameTypeExit, exitFrame.Type)
require.EqualValues(t, 0, exitFrame.Exit.Code)
select {
case <-stalledSubscriber.closed:
default:
t.Fatal("the stalled reconnectable subscriber was not dropped")
}
reconnectedSubscriber, err := session.attach()
require.NoError(t, err)
session.sendHistory(reconnectedSubscriber, uint64(cap(stalledSubscriber.frames)))
require.Equal(t, execstream.FrameTypeStdout, (<-reconnectedSubscriber.frames).Type)
require.Equal(t, execstream.FrameTypeExit, (<-reconnectedSubscriber.frames).Type)
require.Equal(t, execstream.FrameTypeNoMoreHistory, (<-reconnectedSubscriber.frames).Type)
}
func TestExecSessionDetachKeepsProcessAlive(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
exec := session.exec.(*fakeExec)
subscriber, err := session.attach()
require.NoError(t, err)
session.detach(subscriber)
require.False(t, session.closed)
require.EqualValues(t, 0, exec.closeCalls.Load())
}
func TestLegacyExecSessionDetachStopsProcess(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
session.policy = legacyExecSessionPolicy
exec := session.exec.(*fakeExec)
subscriber, err := session.attach()
require.NoError(t, err)
session.detach(subscriber)
require.True(t, session.closed)
require.EqualValues(t, 1, exec.closeCalls.Load())
}
func TestLegacyExecSessionDoesNotRetainReplayHistory(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
session.policy = legacyExecSessionPolicy
session.recordFrame(&execstream.Frame{Type: execstream.FrameTypeStdout, Data: []byte("out")})
require.Empty(t, session.replay.frames)
require.Zero(t, session.replay.nextWatermark)
}
func TestExecSessionCloseIfUnusedClosesIdleSession(t *testing.T) {
registry := newExecSessionRegistry()
key := execSessionKey{vmName: "vm", sessionID: "session"}
session := newManualExecSessionForTest(key, registry)
exec := session.exec.(*fakeExec)
registry.sessions[key] = session
session.closeIfUnused()
require.True(t, session.closed)
require.EqualValues(t, 1, exec.closeCalls.Load())
}
func TestExecSessionCloseIfUnusedKeepsAttachedSession(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
_, err := session.attach()
require.NoError(t, err)
session.closeIfUnused()
require.False(t, session.closed)
}
func TestExecSessionCloseStopsProcessAndRemovesRegistryEntry(t *testing.T) {
registry := newExecSessionRegistry()
key := execSessionKey{vmName: "vm", sessionID: "session"}
session := newManualExecSessionForTest(key, registry)
exec := session.exec.(*fakeExec)
registry.sessions[key] = session
session.close()
require.True(t, session.closed)
require.EqualValues(t, 1, exec.closeCalls.Load())
_, ok := registry.get(key)
require.False(t, ok)
}
func TestExecSessionFinishedEntryExpiresAfterTTL(t *testing.T) {
registry := newExecSessionRegistry()
key := execSessionKey{vmName: "vm", sessionID: "session"}
session := newManualExecSessionForTest(key, registry)
session.retentionTTL = 10 * time.Millisecond
registry.sessions[key] = session
session.markFinished()
require.Eventually(t, func() bool {
_, ok := registry.get(key)
return !ok
}, time.Second, 10*time.Millisecond)
}
func TestExecSessionFinishKeepsReconnectableSubscriberOpen(t *testing.T) {
registry := newExecSessionRegistry()
session := newManualExecSessionForTest(execSessionKey{vmName: "vm", sessionID: "session"}, registry)
subscriber, err := session.attach()
require.NoError(t, err)
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.Equal(t, execstream.FrameTypeStdout, (<-subscriber.frames).Type)
require.Equal(t, execstream.FrameTypeExit, (<-subscriber.frames).Type)
session.sendHistory(subscriber, 0)
noMoreHistory, ok := <-subscriber.frames
require.True(t, ok)
require.Equal(t, execstream.FrameTypeNoMoreHistory, noMoreHistory.Type)
require.EqualValues(t, 2, noMoreHistory.Watermark)
}
-6
View File
@@ -66,12 +66,6 @@ func WithWorkerOfflineTimeout(workerOfflineTimeout time.Duration) Option {
}
}
func WithExecSessionRetentionTTL(execSessionRetentionTTL time.Duration) Option {
return func(controller *Controller) {
controller.execSessionRetentionTTL = execSessionRetentionTTL
}
}
func WithExecSSHConnectionKeepaliveInterval(execSSHConnectionKeepaliveInterval time.Duration) Option {
return func(controller *Controller) {
controller.execSSHConnectionKeepaliveInterval = execSSHConnectionKeepaliveInterval