diff --git a/cli/ssh.go b/cli/ssh.go index 39a14e0750..61cb99b087 100644 --- a/cli/ssh.go +++ b/cli/ssh.go @@ -24,6 +24,7 @@ import ( "github.com/gofrs/flock" "github.com/google/uuid" "github.com/mattn/go-isatty" + "github.com/shirou/gopsutil/v4/process" "github.com/spf13/afero" gossh "golang.org/x/crypto/ssh" gosshagent "golang.org/x/crypto/ssh/agent" @@ -84,6 +85,9 @@ func (r *RootCmd) ssh() *serpent.Command { containerName string containerUser string + + // Used in tests to simulate the parent exiting. + testForcePPID int64 ) cmd := &serpent.Command{ Annotations: workspaceCommand, @@ -175,6 +179,24 @@ func (r *RootCmd) ssh() *serpent.Command { ctx, cancel := context.WithCancel(ctx) defer cancel() + // When running as a ProxyCommand (stdio mode), monitor the parent process + // and exit if it dies to avoid leaving orphaned processes. This is + // particularly important when editors like VSCode/Cursor spawn SSH + // connections and then crash or are killed - we don't want zombie + // `coder ssh` processes accumulating. + // Note: using gopsutil to check the parent process as this handles + // windows processes as well in a standard way. + if stdio { + ppid := int32(os.Getppid()) // nolint:gosec + checkParentInterval := 10 * time.Second // Arbitrary interval to not be too frequent + if testForcePPID > 0 { + ppid = int32(testForcePPID) // nolint:gosec + checkParentInterval = 100 * time.Millisecond // Shorter interval for testing + } + ctx, cancel = watchParentContext(ctx, quartz.NewReal(), ppid, process.PidExistsWithContext, checkParentInterval) + defer cancel() + } + // Prevent unnecessary logs from the stdlib from messing up the TTY. // See: https://github.com/coder/coder/issues/13144 log.SetOutput(io.Discard) @@ -775,6 +797,12 @@ func (r *RootCmd) ssh() *serpent.Command { Value: serpent.BoolOf(&forceNewTunnel), Hidden: true, }, + { + Flag: "test.force-ppid", + Description: "Override the parent process ID to simulate a different parent process. ONLY USE THIS IN TESTS.", + Value: serpent.Int64Of(&testForcePPID), + Hidden: true, + }, sshDisableAutostartOption(serpent.BoolOf(&disableAutostart)), } return cmd @@ -1662,3 +1690,33 @@ func normalizeWorkspaceInput(input string) string { return input // Fallback } } + +// watchParentContext returns a context that is canceled when the parent process +// dies. It polls using the provided clock and checks if the parent is alive +// using the provided pidExists function. +func watchParentContext(ctx context.Context, clock quartz.Clock, originalPPID int32, pidExists func(context.Context, int32) (bool, error), interval time.Duration) (context.Context, context.CancelFunc) { + ctx, cancel := context.WithCancel(ctx) // intentionally shadowed + + go func() { + ticker := clock.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + alive, err := pidExists(ctx, originalPPID) + // If we get an error checking the parent process (e.g., permission + // denied, the process is in an unknown state), we assume the parent + // is still alive to avoid disrupting the SSH connection. We only + // cancel when we definitively know the parent is gone (alive=false, err=nil). + if !alive && err == nil { + cancel() + return + } + } + } + }() + + return ctx, cancel +} diff --git a/cli/ssh_internal_test.go b/cli/ssh_internal_test.go index da6e36b96a..ee37638a66 100644 --- a/cli/ssh_internal_test.go +++ b/cli/ssh_internal_test.go @@ -312,6 +312,102 @@ type fakeCloser struct { err error } +func TestWatchParentContext(t *testing.T) { + t.Parallel() + + t.Run("CancelsWhenParentDies", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + mClock := quartz.NewMock(t) + trap := mClock.Trap().NewTicker() + defer trap.Close() + + parentAlive := true + childCtx, cancel := watchParentContext(ctx, mClock, 1234, func(context.Context, int32) (bool, error) { + return parentAlive, nil + }, testutil.WaitShort) + defer cancel() + + // Wait for the ticker to be created + trap.MustWait(ctx).MustRelease(ctx) + + // When: we simulate parent death and advance the clock + parentAlive = false + mClock.AdvanceNext() + + // Then: The context should be canceled + _ = testutil.TryReceive(ctx, t, childCtx.Done()) + }) + + t.Run("DoesNotCancelWhenParentAlive", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + mClock := quartz.NewMock(t) + trap := mClock.Trap().NewTicker() + defer trap.Close() + + childCtx, cancel := watchParentContext(ctx, mClock, 1234, func(context.Context, int32) (bool, error) { + return true, nil // Parent always alive + }, testutil.WaitShort) + defer cancel() + + // Wait for the ticker to be created + trap.MustWait(ctx).MustRelease(ctx) + + // When: we advance the clock several times with the parent alive + for range 3 { + mClock.AdvanceNext() + } + + // Then: context should not be canceled + require.NoError(t, childCtx.Err()) + }) + + t.Run("RespectsParentContext", func(t *testing.T) { + t.Parallel() + ctx, cancelParent := context.WithCancel(context.Background()) + mClock := quartz.NewMock(t) + + childCtx, cancel := watchParentContext(ctx, mClock, 1234, func(context.Context, int32) (bool, error) { + return true, nil + }, testutil.WaitShort) + defer cancel() + + // When: we cancel the parent context + cancelParent() + + // Then: The context should be canceled + require.ErrorIs(t, childCtx.Err(), context.Canceled) + }) + + t.Run("DoesNotCancelOnError", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + mClock := quartz.NewMock(t) + trap := mClock.Trap().NewTicker() + defer trap.Close() + + // Simulate an error checking parent status (e.g., permission denied). + // We should not cancel the context in this case to avoid disrupting + // the SSH connection. + childCtx, cancel := watchParentContext(ctx, mClock, 1234, func(context.Context, int32) (bool, error) { + return false, xerrors.New("permission denied") + }, testutil.WaitShort) + defer cancel() + + // Wait for the ticker to be created + trap.MustWait(ctx).MustRelease(ctx) + + // When: we advance clock several times + for range 3 { + mClock.AdvanceNext() + } + + // Context should NOT be canceled since we got an error (not a definitive "not alive") + require.NoError(t, childCtx.Err(), "context was canceled even though pidExists returned an error") + }) +} + func (c *fakeCloser) Close() error { *c.closes = append(*c.closes, c) return c.err diff --git a/cli/ssh_test.go b/cli/ssh_test.go index 33e3091674..415b3214b3 100644 --- a/cli/ssh_test.go +++ b/cli/ssh_test.go @@ -1122,6 +1122,97 @@ func TestSSH(t *testing.T) { } }) + // This test ensures that the SSH session exits when the parent process dies. + t.Run("StdioExitOnParentDeath", func(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong) + defer cancel() + + // sleepStart -> agentReady -> sessionStarted -> sleepKill -> sleepDone -> cmdDone + sleepStart := make(chan int) + agentReady := make(chan struct{}) + sessionStarted := make(chan struct{}) + sleepKill := make(chan struct{}) + sleepDone := make(chan struct{}) + + // Start a sleep process which we will pretend is the parent. + go func() { + sleepCmd := exec.Command("sleep", "infinity") + if !assert.NoError(t, sleepCmd.Start(), "failed to start sleep command") { + return + } + sleepStart <- sleepCmd.Process.Pid + defer close(sleepDone) + <-sleepKill + sleepCmd.Process.Kill() + _ = sleepCmd.Wait() + }() + + client, workspace, agentToken := setupWorkspaceForAgent(t) + go func() { + defer close(agentReady) + _ = agenttest.New(t, client.URL, agentToken) + coderdtest.NewWorkspaceAgentWaiter(t, client, workspace.ID).WaitFor(coderdtest.AgentsReady) + }() + + clientOutput, clientInput := io.Pipe() + serverOutput, serverInput := io.Pipe() + defer func() { + for _, c := range []io.Closer{clientOutput, clientInput, serverOutput, serverInput} { + _ = c.Close() + } + }() + + // Start a connection to the agent once it's ready + go func() { + <-agentReady + conn, channels, requests, err := ssh.NewClientConn(&testutil.ReaderWriterConn{ + Reader: serverOutput, + Writer: clientInput, + }, "", &ssh.ClientConfig{ + // #nosec + HostKeyCallback: ssh.InsecureIgnoreHostKey(), + }) + if !assert.NoError(t, err, "failed to create SSH client connection") { + return + } + defer conn.Close() + + sshClient := ssh.NewClient(conn, channels, requests) + defer sshClient.Close() + + session, err := sshClient.NewSession() + if !assert.NoError(t, err, "failed to create SSH session") { + return + } + close(sessionStarted) + <-sleepDone + assert.NoError(t, session.Close()) + }() + + // Wait for our "parent" process to start + sleepPid := testutil.RequireReceive(ctx, t, sleepStart) + // Wait for the agent to be ready + testutil.SoftTryReceive(ctx, t, agentReady) + inv, root := clitest.New(t, "ssh", "--stdio", workspace.Name, "--test.force-ppid", fmt.Sprintf("%d", sleepPid)) + clitest.SetupConfig(t, client, root) + inv.Stdin = clientOutput + inv.Stdout = serverInput + inv.Stderr = io.Discard + + // Start the command + clitest.Start(t, inv.WithContext(ctx)) + + // Wait for a session to be established + testutil.SoftTryReceive(ctx, t, sessionStarted) + // Now kill the fake "parent" + close(sleepKill) + // The sleep process should exit + testutil.SoftTryReceive(ctx, t, sleepDone) + // And then the command should exit. This is tracked by clitest.Start. + }) + t.Run("ForwardAgent", func(t *testing.T) { if runtime.GOOS == "windows" { t.Skip("Test not supported on windows")