mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: pty.Start respects context on Windows too (#7373)
* fix: pty.Start respects context on Windows too Signed-off-by: Spike Curtis <spike@coder.com> * Fix windows imports; rename ToExec -> AsExec Signed-off-by: Spike Curtis <spike@coder.com> * Fix import in windows test Signed-off-by: Spike Curtis <spike@coder.com> --------- Signed-off-by: Spike Curtis <spike@coder.com>
This commit is contained in:
+4
-12
@@ -216,11 +216,12 @@ func (a *agent) collectMetadata(ctx context.Context, md codersdk.WorkspaceAgentM
|
||||
// if it can guarantee the clocks are synchronized.
|
||||
CollectedAt: time.Now(),
|
||||
}
|
||||
cmd, err := a.sshServer.CreateCommand(ctx, md.Script, nil)
|
||||
cmdPty, err := a.sshServer.CreateCommand(ctx, md.Script, nil)
|
||||
if err != nil {
|
||||
result.Error = fmt.Sprintf("create cmd: %+v", err)
|
||||
return result
|
||||
}
|
||||
cmd := cmdPty.AsExec()
|
||||
|
||||
cmd.Stdout = &out
|
||||
cmd.Stderr = &out
|
||||
@@ -842,10 +843,11 @@ func (a *agent) runScript(ctx context.Context, lifecycle, script string) error {
|
||||
}()
|
||||
}
|
||||
|
||||
cmd, err := a.sshServer.CreateCommand(ctx, script, nil)
|
||||
cmdPty, err := a.sshServer.CreateCommand(ctx, script, nil)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("create command: %w", err)
|
||||
}
|
||||
cmd := cmdPty.AsExec()
|
||||
cmd.Stdout = writer
|
||||
cmd.Stderr = writer
|
||||
err = cmd.Run()
|
||||
@@ -1044,16 +1046,6 @@ func (a *agent) handleReconnectingPTY(ctx context.Context, logger slog.Logger, m
|
||||
circularBuffer: circularBuffer,
|
||||
}
|
||||
a.reconnectingPTYs.Store(msg.ID, rpty)
|
||||
go func() {
|
||||
// CommandContext isn't respected for Windows PTYs right now,
|
||||
// so we need to manually track the lifecycle.
|
||||
// When the context has been completed either:
|
||||
// 1. The timeout completed.
|
||||
// 2. The parent context was canceled.
|
||||
<-ctx.Done()
|
||||
logger.Debug(ctx, "context done", slog.Error(ctx.Err()))
|
||||
_ = process.Kill()
|
||||
}()
|
||||
// We don't need to separately monitor for the process exiting.
|
||||
// When it exits, our ptty.OutputReader() will return EOF after
|
||||
// reading all process output.
|
||||
|
||||
+1
-2
@@ -12,7 +12,6 @@ import (
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"os"
|
||||
"os/exec"
|
||||
"os/user"
|
||||
"path"
|
||||
"path/filepath"
|
||||
@@ -1697,7 +1696,7 @@ func setupSSHCommand(t *testing.T, beforeArgs []string, afterArgs []string) (*pt
|
||||
"host",
|
||||
)
|
||||
args = append(args, afterArgs...)
|
||||
cmd := exec.Command("ssh", args...)
|
||||
cmd := pty.Command("ssh", args...)
|
||||
return ptytest.Start(t, cmd)
|
||||
}
|
||||
|
||||
|
||||
@@ -255,7 +255,7 @@ func (s *Server) sessionStart(session ssh.Session, extraEnv []string) (retErr er
|
||||
if isPty {
|
||||
return s.startPTYSession(session, cmd, sshPty, windowSize)
|
||||
}
|
||||
return startNonPTYSession(session, cmd)
|
||||
return startNonPTYSession(session, cmd.AsExec())
|
||||
}
|
||||
|
||||
func startNonPTYSession(session ssh.Session, cmd *exec.Cmd) error {
|
||||
@@ -287,7 +287,7 @@ type ptySession interface {
|
||||
RawCommand() string
|
||||
}
|
||||
|
||||
func (s *Server) startPTYSession(session ptySession, cmd *exec.Cmd, sshPty ssh.Pty, windowSize <-chan ssh.Window) (retErr error) {
|
||||
func (s *Server) startPTYSession(session ptySession, cmd *pty.Cmd, sshPty ssh.Pty, windowSize <-chan ssh.Window) (retErr error) {
|
||||
ctx := session.Context()
|
||||
// Disable minimal PTY emulation set by gliderlabs/ssh (NL-to-CRNL).
|
||||
// See https://github.com/coder/coder/issues/3371.
|
||||
@@ -413,7 +413,7 @@ func (s *Server) sftpHandler(session ssh.Session) {
|
||||
// CreateCommand processes raw command input with OpenSSH-like behavior.
|
||||
// If the script provided is empty, it will default to the users shell.
|
||||
// This injects environment variables specified by the user at launch too.
|
||||
func (s *Server) CreateCommand(ctx context.Context, script string, env []string) (*exec.Cmd, error) {
|
||||
func (s *Server) CreateCommand(ctx context.Context, script string, env []string) (*pty.Cmd, error) {
|
||||
currentUser, err := user.Current()
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("get current user: %w", err)
|
||||
@@ -449,7 +449,7 @@ func (s *Server) CreateCommand(ctx context.Context, script string, env []string)
|
||||
}
|
||||
}
|
||||
|
||||
cmd := exec.CommandContext(ctx, shell, args...)
|
||||
cmd := pty.CommandContext(ctx, shell, args...)
|
||||
cmd.Dir = manifest.Directory
|
||||
|
||||
// If the metadata directory doesn't exist, we run the command
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"os/exec"
|
||||
"testing"
|
||||
|
||||
gliderssh "github.com/gliderlabs/ssh"
|
||||
@@ -15,6 +14,7 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/pty"
|
||||
"github.com/coder/coder/testutil"
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
@@ -52,7 +52,7 @@ func Test_sessionStart_orphan(t *testing.T) {
|
||||
close(windowSize)
|
||||
// the command gets the session context so that Go will terminate it when
|
||||
// the session expires.
|
||||
cmd := exec.CommandContext(sessionCtx, "sh", "-c", longScript)
|
||||
cmd := pty.CommandContext(sessionCtx, "sh", "-c", longScript)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
|
||||
Reference in New Issue
Block a user