Return proper exit code on ssh with TTY (#3192)

* Return proper exit code on ssh with TTY

Signed-off-by: Spike Curtis <spike@coder.com>

* Fix revive lint

Signed-off-by: Spike Curtis <spike@coder.com>

* Fix Windows exit code for missing command

Signed-off-by: Spike Curtis <spike@coder.com>

* Fix close error handling on agent TTY

Signed-off-by: Spike Curtis <spike@coder.com>
This commit is contained in:
Spike Curtis
2022-07-27 14:23:28 -05:00
committed by GitHub
parent a37e61a099
commit 36ffdce065
11 changed files with 184 additions and 27 deletions
+10
View File
@@ -29,6 +29,16 @@ type PTY interface {
Resize(height uint16, width uint16) error
}
// Process represents a process running in a PTY
type Process interface {
// Wait for the command to complete. Returned error is as for exec.Cmd.Wait()
Wait() error
// Kill the command process. Returned error is as for os.Process.Kill()
Kill() error
}
// WithFlags represents a PTY whose flags can be inspected, in particular
// to determine whether local echo is enabled.
type WithFlags interface {
+29
View File
@@ -5,6 +5,8 @@ package pty
import (
"os"
"os/exec"
"runtime"
"sync"
"github.com/creack/pty"
@@ -27,6 +29,15 @@ type otherPty struct {
pty, tty *os.File
}
type otherProcess struct {
pty *os.File
cmd *exec.Cmd
// cmdDone protects access to cmdErr: anything reading cmdErr should read from cmdDone first.
cmdDone chan any
cmdErr error
}
func (p *otherPty) Input() ReadWriter {
return ReadWriter{
Reader: p.tty,
@@ -66,3 +77,21 @@ func (p *otherPty) Close() error {
}
return nil
}
func (p *otherProcess) Wait() error {
<-p.cmdDone
return p.cmdErr
}
func (p *otherProcess) Kill() error {
return p.cmd.Process.Kill()
}
func (p *otherProcess) waitInternal() {
// The GC can garbage collect the TTY FD before the command
// has finished running. See:
// https://github.com/creack/pty/issues/127#issuecomment-932764012
p.cmdErr = p.cmd.Wait()
runtime.KeepAlive(p.pty)
close(p.cmdDone)
}
+30
View File
@@ -5,6 +5,7 @@ package pty
import (
"os"
"os/exec"
"sync"
"unsafe"
@@ -66,6 +67,13 @@ type ptyWindows struct {
closed bool
}
type windowsProcess struct {
// cmdDone protects access to cmdErr: anything reading cmdErr should read from cmdDone first.
cmdDone chan any
cmdErr error
proc *os.Process
}
func (p *ptyWindows) Output() ReadWriter {
return ReadWriter{
Reader: p.outputRead,
@@ -111,3 +119,25 @@ func (p *ptyWindows) Close() error {
return nil
}
func (p *windowsProcess) waitInternal() {
defer close(p.cmdDone)
state, err := p.proc.Wait()
if err != nil {
p.cmdErr = err
return
}
if !state.Success() {
p.cmdErr = &exec.ExitError{ProcessState: state}
return
}
}
func (p *windowsProcess) Wait() error {
<-p.cmdDone
return p.cmdErr
}
func (p *windowsProcess) Kill() error {
return p.proc.Kill()
}
+1 -2
View File
@@ -5,7 +5,6 @@ import (
"bytes"
"context"
"io"
"os"
"os/exec"
"runtime"
"strings"
@@ -27,7 +26,7 @@ func New(t *testing.T) *PTY {
return create(t, ptty, "cmd")
}
func Start(t *testing.T, cmd *exec.Cmd) (*PTY, *os.Process) {
func Start(t *testing.T, cmd *exec.Cmd) (*PTY, pty.Process) {
ptty, ps, err := pty.Start(cmd)
require.NoError(t, err)
return create(t, ptty, cmd.Args[0]), ps
+3 -2
View File
@@ -1,10 +1,11 @@
package pty
import (
"os"
"os/exec"
)
func Start(cmd *exec.Cmd) (PTY, *os.Process, error) {
// Start the command in a TTY. The calling code must not use cmd after passing it to the PTY, and
// instead rely on the returned Process to manage the command/process.
func Start(cmd *exec.Cmd) (PTY, Process, error) {
return startPty(cmd)
}
+8 -10
View File
@@ -4,7 +4,6 @@
package pty
import (
"os"
"os/exec"
"runtime"
"strings"
@@ -14,7 +13,7 @@ import (
"golang.org/x/xerrors"
)
func startPty(cmd *exec.Cmd) (PTY, *os.Process, error) {
func startPty(cmd *exec.Cmd) (PTY, Process, error) {
ptty, tty, err := pty.Open()
if err != nil {
return nil, nil, xerrors.Errorf("open: %w", err)
@@ -37,16 +36,15 @@ func startPty(cmd *exec.Cmd) (PTY, *os.Process, error) {
}
return nil, nil, xerrors.Errorf("start: %w", err)
}
go func() {
// The GC can garbage collect the TTY FD before the command
// has finished running. See:
// https://github.com/creack/pty/issues/127#issuecomment-932764012
_ = cmd.Wait()
runtime.KeepAlive(ptty)
}()
oPty := &otherPty{
pty: ptty,
tty: tty,
}
return oPty, cmd.Process, nil
oProcess := &otherProcess{
pty: ptty,
cmd: cmd,
cmdDone: make(chan any),
}
go oProcess.waitInternal()
return oPty, oProcess, nil
}
+18 -1
View File
@@ -7,6 +7,10 @@ import (
"os/exec"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"go.uber.org/goleak"
"github.com/coder/coder/pty/ptytest"
@@ -20,7 +24,20 @@ func TestStart(t *testing.T) {
t.Parallel()
t.Run("Echo", func(t *testing.T) {
t.Parallel()
pty, _ := ptytest.Start(t, exec.Command("echo", "test"))
pty, ps := ptytest.Start(t, exec.Command("echo", "test"))
pty.ExpectMatch("test")
err := ps.Wait()
require.NoError(t, err)
})
t.Run("Kill", func(t *testing.T) {
t.Parallel()
_, ps := ptytest.Start(t, exec.Command("sleep", "30"))
err := ps.Kill()
assert.NoError(t, err)
err = ps.Wait()
var exitErr *exec.ExitError
require.True(t, xerrors.As(err, &exitErr))
assert.NotEqual(t, 0, exitErr.ExitCode())
})
}
+7 -2
View File
@@ -16,7 +16,7 @@ import (
// Allocates a PTY and starts the specified command attached to it.
// See: https://docs.microsoft.com/en-us/windows/console/creating-a-pseudoconsole-session#creating-the-hosted-process
func startPty(cmd *exec.Cmd) (PTY, *os.Process, error) {
func startPty(cmd *exec.Cmd) (PTY, Process, error) {
fullPath, err := exec.LookPath(cmd.Path)
if err != nil {
return nil, nil, err
@@ -83,7 +83,12 @@ func startPty(cmd *exec.Cmd) (PTY, *os.Process, error) {
if err != nil {
return nil, nil, xerrors.Errorf("find process %d: %w", processInfo.ProcessId, err)
}
return pty, process, nil
wp := &windowsProcess{
cmdDone: make(chan any),
proc: process,
}
go wp.waitInternal()
return pty, wp, nil
}
// Taken from: https://github.com/microsoft/hcsshim/blob/7fbdca16f91de8792371ba22b7305bf4ca84170a/internal/exec/exec.go#L476
+15 -1
View File
@@ -8,8 +8,10 @@ import (
"testing"
"github.com/coder/coder/pty/ptytest"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"golang.org/x/xerrors"
)
func TestMain(m *testing.M) {
@@ -20,8 +22,10 @@ func TestStart(t *testing.T) {
t.Parallel()
t.Run("Echo", func(t *testing.T) {
t.Parallel()
pty, _ := ptytest.Start(t, exec.Command("cmd.exe", "/c", "echo", "test"))
pty, ps := ptytest.Start(t, exec.Command("cmd.exe", "/c", "echo", "test"))
pty.ExpectMatch("test")
err := ps.Wait()
require.NoError(t, err)
})
t.Run("Resize", func(t *testing.T) {
t.Parallel()
@@ -29,4 +33,14 @@ func TestStart(t *testing.T) {
err := pty.Resize(100, 50)
require.NoError(t, err)
})
t.Run("Kill", func(t *testing.T) {
t.Parallel()
_, ps := ptytest.Start(t, exec.Command("cmd.exe"))
err := ps.Kill()
assert.NoError(t, err)
err = ps.Wait()
var exitErr *exec.ExitError
require.True(t, xerrors.As(err, &exitErr))
assert.NotEqual(t, 0, exitErr.ExitCode())
})
}