mirror of
https://github.com/cirruslabs/orchard.git
synced 2026-09-30 20:11:20 +02:00
/exec API: drop session reconnection support (#487)
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"])
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user