/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
+6
View File
@@ -70,6 +70,12 @@ linters:
# Not all errors need to be checked
- errcheck
# It's OK to have "magic" numbers in things like time.Sleep()
- mnd
# This is not a library, so it's OK to use dynamic errors
- err113
# It's OK to not initialize some struct fields
- exhaustruct
- exhaustruct_v5
+3 -118
View File
@@ -451,33 +451,11 @@ paths:
parameters:
- in: query
name: command
description: |
Command to execute.
Required when starting a new exec session. May be omitted when reconnecting to an
existing session identified by the `session` parameter.
description: Command to execute.
schema:
type: string
minLength: 1
required: false
- in: query
name: session
description: |
Optional stable exec session identifier. When present, websocket disconnects detach
from the command instead of terminating it, and later requests with the same VM name
and session id may reconnect and request buffered history.
schema:
type: string
minLength: 1
required: false
- in: query
name: cmux_session_id
description: |
Compatibility alias for `session`. Prefer `session` for new Orchard clients.
schema:
type: string
minLength: 1
required: false
required: true
- in: query
name: interactive
description: |
@@ -580,9 +558,7 @@ paths:
'400':
description: Invalid parameters were supplied
'404':
description: VM resource with the given name or reconnectable exec session doesn't exist
'409':
description: Reconnectable exec session already exists with different options
description: VM resource with the given name doesn't exist
'503':
description: Controller failed to establish a connection with the VM
/vms/{name}/ip:
@@ -1083,19 +1059,11 @@ components:
oneOf:
- $ref: '#/components/schemas/ExecClientFrameStdin'
- $ref: '#/components/schemas/ExecClientFrameResize'
- $ref: '#/components/schemas/ExecClientFrameHistory'
- $ref: '#/components/schemas/ExecClientFrameAck'
- $ref: '#/components/schemas/ExecClientFrameDetach'
- $ref: '#/components/schemas/ExecClientFrameClose'
discriminator:
propertyName: type
mapping:
stdin: '#/components/schemas/ExecClientFrameStdin'
resize: '#/components/schemas/ExecClientFrameResize'
history: '#/components/schemas/ExecClientFrameHistory'
ack: '#/components/schemas/ExecClientFrameAck'
detach: '#/components/schemas/ExecClientFrameDetach'
close: '#/components/schemas/ExecClientFrameClose'
ExecClientFrameStdin:
description: Send bytes to the process standard input
type: object
@@ -1129,56 +1097,6 @@ components:
terminal:
rows: 40
cols: 120
ExecClientFrameHistory:
description: Request buffered output strictly newer than the supplied watermark
type: object
required: [ type, watermark ]
properties:
type:
type: string
enum: [ history ]
watermark:
type: integer
format: int64
minimum: 0
example:
type: history
watermark: 42
ExecClientFrameAck:
description: Acknowledge that output has been durably consumed through the supplied watermark
type: object
required: [ type, watermark ]
properties:
type:
type: string
enum: [ ack ]
watermark:
type: integer
format: int64
minimum: 0
example:
type: ack
watermark: 42
ExecClientFrameDetach:
description: Detach this websocket while leaving the remote command running
type: object
required: [ type ]
properties:
type:
type: string
enum: [ detach ]
example:
type: detach
ExecClientFrameClose:
description: Close the reconnectable exec session and terminate the remote command
type: object
required: [ type ]
properties:
type:
type: string
enum: [ close ]
example:
type: close
ExecControllerFrame:
description: WebSocket frame from Orchard Controller to the Orchard Client
oneOf:
@@ -1186,7 +1104,6 @@ components:
- $ref: '#/components/schemas/ExecControllerFrameStderr'
- $ref: '#/components/schemas/ExecControllerFrameExit'
- $ref: '#/components/schemas/ExecControllerFrameError'
- $ref: '#/components/schemas/ExecControllerFrameNoMoreHistory'
discriminator:
propertyName: type
mapping:
@@ -1194,7 +1111,6 @@ components:
stderr: '#/components/schemas/ExecControllerFrameStderr'
exit: '#/components/schemas/ExecControllerFrameExit'
error: '#/components/schemas/ExecControllerFrameError'
no_more_history: '#/components/schemas/ExecControllerFrameNoMoreHistory'
ExecControllerFrameStdout:
description: Standard output from the process
type: object
@@ -1207,10 +1123,6 @@ components:
type: string
format: byte
description: Base64-encoded standard output bytes from the process
watermark:
type: integer
format: int64
description: Monotonic output watermark present on reconnectable sessions
example:
type: stdout
data: aGVsbG8K
@@ -1226,10 +1138,6 @@ components:
type: string
format: byte
description: Base64-encoded standard error bytes from the process
watermark:
type: integer
format: int64
description: Monotonic output watermark present on reconnectable sessions
example:
type: stderr
data: aGVsbG8K
@@ -1249,10 +1157,6 @@ components:
type: integer
format: int32
description: Process exit code
watermark:
type: integer
format: int64
description: Monotonic output watermark present on reconnectable sessions
example:
type: exit
exit:
@@ -1268,28 +1172,9 @@ components:
error:
type: string
description: Error message text
watermark:
type: integer
format: int64
description: Monotonic output watermark present on reconnectable sessions
example:
type: error
error: Failed to establish SSH connection to a VM
ExecControllerFrameNoMoreHistory:
description: Marker indicating that the requested replay range has been fully sent
type: object
required: [ type, watermark ]
properties:
type:
type: string
enum: [ no_more_history ]
watermark:
type: integer
format: int64
description: Highest watermark known to the exec session at the time of replay
example:
type: no_more_history
watermark: 42
ExecTerminalSize:
description: Pseudo-terminal size
type: object
-4
View File
@@ -36,7 +36,6 @@ var noExperimentalRPCV2 bool
var experimentalPingInterval time.Duration
var experimentalDisableDBCompression bool
var workerOfflineTimeout time.Duration
var execSessionRetentionTTL time.Duration
var execSSHConnectionKeepaliveInterval time.Duration
var synthetic bool
@@ -92,8 +91,6 @@ func newRunCommand() *cobra.Command {
"duration (e.g. 60s or 5m30s) after which a worker is considered offline for the purposes "+
"of scheduling (no new VMs will be scheduled on such worker and already assigned VMs will be "+
"marked as failed)")
cmd.Flags().DurationVar(&execSessionRetentionTTL, "exec-session-retention-ttl", 10*time.Minute,
"duration to retain reconnectable exec session history after the command exits")
cmd.Flags().DurationVar(&execSSHConnectionKeepaliveInterval, "exec-ssh-connection-keepalive-interval", 30*time.Second,
"interval between SSH keepalive requests sent by the controller for shared exec connections")
@@ -156,7 +153,6 @@ func runController(cmd *cobra.Command, args []string) (err error) {
controller.WithListenAddr(address),
controller.WithDataDir(dataDir),
controller.WithWorkerOfflineTimeout(workerOfflineTimeout),
controller.WithExecSessionRetentionTTL(execSessionRetentionTTL),
controller.WithExecSSHConnectionKeepaliveInterval(execSSHConnectionKeepaliveInterval),
controller.WithLogger(logger),
}
+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
+11 -17
View File
@@ -10,26 +10,20 @@ import (
type FrameType string
const (
FrameTypeStdin FrameType = "stdin"
FrameTypeResize FrameType = "resize"
FrameTypeStdout FrameType = "stdout"
FrameTypeStderr FrameType = "stderr"
FrameTypeExit FrameType = "exit"
FrameTypeError FrameType = "error"
FrameTypeHistory FrameType = "history"
FrameTypeNoMoreHistory FrameType = "no_more_history"
FrameTypeAck FrameType = "ack"
FrameTypeDetach FrameType = "detach"
FrameTypeClose FrameType = "close"
FrameTypeStdin FrameType = "stdin"
FrameTypeResize FrameType = "resize"
FrameTypeStdout FrameType = "stdout"
FrameTypeStderr FrameType = "stderr"
FrameTypeExit FrameType = "exit"
FrameTypeError FrameType = "error"
)
type Frame struct {
Type FrameType `json:"type"`
Data []byte `json:"data,omitempty"`
Terminal *TerminalSize `json:"terminal,omitempty"`
Exit *Exit `json:"exit,omitempty"`
Error string `json:"error,omitempty"`
Watermark uint64 `json:"watermark,omitempty"`
Type FrameType `json:"type"`
Data []byte `json:"data,omitempty"`
Terminal *TerminalSize `json:"terminal,omitempty"`
Exit *Exit `json:"exit,omitempty"`
Error string `json:"error,omitempty"`
}
type Exit struct {
-15
View File
@@ -7,21 +7,6 @@ import (
"github.com/stretchr/testify/require"
)
func TestFrameRoundTripsWatermark(t *testing.T) {
frame := Frame{
Type: FrameTypeHistory,
Watermark: 42,
}
payload, err := json.Marshal(frame)
require.NoError(t, err)
var decoded Frame
err = json.Unmarshal(payload, &decoded)
require.NoError(t, err)
require.Equal(t, frame, decoded)
}
func TestFrameRoundTripsTerminalSize(t *testing.T) {
frame := Frame{
Type: FrameTypeResize,
-210
View File
@@ -365,176 +365,6 @@ func TestVMExecKeepsSharedSSHClientAfterSessionRejection(t *testing.T) {
require.EqualValues(t, 1, sshServer.SuccessfulConnections())
}
func TestVMExecSessionReconnectHistory(t *testing.T) {
devClient, vmName := prepareForExec(t)
sessionID := uuid.NewString()
wsConn, err := devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
Command: "sh -c 'echo first; sleep 1; echo second'",
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
firstFrame := readFrame(t, wsConn)
require.Equal(t, execstream.FrameTypeStdout, firstFrame.Type)
require.Equal(t, "first\n", string(firstFrame.Data))
require.EqualValues(t, 1, firstFrame.Watermark)
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeDetach})
require.NoError(t, err)
_ = wsConn.CloseNow()
// Let the detached process finish so this test verifies partial replay
// without relying on live-output timing.
time.Sleep(2 * time.Second)
wsConn, err = devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
defer wsConn.CloseNow()
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{
Type: execstream.FrameTypeHistory,
Watermark: firstFrame.Watermark,
})
require.NoError(t, err)
frames := readFramesUntilExit(t, wsConn)
require.Len(t, framesByType(frames, execstream.FrameTypeStdout), 1)
require.Equal(t, "second\n", string(framesByType(frames, execstream.FrameTypeStdout)[0].Data))
require.EqualValues(t, 0, framesByType(frames, execstream.FrameTypeExit)[0].Exit.Code)
}
func TestVMExecSessionReconnectAfterExit(t *testing.T) {
devClient, vmName := prepareForExec(t)
sessionID := uuid.NewString()
wsConn, err := devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
Command: "sh -c 'echo replay-me'",
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeDetach})
require.NoError(t, err)
_ = wsConn.CloseNow()
time.Sleep(time.Second)
wsConn, err = devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
defer wsConn.CloseNow()
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeHistory})
require.NoError(t, err)
frames := readFramesUntilExit(t, wsConn)
require.Equal(t, "replay-me\n", string(framesByType(frames, execstream.FrameTypeStdout)[0].Data))
require.EqualValues(t, 0, framesByType(frames, execstream.FrameTypeExit)[0].Exit.Code)
}
func TestVMExecSessionReplayPreservesStreams(t *testing.T) {
devClient, vmName := prepareForExec(t)
sessionID := uuid.NewString()
wsConn, err := devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
Command: "sh -c 'echo out1; sleep 1; echo err1 >&2; sleep 1; echo out2; sleep 1; echo err2 >&2'",
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeDetach})
require.NoError(t, err)
_ = wsConn.CloseNow()
time.Sleep(4 * time.Second)
wsConn, err = devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
defer wsConn.CloseNow()
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeHistory})
require.NoError(t, err)
frames := readFramesUntilExit(t, wsConn)
require.Equal(t, []execstream.FrameType{
execstream.FrameTypeStdout,
execstream.FrameTypeStderr,
execstream.FrameTypeStdout,
execstream.FrameTypeStderr,
execstream.FrameTypeExit,
}, frameTypes(frames))
require.Equal(t, "out1\n", string(frames[0].Data))
require.Equal(t, "err1\n", string(frames[1].Data))
require.Equal(t, "out2\n", string(frames[2].Data))
require.Equal(t, "err2\n", string(frames[3].Data))
}
func TestVMExecSessionStdinSurvivesReconnect(t *testing.T) {
devClient, vmName := prepareForExec(t)
sessionID := uuid.NewString()
wsConn, err := devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
Command: "/bin/cat",
Interactive: true,
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{
Type: execstream.FrameTypeStdin,
Data: []byte("one\n"),
})
require.NoError(t, err)
frame := readFrame(t, wsConn)
require.Equal(t, execstream.FrameTypeStdout, frame.Type)
require.Equal(t, "one\n", string(frame.Data))
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeDetach})
require.NoError(t, err)
_ = wsConn.CloseNow()
wsConn, err = devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{
WaitSeconds: 30,
Session: sessionID,
})
require.NoError(t, err)
defer wsConn.CloseNow()
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{
Type: execstream.FrameTypeStdin,
Data: []byte("two\n"),
})
require.NoError(t, err)
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{
Type: execstream.FrameTypeStdin,
Data: []byte{},
})
require.NoError(t, err)
err = execstream.WriteFrame(t.Context(), wsConn, &execstream.Frame{Type: execstream.FrameTypeHistory})
require.NoError(t, err)
frames := readFramesUntilExit(t, wsConn)
stdoutFrames := framesByType(frames, execstream.FrameTypeStdout)
require.Len(t, stdoutFrames, 2)
require.Equal(t, "one\n", string(stdoutFrames[0].Data))
require.Equal(t, "two\n", string(stdoutFrames[1].Data))
require.EqualValues(t, 0, framesByType(frames, execstream.FrameTypeExit)[0].Exit.Code)
}
func prepareForExec(t *testing.T) (*client.Client, string) {
devClient, _, _ := devcontroller.StartIntegrationTestEnvironment(t)
@@ -635,43 +465,3 @@ func readFrameErr(ctx context.Context, wsConn *websocket.Conn) (*execstream.Fram
return &frame, nil
}
func readFramesUntilExit(t *testing.T, wsConn *websocket.Conn) []*execstream.Frame {
t.Helper()
var frames []*execstream.Frame
for {
frame := readFrame(t, wsConn)
if frame.Type == execstream.FrameTypeNoMoreHistory {
continue
}
frames = append(frames, frame)
if frame.Type == execstream.FrameTypeExit {
return frames
}
}
}
func framesByType(frames []*execstream.Frame, frameType execstream.FrameType) []*execstream.Frame {
var result []*execstream.Frame
for _, frame := range frames {
if frame.Type == frameType {
result = append(result, frame)
}
}
return result
}
func frameTypes(frames []*execstream.Frame) []execstream.FrameType {
var result []execstream.FrameType
for _, frame := range frames {
result = append(result, frame.Type)
}
return result
}
-4
View File
@@ -71,7 +71,6 @@ type ExecSessionOptions struct {
Env map[string]string
Workdir string
WaitSeconds uint16
Session string
}
func (service *VMsService) Create(ctx context.Context, vm *v1.VM) error {
@@ -253,9 +252,6 @@ func (service *VMsService) ExecSession(
if options.Workdir != "" {
params["workdir"] = options.Workdir
}
if options.Session != "" {
params["session"] = options.Session
}
return service.client.wsRequestRaw(ctx, fmt.Sprintf("vms/%s/exec", url.PathEscape(name)),
params)
+1 -3
View File
@@ -20,7 +20,7 @@ func TestHTTPClientForWebSocketHonorsWait(t *testing.T) {
require.Same(t, devClient.httpClient.Transport, httpClient.Transport)
}
func TestExecSessionBuildsReconnectableQuery(t *testing.T) {
func TestExecSessionBuildsQuery(t *testing.T) {
var query map[string][]string
server := httptest.NewServer(
@@ -51,7 +51,6 @@ func TestExecSessionBuildsReconnectableQuery(t *testing.T) {
Env: map[string]string{"GREETING": "hello"},
Workdir: "/tmp",
WaitSeconds: 7,
Session: "resume-me",
})
require.NoError(t, err)
defer conn.CloseNow()
@@ -64,5 +63,4 @@ func TestExecSessionBuildsReconnectableQuery(t *testing.T) {
require.Equal(t, []string{"hello"}, query["env[GREETING]"])
require.Equal(t, []string{"/tmp"}, query["workdir"])
require.Equal(t, []string{"7"}, query[waitParameterName])
require.Equal(t, []string{"resume-me"}, query["session"])
}