diff --git a/.golangci.yml b/.golangci.yml index 1ebf01b..ae3c213 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -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 diff --git a/api/openapi.yaml b/api/openapi.yaml index 5796e4b..76dbb06 100644 --- a/api/openapi.yaml +++ b/api/openapi.yaml @@ -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 diff --git a/internal/command/controller/run.go b/internal/command/controller/run.go index 8c37222..22431ab 100644 --- a/internal/command/controller/run.go +++ b/internal/command/controller/run.go @@ -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), } diff --git a/internal/controller/api_vms_exec.go b/internal/controller/api_vms_exec.go index 235976e..9125898 100644 --- a/internal/controller/api_vms_exec.go +++ b/internal/controller/api_vms_exec.go @@ -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) } diff --git a/internal/controller/api_vms_exec_test.go b/internal/controller/api_vms_exec_test.go index cfdb6ef..d0047c1 100644 --- a/internal/controller/api_vms_exec_test.go +++ b/internal/controller/api_vms_exec_test.go @@ -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") } diff --git a/internal/controller/controller.go b/internal/controller/controller.go index 2ef46f3..ed69543 100644 --- a/internal/controller/controller.go +++ b/internal/controller/controller.go @@ -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 { diff --git a/internal/controller/exec_sessions.go b/internal/controller/exec_sessions.go deleted file mode 100644 index e887fbb..0000000 --- a/internal/controller/exec_sessions.go +++ /dev/null @@ -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 - } -} diff --git a/internal/controller/exec_sessions_test.go b/internal/controller/exec_sessions_test.go deleted file mode 100644 index e8bfca9..0000000 --- a/internal/controller/exec_sessions_test.go +++ /dev/null @@ -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) -} diff --git a/internal/controller/option.go b/internal/controller/option.go index e29b99a..4bec9e2 100644 --- a/internal/controller/option.go +++ b/internal/controller/option.go @@ -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 diff --git a/internal/execstream/frame.go b/internal/execstream/frame.go index 9175b50..f2c5958 100644 --- a/internal/execstream/frame.go +++ b/internal/execstream/frame.go @@ -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 { diff --git a/internal/execstream/frame_test.go b/internal/execstream/frame_test.go index 3b08c73..d7bccfa 100644 --- a/internal/execstream/frame_test.go +++ b/internal/execstream/frame_test.go @@ -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, diff --git a/internal/tests/exec_test.go b/internal/tests/exec_test.go index 2c3b314..149847e 100644 --- a/internal/tests/exec_test.go +++ b/internal/tests/exec_test.go @@ -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 -} diff --git a/pkg/client/vms.go b/pkg/client/vms.go index d827209..2a71471 100644 --- a/pkg/client/vms.go +++ b/pkg/client/vms.go @@ -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) diff --git a/pkg/client/vms_test.go b/pkg/client/vms_test.go index 0e4ca23..87bab43 100644 --- a/pkg/client/vms_test.go +++ b/pkg/client/vms_test.go @@ -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"]) }