diff --git a/internal/remote/batch_tty_test.go b/internal/remote/batch_tty_test.go new file mode 100644 index 0000000..a87c27d --- /dev/null +++ b/internal/remote/batch_tty_test.go @@ -0,0 +1,81 @@ +package remote + +import ( + "bytes" + "context" + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "syscall" + "testing" + "time" + + "github.com/creack/pty" +) + +// a batch ssh run must not reach the terminal: tssh ignores BatchMode and +// reads host key answers from /dev/tty, stealing keys typed into the picker. +// the check re-runs this test under a pty it owns, with a fake tssh that +// reports whether it could open /dev/tty +func TestBatchSSHHasNoTerminal(t *testing.T) { + if os.Getenv("TOBI_TEST_TTY_PROBE") == "1" { + _, err := os.Open("/dev/tty") + out, _ := sshRun(context.Background(), "host", nil, sshOpts{cmd: "true", batch: true, timeout: "1"}) + fmt.Printf("helper:%v ssh:%s", err == nil, strings.TrimSpace(out)) + return + } + dir := t.TempDir() + fake := "#!/bin/sh\n(exec 3/dev/null && echo tty || echo notty\n" + if err := os.WriteFile(filepath.Join(dir, "tssh"), []byte(fake), 0755); err != nil { + t.Fatal(err) + } + master, tty, err := pty.Open() + if err != nil { + t.Fatal(err) + } + defer master.Close() + defer tty.Close() + var out bytes.Buffer + c := exec.Command(os.Args[0], "-test.run=^TestBatchSSHHasNoTerminal$") + c.Env = append(os.Environ(), "TOBI_TEST_TTY_PROBE=1", "PATH="+dir+":/usr/bin:/bin") + c.Stdin, c.Stdout, c.Stderr = tty, &out, &out + c.SysProcAttr = &syscall.SysProcAttr{Setsid: true, Setctty: true} + if err := c.Run(); err != nil { + t.Fatalf("%v: %s", err, out.String()) + } + if !strings.Contains(out.String(), "helper:true ssh:notty") { + t.Fatalf("want the helper on a terminal and ssh off it, got %q", out.String()) + } +} + +// a probe that times out takes its ProxyCommand with it: killing only ssh +// orphaned a tailscale nc per unreachable host +func TestBatchSSHTimeoutKillsItsChildren(t *testing.T) { + dir := t.TempDir() + pidFile := filepath.Join(dir, "child.pid") + fake := "#!/bin/sh\nsleep 60 &\necho $! > " + pidFile + "\nwait\n" + if err := os.WriteFile(filepath.Join(dir, "tssh"), []byte(fake), 0755); err != nil { + t.Fatal(err) + } + t.Setenv("PATH", dir+":/usr/bin:/bin") + ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond) + defer cancel() + if _, err := sshRun(ctx, "host", nil, sshOpts{cmd: "true", batch: true, timeout: "1"}); !errors.Is(err, ErrUnreachable) { + t.Fatalf("a timed out run should be unreachable, got %v", err) + } + data, err := os.ReadFile(pidFile) + if err != nil { + t.Fatal(err) + } + pid, _ := strconv.Atoi(strings.TrimSpace(string(data))) + for deadline := time.Now().Add(time.Second); syscall.Kill(pid, 0) == nil; time.Sleep(10 * time.Millisecond) { + if time.Now().After(deadline) { + _ = syscall.Kill(pid, syscall.SIGKILL) + t.Fatalf("the fake ssh's child %d outlived the timeout", pid) + } + } +} diff --git a/internal/remote/remote.go b/internal/remote/remote.go index 25e1b55..d46e312 100644 --- a/internal/remote/remote.go +++ b/internal/remote/remote.go @@ -18,6 +18,7 @@ import ( "strconv" "strings" "sync" + "syscall" "time" "github.com/klauspost/compress/zstd" @@ -203,6 +204,14 @@ func sshCommand(ctx context.Context, dest string, opts sshOpts) *exec.Cmd { } c := exec.CommandContext(ctx, bin, append(args, dest, opts.cmd)...) c.WaitDelay = 500 * time.Millisecond + if opts.batch { + // no controlling terminal: tssh ignores BatchMode and prompts for + // unknown host keys on /dev/tty, eating the picker's keystrokes. + // the new session is also a process group, so a timeout takes the + // ProxyCommand down with ssh instead of orphaning it + c.SysProcAttr = &syscall.SysProcAttr{Setsid: true} + c.Cancel = func() error { return syscall.Kill(-c.Process.Pid, syscall.SIGKILL) } + } return c }