diff --git a/AGENTS.md b/AGENTS.md index 689d99e..525cb9a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -13,16 +13,17 @@ Module path: `github.com/tphakala/et-go`. Go version: see `go.mod`. ## Layout -| Path | Role | -| ----------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `cmd/et` | Entry point and flag parsing. | -| `internal/bootstrap` | Runs the system ssh to start etterminal; validates the config, parses the IDPASSKEY credentials it prints and redacts the passkey. | -| `internal/etcp` | Reliable, ordered, encrypted packet connection over replaceable TCP links: replay ring, recover exchange, liveness probes, reconnect backoff, write backpressure. | -| `internal/etservertest` | Test-only fake etserver (written independently from upstream semantics) and an in-memory `net.Pipe` network with cut and refuse controls, for synctest-driven etcp tests. | -| `internal/protocol` | Wire messages generated from upstream's `.proto` files (opaque API), plus `Header`, `Version` and `Packet`. Regenerate with `go generate ./internal/protocol` (needs protoc 3.21.12). | -| `internal/seal` | One direction of the libsodium-compatible encrypted stream: secretbox with a counter nonce. | -| `internal/wire` | Handshake message framing, stream frame framing, packet layout and size limits. | -| `rules/` | ruleguard matchers used by golangci-lint (build tag `ruleguard`). | +| Path | Role | +| ----------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `cmd/et` | Entry point and flag parsing. | +| `internal/bootstrap` | Runs the system ssh to start etterminal; validates the config, parses the IDPASSKEY credentials it prints and redacts the passkey. | +| `internal/console` | Local console in raw mode per platform: raw VT input as UTF-8, output, window size and resize events, and a Close that unblocks a pending Read. Windows uses UTF-16 carry codecs shared with every platform's tests. | +| `internal/etcp` | Reliable, ordered, encrypted packet connection over replaceable TCP links: replay ring, recover exchange, liveness probes, reconnect backoff, write backpressure. | +| `internal/etservertest` | Test-only fake etserver (written independently from upstream semantics) and an in-memory `net.Pipe` network with cut and refuse controls, for synctest-driven etcp tests. | +| `internal/protocol` | Wire messages generated from upstream's `.proto` files (opaque API), plus `Header`, `Version` and `Packet`. Regenerate with `go generate ./internal/protocol` (needs protoc 3.21.12). | +| `internal/seal` | One direction of the libsodium-compatible encrypted stream: secretbox with a counter nonce. | +| `internal/wire` | Handshake message framing, stream frame framing, packet layout and size limits. | +| `rules/` | ruleguard matchers used by golangci-lint (build tag `ruleguard`). | Update this table when a package is added. diff --git a/go.mod b/go.mod index 81574c8..3e3f54c 100644 --- a/go.mod +++ b/go.mod @@ -5,13 +5,14 @@ go 1.27.0 require ( github.com/quasilyte/go-ruleguard/dsl v0.3.23 golang.org/x/crypto v0.57.0 + golang.org/x/sys v0.48.0 + golang.org/x/term v0.46.0 google.golang.org/protobuf v1.36.12 ) require ( golang.org/x/mod v0.41.0 // indirect golang.org/x/sync v0.23.0 // indirect - golang.org/x/sys v0.48.0 // indirect golang.org/x/tools v0.50.0 // indirect ) diff --git a/go.sum b/go.sum index 45f5f5a..fb6da53 100644 --- a/go.sum +++ b/go.sum @@ -10,6 +10,8 @@ golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0= golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= +golang.org/x/term v0.46.0 h1:3+OXuTbaKDgwk8jTi3aSLHRlmWqHEUDUtxnbFigO4YE= +golang.org/x/term v0.46.0/go.mod h1:+K02xbkittuwc0Am4abfA3Fc+XRGXkvBXNO88NCXPoc= golang.org/x/tools v0.50.0 h1:c2ifzfcuY7L90lZ2aKd8S4K2NpASF08SZx9ZuJkHmSU= golang.org/x/tools v0.50.0/go.mod h1:7ulVMw3831Mwi5EZD6RomGyffr4VFjuNYXf2BbCEAV0= google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= diff --git a/internal/console/console.go b/internal/console/console.go new file mode 100644 index 0000000..eab80a7 --- /dev/null +++ b/internal/console/console.go @@ -0,0 +1,22 @@ +// Package console is the local console or terminal that et puts into raw mode: +// raw VT input as UTF-8 bytes, remote output written back, the window size, +// and resize events. Each platform implements the same concrete Console type +// (console_unix.go, console_windows.go). +// +// On Unix, et must run as the foreground job of its controlling terminal. +// A background job that reads or changes the terminal is stopped by job +// control (SIGTTIN, SIGTTOU) until it is brought to the foreground. +package console + +import "errors" + +// Size is the terminal size in character cells, plus pixels when the +// platform reports them (0 otherwise). +type Size struct { + Rows, Cols int + Width, Height int +} + +// ErrNotTerminal is returned by Open when stdin or stdout is not a terminal, +// or when the process has no controlling terminal to open. +var ErrNotTerminal = errors.New("console: stdin or stdout is not a terminal") diff --git a/internal/console/console_linux_test.go b/internal/console/console_linux_test.go new file mode 100644 index 0000000..1589c5d --- /dev/null +++ b/internal/console/console_linux_test.go @@ -0,0 +1,897 @@ +package console + +import ( + "context" + "errors" + "io" + "io/fs" + "os" + "path/filepath" + "testing" + "time" + + "golang.org/x/sys/unix" + "golang.org/x/term" +) + +// waitTimeout bounds every wait on a real file descriptor; synctest cannot +// drive real terminal I/O. +const waitTimeout = 5 * time.Second + +func mustConsole(t *testing.T, tty *os.File) *Console { + t.Helper() + c, err := newConsole(tty) + if err != nil { + t.Fatalf("newConsole: %v", err) + } + return c +} + +func TestOpenRejectsNonTerminal(t *testing.T) { + _, slave := openPTY(t) + r, w, err := os.Pipe() + if err != nil { + t.Fatalf("pipe: %v", err) + } + closeAtEnd(t, r) + closeAtEnd(t, w) + + tests := []struct { + name string + stdin, stdout *os.File + }{ + {"both pipes", r, w}, + {"stdin pipe", r, slave}, + {"stdout pipe", slave, w}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c, err := open(slave.Name(), tt.stdin, tt.stdout) + if !errors.Is(err, ErrNotTerminal) { + if c != nil { + _ = c.Close() + } + t.Fatalf("open error = %v, want ErrNotTerminal", err) + } + }) + } +} + +// openThroughPath builds the Console the way Open does, reopening the +// terminal by path. open does not pass O_NOCTTY, so a session leader +// without a controlling terminal would acquire the test pty as one; the +// helper skips in that case rather than change the process's terminal. +func openThroughPath(t *testing.T, slave *os.File) *Console { + t.Helper() + if sid, err := unix.Getsid(0); err == nil && sid == os.Getpid() { + t.Skip("the test process is a session leader; opening the pty without O_NOCTTY could make it the controlling terminal") + } + c, err := open(slave.Name(), slave, slave) + if err != nil { + t.Fatalf("open: %v", err) + } + return c +} + +// TestOpenRejectsNonTerminalTTYPath covers a tty path that opens but is not +// a terminal: open must report ErrNotTerminal, not a bare ioctl error. +func TestOpenRejectsNonTerminalTTYPath(t *testing.T) { + _, slave := openPTY(t) + regular := filepath.Join(t.TempDir(), "not-a-tty") + if err := os.WriteFile(regular, nil, 0o600); err != nil { + t.Fatalf("create regular file: %v", err) + } + c, err := open(regular, slave, slave) + if !errors.Is(err, ErrNotTerminal) { + if c != nil { + _ = c.Close() + } + t.Fatalf("open(regular file) error = %v, want ErrNotTerminal", err) + } +} + +func TestOpenWithoutControllingTerminal(t *testing.T) { + _, slave := openPTY(t) + _, err := open("/nonexistent/et-go-tty", slave, slave) + if !errors.Is(err, ErrNotTerminal) { + t.Fatalf("open error = %v, want ErrNotTerminal", err) + } + if !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("open error = %v, want the underlying open error kept", err) + } +} + +func TestOpenAcceptsTerminal(t *testing.T) { + _, slave := openPTY(t) + c := openThroughPath(t, slave) + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } +} + +func TestReadWrite(t *testing.T) { + master, slave := openPTY(t) + c := mustConsole(t, slave) + if _, err := c.MakeRaw(); err != nil { + t.Fatalf("MakeRaw: %v", err) + } + + if _, err := master.WriteString("\x1b[A"); err != nil { + t.Fatalf("write to master: %v", err) + } + got := make([]byte, 16) + n, err := c.Read(got) + if err != nil || string(got[:n]) != "\x1b[A" { + t.Fatalf("Read = %q, %v; want the arrow-up sequence unchanged", got[:n], err) + } + + if _, err := c.Write([]byte("out")); err != nil { + t.Fatalf("Write: %v", err) + } + out := make([]byte, 3) + if _, err := io.ReadFull(master, out); err != nil || string(out) != "out" { + t.Fatalf("master read %q, %v; want %q", out, err, "out") + } +} + +func TestMakeRawAndRestore(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + before := termios(t, slave) + if before.Lflag&(unix.ICANON|unix.ECHO) == 0 { + t.Fatal("a fresh pty should start in cooked mode") + } + + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if raw := termios(t, slave); raw.Lflag&(unix.ICANON|unix.ECHO) != 0 { + t.Fatalf("after MakeRaw Lflag = %#x, want ICANON and ECHO cleared", raw.Lflag) + } + + if err := restore(); err != nil { + t.Fatalf("restore: %v", err) + } + if err := restore(); err != nil { + t.Fatalf("second restore must be a no-op, got %v", err) + } + if after := termios(t, slave); after.Lflag != before.Lflag { + t.Fatalf("after restore Lflag = %#x, want %#x", after.Lflag, before.Lflag) + } +} + +// TestRestoreReturnsToOpenBaseline models ssh being interrupted at a +// password prompt between Open and MakeRaw, leaving echo off: restore must +// return to the state recorded at Open, not to the degraded one. +func TestRestoreReturnsToOpenBaseline(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + baseline := termios(t, slave) + + degraded := *baseline + degraded.Lflag &^= unix.ECHO + if err := controlFile(slave, func(fd int) error { + return unix.IoctlSetTermios(fd, unix.TCSETS, °raded) + }); err != nil { + t.Fatalf("degrade termios: %v", err) + } + + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := restore(); err != nil { + t.Fatalf("restore: %v", err) + } + if after := termios(t, slave); after.Lflag != baseline.Lflag { + t.Fatalf("after restore Lflag = %#x, want the Open baseline %#x (echo on)", after.Lflag, baseline.Lflag) + } + + // restore is repeatable: after the mode changes again, a second call + // returns to the baseline too. + setTermios(t, slave, °raded) + if err := restore(); err != nil { + t.Fatalf("second restore: %v", err) + } + if after := termios(t, slave); after.Lflag != baseline.Lflag { + t.Fatalf("after a second restore Lflag = %#x, want the Open baseline %#x", after.Lflag, baseline.Lflag) + } +} + +// probeTTY opens a second descriptor on the terminal behind slave, to +// inspect and change it after the Console closed its own. +func probeTTY(t *testing.T, slave *os.File) *os.File { + t.Helper() + probe, err := os.OpenFile(slave.Name(), os.O_RDWR|unix.O_NOCTTY, 0) + if err != nil { + t.Fatalf("reopen slave: %v", err) + } + closeAtEnd(t, probe) + return probe +} + +func setTermios(t *testing.T, f *os.File, tio *unix.Termios) { + t.Helper() + if err := controlFile(f, func(fd int) error { + return unix.IoctlSetTermios(fd, unix.TCSETS, tio) + }); err != nil { + t.Fatalf("tcsets: %v", err) + } +} + +// withEchoOff clears ECHO on f and returns the termios it had before. +func withEchoOff(t *testing.T, f *os.File) *unix.Termios { + t.Helper() + before := termios(t, f) + degraded := *before + degraded.Lflag &^= unix.ECHO + setTermios(t, f, °raded) + return before +} + +// TestCloseRestoresBaselineWithoutMakeRaw models ssh interrupted at a +// password prompt with echo off and et closing before MakeRaw ran: Close +// must still return the terminal to the Open baseline. +func TestCloseRestoresBaselineWithoutMakeRaw(t *testing.T) { + _, slave := openPTY(t) + probe := probeTTY(t, slave) + c := mustConsole(t, slave) + baseline := withEchoOff(t, probe) + + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if after := termios(t, probe); after.Lflag != baseline.Lflag { + t.Fatalf("after Close Lflag = %#x, want the Open baseline %#x (echo on)", after.Lflag, baseline.Lflag) + } +} + +// TestCloseTwice checks that a Close that starts while another is still +// restoring waits for it and returns nil, and that a Close after both +// returns nil at once. +func TestCloseTwice(t *testing.T) { + _, slave := openPTY(t) + probe := probeTTY(t, slave) + c := mustConsole(t, slave) + baseline := withEchoOff(t, probe) // so the first Close has a mode to set + + entered := make(chan struct{}) + release := make(chan struct{}) + set := c.setState + c.setState = func(fd int, s *term.State) error { + close(entered) + <-release + return set(fd, s) + } + + first := make(chan error, 1) + go func() { first <- c.Close() }() + select { + case <-entered: + case <-time.After(waitTimeout): + t.Fatal("the first Close never reached the terminal restore") + } + + second := make(chan error, 1) + go func() { second <- c.Close() }() + select { + case err := <-second: + close(release) + t.Fatalf("a second Close returned (%v) while the first was still restoring", err) + case <-time.After(200 * time.Millisecond): + } + close(release) + + for name, ch := range map[string]chan error{"first": first, "second": second} { + select { + case err := <-ch: + if err != nil { + t.Fatalf("%s Close = %v, want nil", name, err) + } + case <-time.After(waitTimeout): + t.Fatalf("%s Close never returned", name) + } + } + if after := termios(t, probe); after.Lflag != baseline.Lflag { + t.Fatalf("after Close Lflag = %#x, want the Open baseline %#x", after.Lflag, baseline.Lflag) + } + if err := c.Close(); err != nil { + t.Fatalf("Close after Close = %v, want nil", err) + } +} + +func TestWriteAfterClose(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if n, err := c.Write([]byte("x")); n != 0 || !errors.Is(err, os.ErrClosed) { + t.Fatalf("Write after Close = %d, %v; want 0, os.ErrClosed", n, err) + } +} + +func TestMakeRawAfterClose(t *testing.T) { + _, slave := openPTY(t) + probe := probeTTY(t, slave) + c := mustConsole(t, slave) + before := termios(t, probe) + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + restore, err := c.MakeRaw() + if restore != nil || !errors.Is(err, os.ErrClosed) { + t.Fatalf("MakeRaw after Close = (restore set: %t), %v; want nil, os.ErrClosed", restore != nil, err) + } + if after := termios(t, probe); after.Lflag != before.Lflag { + t.Fatalf("MakeRaw after Close changed Lflag to %#x, want %#x", after.Lflag, before.Lflag) + } +} + +// TestRestoreAfterClose checks that a restore func run after Close returns +// nil and leaves the terminal alone: Close already restored it. +func TestRestoreAfterClose(t *testing.T) { + _, slave := openPTY(t) + probe := probeTTY(t, slave) + c := mustConsole(t, slave) + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + // Another program now owns the terminal and turned echo off. + withEchoOff(t, probe) + changed := termios(t, probe) + + if err := restore(); err != nil { + t.Fatalf("restore after Close = %v, want nil", err) + } + if after := termios(t, probe); after.Lflag != changed.Lflag { + t.Fatalf("restore after Close changed Lflag to %#x, want it left at %#x", after.Lflag, changed.Lflag) + } +} + +// TestCloseRestoresAfterSecondMakeRaw covers raw mode entered again after a +// restore: Close must still return to the baseline. +func TestCloseRestoresAfterSecondMakeRaw(t *testing.T) { + _, slave := openPTY(t) + probe := probeTTY(t, slave) + c := mustConsole(t, slave) + baseline := termios(t, probe) + + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := restore(); err != nil { + t.Fatalf("restore: %v", err) + } + if _, err := c.MakeRaw(); err != nil { + t.Fatalf("second MakeRaw: %v", err) + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if after := termios(t, probe); after.Lflag != baseline.Lflag { + t.Fatalf("after Close Lflag = %#x, want the Open baseline %#x", after.Lflag, baseline.Lflag) + } +} + +func TestSizeAfterClose(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := c.Size(); !errors.Is(err, os.ErrClosed) { + t.Fatalf("Size after Close = %v, want an error wrapping os.ErrClosed", err) + } +} + +// TestRestoreSkipsUnchangedMode checks that restore and Close do not set a +// terminal that is already at the baseline. Setting it from a background +// job would stop et with SIGTTOU, but the test pty is opened with O_NOCTTY +// and is not the test's controlling terminal, so job control never applies +// here; the set calls are counted through c.setState instead. +func TestRestoreSkipsUnchangedMode(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + sets := 0 + set := c.setState + c.setState = func(fd int, s *term.State) error { + sets++ + return set(fd, s) + } + + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := restore(); err != nil { + t.Fatalf("restore: %v", err) + } + if sets != 1 { + t.Fatalf("restore from raw mode made %d set calls, want 1", sets) + } + if err := restore(); err != nil { + t.Fatalf("second restore: %v", err) + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if sets != 1 { + t.Fatalf("restore and Close at the baseline made %d more set calls, want 0", sets-1) + } +} + +// TestCloseReportsRestoreError checks that a failed restore reaches the +// caller of the first Close, and that a later Close returns nil. +func TestCloseReportsRestoreError(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + withEchoOff(t, slave) // so Close has a mode to set + errSet := errors.New("set terminal state failed") + c.setState = func(int, *term.State) error { return errSet } + + if err := c.Close(); !errors.Is(err, errSet) { + t.Fatalf("Close = %v, want an error wrapping the failed restore", err) + } + // The failed restore must not stop Close from closing the terminal. + if _, err := c.Write([]byte("x")); !errors.Is(err, os.ErrClosed) { + t.Fatalf("Write after a Close whose restore failed = %v, want an error wrapping os.ErrClosed", err) + } + if err := c.Close(); err != nil { + t.Fatalf("second Close = %v, want nil", err) + } +} + +// TestSizeConcurrentWithClose runs Size in a loop while Close runs, many +// times over: every error Size returns must wrap os.ErrClosed. Under -race +// it catches a Size that reads the closed flag without the lock (a data +// race on it). A Size that takes the lock for the check but drops it +// before its ioctl is caught only when the ioctl happens to land after the +// descriptor closed ("use of closed file"), so that shape is caught +// probabilistically, not on every run. +func TestSizeConcurrentWithClose(t *testing.T) { + for i := range 50 { + _, slave := openPTY(t) + c := mustConsole(t, slave) + started := make(chan struct{}) + sizeErr := make(chan error, 1) + go func() { + close(started) + for { + if _, err := c.Size(); err != nil { + sizeErr <- err + return + } + } + }() + <-started + if err := c.Close(); err != nil { + t.Fatalf("round %d: Close = %v, want nil", i, err) + } + select { + case err := <-sizeErr: + if !errors.Is(err, os.ErrClosed) { + t.Fatalf("round %d: Size racing Close = %v, want an error wrapping os.ErrClosed", i, err) + } + case <-time.After(waitTimeout): + t.Fatalf("round %d: Size kept succeeding after Close", i) + } + } +} + +// TestRestoreConcurrentWithClose runs a restore func and Close at the same +// time, many times over so the two land in both orders: restore must run +// under the lock Close holds, so it either restores before Close or sees +// the Console closed and returns nil, never touching a closed descriptor. +func TestRestoreConcurrentWithClose(t *testing.T) { + for i := range 50 { + _, slave := openPTY(t) + c := mustConsole(t, slave) + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("round %d: MakeRaw: %v", i, err) + } + started := make(chan struct{}) + restored := make(chan error, 1) + go func() { + close(started) + restored <- restore() + }() + <-started // start Close as restore starts, so the two overlap + if err := c.Close(); err != nil { + t.Fatalf("round %d: Close = %v, want nil", i, err) + } + select { + case err := <-restored: + if err != nil { + t.Fatalf("round %d: restore concurrent with Close = %v, want nil", i, err) + } + case <-time.After(waitTimeout): + t.Fatalf("round %d: restore never returned", i) + } + } +} + +func TestCloseRestoresAndUnblocksRead(t *testing.T) { + _, slave := openPTY(t) + // A second descriptor on the same terminal to inspect it after Close. + probe := probeTTY(t, slave) + before := termios(t, probe) + + c := openThroughPath(t, slave) + if _, err := c.MakeRaw(); err != nil { + t.Fatalf("MakeRaw: %v", err) + } + // SetReadDeadline succeeds only on a pollable fd (a non-pollable one + // returns os.ErrNoDeadline). A pollable fd is what lets Close unblock a + // pending Read: the Read parks in the poller, which Close wakes. + if err := c.tty.SetReadDeadline(time.Time{}); err != nil { + t.Fatalf("SetReadDeadline: %v, want the tty fd to be pollable", err) + } + readErr := make(chan error, 1) + go func() { + _, err := c.Read(make([]byte, 8)) + readErr <- err + }() + + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + select { + case err := <-readErr: + if !errors.Is(err, os.ErrClosed) { + t.Fatalf("pending Read returned %v, want os.ErrClosed", err) + } + case <-time.After(waitTimeout): + t.Fatal("Close did not unblock a pending Read") + } + if after := termios(t, probe); after.Lflag != before.Lflag { + t.Fatalf("Close left Lflag = %#x, want restored %#x", after.Lflag, before.Lflag) + } +} + +func TestSize(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + want := unix.Winsize{Row: 40, Col: 132, Xpixel: 1320, Ypixel: 800} + setWinsize(t, slave, &want) + + got, err := c.Size() + if err != nil { + t.Fatalf("Size: %v", err) + } + if got != (Size{Rows: 40, Cols: 132, Width: 1320, Height: 800}) { + t.Fatalf("Size() = %+v, want 40x132 with 1320x800 pixels", got) + } +} + +func TestResizesYieldsChangesOnly(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + setWinsize(t, slave, &unix.Winsize{Row: 24, Col: 80}) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + sizes := make(chan Size, 4) + done := make(chan struct{}) + resizes := c.Resizes(ctx) // baseline 24x80 is taken here + go func() { + defer close(done) + for sz := range resizes { + sizes <- sz + } + }() + + tick := time.NewTicker(20 * time.Millisecond) + defer tick.Stop() + winch := func() { + t.Helper() + if err := unix.Kill(os.Getpid(), unix.SIGWINCH); err != nil { + t.Fatalf("kill: %v", err) + } + } + + // Phase 1: SIGWINCH with the size unchanged, repeated for 300 ms so the + // iterator has certainly registered for the signal. Nothing may be yielded. + quiet := time.After(300 * time.Millisecond) +unchanged: + for { + winch() + select { + case sz := <-sizes: + t.Fatalf("an unchanged size was yielded: %+v", sz) + case <-tick.C: + case <-quiet: + break unchanged + } + } + + // Phase 2: change the size and signal until the new size arrives. + setWinsize(t, slave, &unix.Winsize{Row: 50, Col: 120}) + deadline := time.After(waitTimeout) + for { + winch() + select { + case sz := <-sizes: + if sz != (Size{Rows: 50, Cols: 120}) { + t.Fatalf("Resizes yielded %+v, want 50x120", sz) + } + cancel() + waitDone(t, done, "Resizes after cancel") + if len(sizes) != 0 { + t.Fatalf("the new size was yielded more than once: %+v", <-sizes) + } + return + case <-tick.C: + case <-deadline: + t.Fatal("Resizes never yielded the new size") + } + } +} + +// TestResizesYieldsChangeBeforeRangingStarts covers the gap between the +// Resizes call and the start of ranging over its result: SIGWINCH is only +// registered once the returned sequence starts running, so a resize that +// happens earlier must still be caught by the check Resizes makes right +// after registering, not lost until the next signal. +func TestResizesYieldsChangeBeforeRangingStarts(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + setWinsize(t, slave, &unix.Winsize{Row: 24, Col: 80}) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + resizes := c.Resizes(ctx) // baseline 24x80 is taken here + + // Change the size before ranging starts. No SIGWINCH is sent to this + // process: TIOCSWINSZ on Linux signals a pty's foreground process + // group, and this test process was never made that group (the pty was + // opened with O_NOCTTY and never became a controlling terminal), so + // the kernel does not deliver one here either (confirmed empirically + // against this behavior: signal.Notify(SIGWINCH) plus a bare + // TIOCSWINSZ on such a pty times out with nothing received). The only + // way the new size can reach the consumer below is the immediate + // post-Notify check inside Resizes. + setWinsize(t, slave, &unix.Winsize{Row: 50, Col: 120}) + + sizes := make(chan Size, 4) + done := make(chan struct{}) + go func() { + defer close(done) + for sz := range resizes { + sizes <- sz + } + }() + + select { + case sz := <-sizes: + if sz != (Size{Rows: 50, Cols: 120}) { + t.Fatalf("Resizes yielded %+v, want 50x120", sz) + } + case <-time.After(waitTimeout): + t.Fatal("Resizes never yielded the size that changed before ranging started") + } + cancel() + waitDone(t, done, "Resizes after cancel") +} + +// waitDone waits, bounded, for a consumer goroutine to close done. +func waitDone(t *testing.T, done <-chan struct{}, what string) { + t.Helper() + select { + case <-done: + case <-time.After(waitTimeout): + t.Fatalf("%s: the range over Resizes never finished", what) + } +} + +// TestResizesStopsAfterCancel checks that nothing is yielded once ctx has +// ended, even a change made before it ended: here the size changes and ctx +// is cancelled before ranging starts, so the check Resizes makes right +// after registering for SIGWINCH sees a change it must not yield. +func TestResizesStopsAfterCancel(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + setWinsize(t, slave, &unix.Winsize{Row: 24, Col: 80}) + + ctx, cancel := context.WithCancel(t.Context()) + resizes := c.Resizes(ctx) // baseline 24x80 is taken here + setWinsize(t, slave, &unix.Winsize{Row: 50, Col: 120}) + cancel() + + sizes := make(chan Size, 4) + done := make(chan struct{}) + go func() { + defer close(done) + for sz := range resizes { + sizes <- sz + } + }() + waitDone(t, done, "cancelled Resizes") + if len(sizes) != 0 { + t.Fatalf("Resizes yielded %+v after its context ended", <-sizes) + } +} + +// TestResizesStopsWhenConsumerBreaks checks that breaking out of the range +// ends it while ctx is still live, at both yield sites: the check right +// after SIGWINCH registration (a change made before ranging starts) and +// the signal loop (a change signalled while ranging). +func TestResizesStopsWhenConsumerBreaks(t *testing.T) { + winch := func(t *testing.T) { + t.Helper() + if err := unix.Kill(os.Getpid(), unix.SIGWINCH); err != nil { + t.Fatalf("kill: %v", err) + } + } + for _, tc := range []struct { + name string + // inLoop changes the size only after ranging has run for a while, + // so the signal loop yields it rather than the post-registration + // check. + inLoop bool + }{ + {name: "post_registration_check"}, + {name: "signal_loop", inLoop: true}, + } { + t.Run(tc.name, func(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + setWinsize(t, slave, &unix.Winsize{Row: 24, Col: 80}) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() // ctx stays live until the test has checked + resizes := c.Resizes(ctx) + if !tc.inLoop { + setWinsize(t, slave, &unix.Winsize{Row: 50, Col: 120}) + } + + got := make(chan Size, 1) + done := make(chan struct{}) + go func() { + defer close(done) + for sz := range resizes { + got <- sz + break + } + }() + + if tc.inLoop { + // Signal an unchanged size for 300 ms so the post-registration + // check has certainly run and found nothing, as in + // TestResizesYieldsChangesOnly, then change it. + tick := time.NewTicker(20 * time.Millisecond) + defer tick.Stop() + quiet := time.After(300 * time.Millisecond) + settle: + for { + winch(t) + select { + case sz := <-got: + t.Fatalf("an unchanged size was yielded: %+v", sz) + case <-tick.C: + case <-quiet: + break settle + } + } + setWinsize(t, slave, &unix.Winsize{Row: 50, Col: 120}) + deadline := time.After(waitTimeout) + signal: + for { + winch(t) + select { + case sz := <-got: + got <- sz // hand it on to the check below + break signal + case <-tick.C: + case <-deadline: + t.Fatal("Resizes never yielded the new size") + } + } + } + + select { + case sz := <-got: + if sz != (Size{Rows: 50, Cols: 120}) { + t.Fatalf("Resizes yielded %+v, want 50x120", sz) + } + case <-time.After(waitTimeout): + t.Fatal("Resizes never yielded the new size") + } + waitDone(t, done, "break with ctx live") + }) + } +} + +// TestResizesEndsAfterClose checks that the range over Resizes ends once +// the Console is closed while ctx is still live: at the start of ranging +// when Close ran before it, and at the next SIGWINCH when Close runs while +// the signal loop waits. +func TestResizesEndsAfterClose(t *testing.T) { + for _, tc := range []struct { + name string + // inLoop closes only after ranging has run for a while, so the + // signal loop, not the post-registration check, sees the Close. + inLoop bool + }{ + {name: "closed_before_ranging"}, + {name: "signal_loop", inLoop: true}, + } { + t.Run(tc.name, func(t *testing.T) { + _, slave := openPTY(t) + c := mustConsole(t, slave) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() // ctx stays live until the test has checked + resizes := c.Resizes(ctx) + if !tc.inLoop { + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + } + + sizes := make(chan Size, 4) + done := make(chan struct{}) + go func() { + defer close(done) + for sz := range resizes { + sizes <- sz + } + }() + + if tc.inLoop { + // Signal an unchanged size for 300 ms so the iterator has + // certainly registered and is waiting in the signal loop, + // as in TestResizesYieldsChangesOnly, then close. + if winchUntil(t, done, 300*time.Millisecond) { + t.Fatal("the range over Resizes ended before Close") + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + // The loop sees the Close only when a signal wakes it. + if !winchUntil(t, done, waitTimeout) { + t.Fatal("the range over Resizes did not end after Close") + } + } + waitDone(t, done, "Resizes after Close") + if len(sizes) != 0 { + t.Fatalf("Resizes yielded %+v around Close, want nothing", <-sizes) + } + }) + } +} + +// winchUntil sends this process SIGWINCH every 20 ms until done is closed, +// which it reports as true, or d passes, which it reports as false. +func winchUntil(t *testing.T, done <-chan struct{}, d time.Duration) bool { + t.Helper() + tick := time.NewTicker(20 * time.Millisecond) + defer tick.Stop() + deadline := time.After(d) + for { + if err := unix.Kill(os.Getpid(), unix.SIGWINCH); err != nil { + t.Fatalf("kill: %v", err) + } + select { + case <-done: + return true + case <-tick.C: + case <-deadline: + return false + } + } +} + +func setWinsize(t *testing.T, f *os.File, ws *unix.Winsize) { + t.Helper() + if err := controlFile(f, func(fd int) error { + return unix.IoctlSetWinsize(fd, unix.TIOCSWINSZ, ws) + }); err != nil { + t.Fatalf("set winsize: %v", err) + } +} diff --git a/internal/console/console_unix.go b/internal/console/console_unix.go new file mode 100644 index 0000000..39ed90f --- /dev/null +++ b/internal/console/console_unix.go @@ -0,0 +1,300 @@ +//go:build unix + +package console + +import ( + "context" + "errors" + "fmt" + "iter" + "os" + "os/signal" + "sync" + + "golang.org/x/sys/unix" + "golang.org/x/term" +) + +// Console is the controlling terminal, opened as /dev/tty. +// +// Opening /dev/tty creates an open file description of its own, so the +// non-blocking mode Go sets on it (which is what lets Close unblock a +// pending Read; MEASURED on Linux against a pty in the package tests, not +// yet measured on darwin) never leaks into the shell's stdin, even if et +// crashes. Fd is never called on the file: per the os.File.Fd docs its +// deadline methods would then stop working. All ioctls go through +// SyscallConn. +// +// Read and Write each have one owner goroutine: Read is not safe to call +// concurrently with itself, nor Write with itself; Read and Write may run +// concurrently with each other and with Close. +type Console struct { + tty *os.File + base *term.State // terminal state when the console was opened + + // setState applies a terminal state. It is term.Restore, held in a + // field so a test can count the calls restore makes. + setState func(fd int, s *term.State) error + + // mu serialises MakeRaw, restore, Size and Close. Close holds it for + // its whole body, so a second Close waits for the first to finish. + mu sync.Mutex + closed bool // guarded by mu +} + +// Open checks that stdin and stdout are terminals, then opens the +// controlling terminal (/dev/tty) as the console. It returns an error +// satisfying errors.Is(err, ErrNotTerminal) if either is not a terminal or +// the process has no controlling terminal. +// +// Open records the terminal state as the baseline that restore returns to. +// The intended caller opens the console before running anything that may +// leave it in a changed mode, such as ssh between Open and MakeRaw, where +// an interrupted password prompt can leave echo off. +func Open() (*Console, error) { + return open("/dev/tty", os.Stdin, os.Stdout) +} + +// open is Open with the terminal path and standard files as parameters, so +// tests can supply pipes and a missing path. +func open(ttyPath string, stdin, stdout *os.File) (*Console, error) { + if !isTerminal(stdin) || !isTerminal(stdout) { + return nil, ErrNotTerminal + } + tty, err := os.OpenFile(ttyPath, os.O_RDWR, 0) + if err != nil { + return nil, fmt.Errorf("%w: open %s: %w", ErrNotTerminal, ttyPath, err) + } + c, err := newConsole(tty) + if err != nil { + _ = tty.Close() + return nil, err + } + return c, nil +} + +// newConsole wraps an already open terminal file and records its current +// state as the baseline. Tests pass a pty slave. +func newConsole(tty *os.File) (*Console, error) { + c := &Console{tty: tty, setState: term.Restore} + err := c.control(func(fd int) error { + var gerr error + c.base, gerr = term.GetState(fd) + return gerr + }) + if err != nil { + return nil, fmt.Errorf("%w: read terminal state: %w", ErrNotTerminal, err) + } + return c, nil +} + +// MakeRaw switches the terminal to raw mode. It may be called again after +// restore; Close still returns the terminal to the Open baseline. After +// Close it returns an error satisfying errors.Is(err, os.ErrClosed) and +// leaves the terminal alone. +// +// The returned restore func returns the terminal to the state recorded by +// Open, not to whatever mode was current when MakeRaw ran. Every restore +// func does this whenever it runs, including one kept from an earlier +// MakeRaw, so running a stale one during a later raw session leaves raw +// mode. It reads the current state first and sets the baseline only when +// they differ, so it is repeatable and leaves a terminal already at the +// baseline untouched. It is safe to call from any goroutine. After Close it +// returns nil without touching the terminal, which Close already restored. +func (c *Console) MakeRaw() (restore func() error, err error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return nil, fmt.Errorf("console: make raw: %w", os.ErrClosed) + } + err = c.control(func(fd int) error { + _, rerr := term.MakeRaw(fd) + return rerr + }) + if err != nil { + return nil, fmt.Errorf("console: make raw: %w", err) + } + return c.restore, nil +} + +// restore is the func MakeRaw returns; see MakeRaw. +func (c *Console) restore() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return nil + } + return c.restoreLocked() +} + +// restoreLocked sets the Open baseline if the current state differs from +// it. Only setting the state can stop a background job (SIGTTOU; reading it +// cannot), so a terminal already at the baseline is left untouched. c.mu +// must be held. +func (c *Console) restoreLocked() error { + err := c.control(func(fd int) error { + cur, err := term.GetState(fd) + if err != nil { + return err + } + if *cur == *c.base { + return nil + } + return c.setState(fd, c.base) + }) + if err != nil { + return fmt.Errorf("console: restore: %w", err) + } + return nil +} + +// Read reads raw input bytes. A pending Read returns os.ErrClosed once Close +// is called (MEASURED on Linux against a pty in the package tests, not yet +// measured on darwin). +func (c *Console) Read(p []byte) (int, error) { return c.tty.Read(p) } + +// Write writes remote output to the terminal. After Close it returns an +// error satisfying errors.Is(err, os.ErrClosed), which os.File reports for +// a closed file. A Write that runs while Close is in progress can still +// reach the terminal until Close closes the descriptor, possibly after the +// mode was restored. This differs from Windows, where Write writes nothing +// once Close has started. +func (c *Console) Write(p []byte) (int, error) { return c.tty.Write(p) } + +// Size returns the current window size, including pixels when known. After +// Close it returns an error satisfying errors.Is(err, os.ErrClosed). +func (c *Console) Size() (Size, error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return Size{}, fmt.Errorf("console: size: %w", os.ErrClosed) + } + var ws *unix.Winsize + err := c.control(func(fd int) error { + var gerr error + ws, gerr = unix.IoctlGetWinsize(fd, unix.TIOCGWINSZ) + return gerr + }) + if err != nil { + return Size{}, fmt.Errorf("console: size: %w", err) + } + return Size{ + Rows: int(ws.Row), + Cols: int(ws.Col), + Width: int(ws.Xpixel), + Height: int(ws.Ypixel), + }, nil +} + +// Resizes yields the window size each time it changes, until ctx ends, the +// Console is closed, or the loop body stops. Changes are measured against +// the size when Resizes is called, which is not yielded itself; callers +// read it with Size. A change that happens before ranging starts is still +// caught: SIGWINCH registration and the first comparison against the +// call-time size both happen as soon as the returned sequence starts +// running, so no resize can fall in the gap. Nothing is yielded once ctx +// has ended, even a change made before it ended. Range over the result +// once. +// +// The sequence ends when ctx ends or the Console is closed. It notices a +// Close when it next reads the size: at once if the Console was closed +// before ranging started, otherwise at the next SIGWINCH, since the loop +// wakes only on that signal and on ctx. On Windows the next poll notices +// it instead. +func (c *Console) Resizes(ctx context.Context) iter.Seq[Size] { + last, _ := c.Size() + return func(yield func(Size) bool) { + sig := make(chan os.Signal, 1) + signal.Notify(sig, unix.SIGWINCH) + defer signal.Stop(sig) + + // A resize between the Resizes call and Notify registering above + // would otherwise be lost: SIGWINCH is ignored by default, so a + // signal that fires in that gap never reaches sig. Check once, + // right after registering, so such a change is still caught. + sz, err := c.Size() + if errors.Is(err, os.ErrClosed) { + return + } + if err == nil && sz != last { + last = sz + if ctx.Err() != nil || !yield(sz) { + return + } + } + + for { + select { + case <-ctx.Done(): + return + case <-sig: + } + sz, err := c.Size() + if errors.Is(err, os.ErrClosed) { + return + } + if err != nil || sz == last { + continue + } + last = sz + // select picks at random when a signal and the cancellation + // are both ready, so check ctx again before yielding. + if ctx.Err() != nil || !yield(sz) { + return + } + } + } +} + +// Close returns the terminal to the state recorded by Open whether or not +// MakeRaw was called, so a mode an interrupted prompt left behind is undone +// too, then closes the terminal. Like a restore func, it sets the baseline +// only when the current state differs from it. Closing the terminal +// unblocks a pending Read, which returns os.ErrClosed (MEASURED on Linux +// against a pty in the package tests, not yet measured on darwin). +// +// Close is idempotent. A second Close that runs while the first is still in +// progress waits for the first to finish, then returns nil; a Close after +// that returns nil at once. Only the first Close reports errors. +// +// Close does not discard typeahead the reader has not consumed; it stays +// queued for the next reader of the terminal. OpenSSH does the same: its +// leave_raw_mode (sshtty.c) restores with tcsetattr TCSADRAIN, which does +// not flush input. On Windows, Close flushes the input buffer instead. +func (c *Console) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return nil + } + c.closed = true + return errors.Join(c.restoreLocked(), c.tty.Close()) +} + +// control runs f with the terminal's file descriptor. +func (c *Console) control(f func(fd int) error) error { + return controlFile(c.tty, f) +} + +// controlFile runs f with file's descriptor through SyscallConn, never Fd. +func controlFile(file *os.File, f func(fd int) error) error { + rc, err := file.SyscallConn() + if err != nil { + return err + } + var ferr error + if err := rc.Control(func(fd uintptr) { ferr = f(int(fd)) }); err != nil { + return err + } + return ferr +} + +// isTerminal reports whether f is a terminal; an error reading it counts as no. +func isTerminal(f *os.File) bool { + var ok bool + err := controlFile(f, func(fd int) error { + ok = term.IsTerminal(fd) + return nil + }) + return err == nil && ok +} diff --git a/internal/console/console_windows.go b/internal/console/console_windows.go new file mode 100644 index 0000000..6dddd28 --- /dev/null +++ b/internal/console/console_windows.go @@ -0,0 +1,540 @@ +package console + +import ( + "context" + "errors" + "fmt" + "io" + "iter" + "os" + "sync" + "sync/atomic" + "time" + + "golang.org/x/sys/windows" +) + +const ( + // readUnits is the UTF-16 buffer size for one ReadConsoleW call. + readUnits = 4096 + // writeUnits caps one WriteConsoleW call: a write larger than the + // available heap fails with ERROR_NOT_ENOUGH_MEMORY (Microsoft docs, + // WriteConsole, nNumberOfCharsToWrite). + writeUnits = 8192 + // resizePoll is how often Resizes checks the window size. + resizePoll = 200 * time.Millisecond + // wakeEvery is how often Close re-injects the wake record while the + // reader has not returned yet. + wakeEvery = 100 * time.Millisecond + // closeWait bounds how long Close waits for the reader to return. + closeWait = time.Second +) + +// Console is the Windows console attached to the process. +// +// Read and Write each have one owner goroutine: Read is not safe to call +// concurrently with itself, nor Write with itself; Read and Write may run +// concurrently with each other and with Close. +type Console struct { + in, out windows.Handle + inBase, outBase uint32 // console modes when the console was opened + + // mu guards the flags below and is held across every console mode + // change, so MakeRaw, restore and Close never interleave their + // SetConsoleMode calls. + mu sync.Mutex + closing bool // guarded by mu: Close has started + reading bool // guarded by mu: a ReadConsoleW call is in flight + + // wmu is held by Write for its whole call and by Close before it + // restores the modes, so a Write in flight when Close starts finishes + // under the modes it started with. + wmu sync.Mutex + + exited chan struct{} // receives once when a reader returns after Close + done chan struct{} // closed when the first Close has finished + + // injectFn appends input records (WriteConsoleInputW), readFn reads + // UTF-16 units (ReadConsoleW) and writeFn writes them (WriteConsoleW). + // nil means the real call; tests set fakes to run without a console. + injectFn func(h windows.Handle, recs []inputRecord) error + readFn func(h windows.Handle, buf *uint16, toread uint32, read *uint32, inputControl *byte) error + writeFn func(h windows.Handle, buf *uint16, n uint32, written *uint32, reserved *byte) error + + // getModeFn and setModeFn read and set a console mode when restoring + // (GetConsoleMode, SetConsoleMode). nil means the real call. + getModeFn func(h windows.Handle, mode *uint32) error + setModeFn func(h windows.Handle, mode uint32) error + + // sizeFn reads the screen buffer info for Size + // (GetConsoleScreenBufferInfo). nil means the real call. + sizeFn func(h windows.Handle, info *windows.ConsoleScreenBufferInfo) error + + // Reader-owned state. + dec utf16Decoder + units []uint16 + pending []byte + + // Writer-owned state. + enc utf8Encoder + buf []uint16 +} + +// Open returns the console attached to stdin and stdout. It returns an error +// satisfying errors.Is(err, ErrNotTerminal) if either is not a console or +// its standard handle cannot be read; the error also wraps the cause. +// +// Open records the console modes as the baseline that restore returns to. +// The intended caller opens the console before running anything that may +// leave it in a changed mode, such as ssh.exe between Open and MakeRaw, +// where an interrupted password prompt can leave echo off. +func Open() (*Console, error) { + in, err := windows.GetStdHandle(windows.STD_INPUT_HANDLE) + if err != nil { + return nil, fmt.Errorf("%w: stdin handle: %w", ErrNotTerminal, err) + } + out, err := windows.GetStdHandle(windows.STD_OUTPUT_HANDLE) + if err != nil { + return nil, fmt.Errorf("%w: stdout handle: %w", ErrNotTerminal, err) + } + c := newConsole(in, out) + if err := windows.GetConsoleMode(in, &c.inBase); err != nil { + return nil, fmt.Errorf("%w: stdin console mode: %w", ErrNotTerminal, err) + } + if err := windows.GetConsoleMode(out, &c.outBase); err != nil { + return nil, fmt.Errorf("%w: stdout console mode: %w", ErrNotTerminal, err) + } + return c, nil +} + +// newConsole returns a Console on the given handles with its channels and +// read buffer set up. Open records the baseline modes; tests that need no +// console pass zero handles. +func newConsole(in, out windows.Handle) *Console { + return &Console{ + in: in, + out: out, + exited: make(chan struct{}, 1), + done: make(chan struct{}), + units: make([]uint16, readUnits), + } +} + +// MakeRaw switches the console to raw VT input and VT output processing, +// starting from the modes recorded by Open. Code pages are not touched: Read +// and Write use the UTF-16 console APIs, and a console code page applies +// only to the 8-bit form of those calls (Microsoft docs, ReadConsole and +// WriteConsole remarks: "uses either Unicode characters or 8-bit +// characters from the console's current code page"). It may be +// called again after restore; Close still returns the console to the Open +// baseline. Once Close has started it returns an error satisfying +// errors.Is(err, os.ErrClosed) and leaves the console alone. +// +// The returned restore func returns the console to the modes recorded by +// Open, not to whatever mode was current when MakeRaw ran. Every restore +// func does this whenever it runs, including one kept from an earlier +// MakeRaw, so running a stale one during a later raw session leaves raw +// mode. It reads the current modes first and sets the baseline only where +// they differ, so it is repeatable and leaves a console already at the +// baseline untouched. It is safe to call from any goroutine. Once Close has +// started it returns nil without touching the console, which Close +// restores. +func (c *Console) MakeRaw() (restore func() error, err error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closing { + return nil, fmt.Errorf("console: make raw: %w", os.ErrClosed) + } + // &^ clears only the listed flags, so ENABLE_EXTENDED_FLAGS and + // ENABLE_QUICK_EDIT_MODE keep whatever GetConsoleMode reported and Quick + // Edit stays as the user set it (MEASURED on win11-qa ConPTY: input mode + // 0x1f7 became 0x3f0, both flags still set). + rawIn := c.inBase&^(windows.ENABLE_PROCESSED_INPUT|windows.ENABLE_LINE_INPUT| + windows.ENABLE_ECHO_INPUT|windows.ENABLE_WINDOW_INPUT) | + windows.ENABLE_VIRTUAL_TERMINAL_INPUT + if err := windows.SetConsoleMode(c.in, rawIn); err != nil { + return nil, fmt.Errorf("console: set input mode: %w", err) + } + + rawOut := c.outBase | windows.ENABLE_VIRTUAL_TERMINAL_PROCESSING | + windows.ENABLE_PROCESSED_OUTPUT | windows.ENABLE_WRAP_AT_EOL_OUTPUT + if err := windows.SetConsoleMode(c.out, rawOut|windows.DISABLE_NEWLINE_AUTO_RETURN); err != nil { + // If the host rejects DISABLE_NEWLINE_AUTO_RETURN, run without it. + if err := windows.SetConsoleMode(c.out, rawOut); err != nil { + rollback := windows.SetConsoleMode(c.in, c.inBase) + if rollback != nil { + rollback = fmt.Errorf("console: roll back input mode: %w", rollback) + } + return nil, errors.Join(fmt.Errorf("console: set output mode: %w", err), rollback) + } + } + + return c.restore, nil +} + +// restore is the func MakeRaw returns; see MakeRaw. +func (c *Console) restore() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closing { + return nil + } + return c.restoreLocked() +} + +// restoreLocked sets each handle back to its Open baseline mode where the +// current mode differs from it. c.mu must be held. +func (c *Console) restoreLocked() error { + err := errors.Join( + c.setModeIfChanged(c.in, c.inBase), + c.setModeIfChanged(c.out, c.outBase), + ) + if err != nil { + return fmt.Errorf("console: restore: %w", err) + } + return nil +} + +func (c *Console) setModeIfChanged(h windows.Handle, mode uint32) error { + get, set := c.getModeFn, c.setModeFn + if get == nil { + get = windows.GetConsoleMode + } + if set == nil { + set = windows.SetConsoleMode + } + var cur uint32 + if err := get(h, &cur); err != nil { + return err + } + if cur == mode { + return nil + } + return set(h, mode) +} + +// Read reads raw VT input as UTF-8. Once Close has started, Read returns +// os.ErrClosed, including a Read that was blocked and any bytes still +// buffered from an earlier console read. A console read that returns no +// units is retried. A Read with an empty p returns 0, nil at once without +// reading the console. +// +// A generated Ctrl+Break does not end a pending read. MEASURED on win11-qa +// under ConPTY (ssh -tt), 2026-09-25: with a handler that consumes +// CTRL_BREAK_EVENT (as OnBreak installs), GenerateConsoleCtrlEvent +// (CTRL_BREAK_EVENT, 0) ran the handler and left ReadConsoleW pending in +// raw and in line mode, so Read has no handling for it. A physical +// Ctrl+Break key press is unmeasured. +func (c *Console) Read(p []byte) (int, error) { + read := c.readFn + if read == nil { + read = windows.ReadConsole + } + for { + c.mu.Lock() + if c.closing { + c.mu.Unlock() + return 0, os.ErrClosed + } + if len(p) == 0 { + // Nothing to read into: do not block in the console. + c.mu.Unlock() + return 0, nil + } + if len(c.pending) > 0 { + c.mu.Unlock() + break + } + c.reading = true + c.mu.Unlock() + + var n uint32 + err := read(c.in, &c.units[0], uint32(len(c.units)), &n, nil) + + c.mu.Lock() + c.reading = false + closing := c.closing + c.mu.Unlock() + if closing { + select { + case c.exited <- struct{}{}: + default: + } + return 0, os.ErrClosed + } + if err != nil { + return 0, fmt.Errorf("console: read: %w", err) + } + c.pending = c.dec.append(c.pending[:0], c.units[:n]) + } + n := copy(p, c.pending) + c.pending = c.pending[n:] + return n, nil +} + +// Write writes remote output, at most writeUnits UTF-16 units per console +// call, and never ends the chunk it hands to a call on a high surrogate. +// An incomplete UTF-8 sequence at the end of p is held until the next +// Write completes it. +// When the console reports a partial write, Write continues from the first +// unit not written, so a pair can be split between calls only if the +// console itself reports writing half of it. Whether any console host does +// that is unmeasured; Write does not back off to a pair boundary. +// +// On a console error Write reports 0 bytes written even if earlier chunks +// reached the screen: the UTF-16 units are not mapped back to input bytes. +// io.Writer allows any n < len(p) with a non-nil error, and callers treat a +// console write error as fatal. +// +// Once Close has started Write returns an error satisfying +// errors.Is(err, os.ErrClosed) and writes nothing. A Write already in +// progress when Close starts finishes first: Close restores the console +// modes only after it returns. This differs from Unix, where a Write that +// runs while Close is in progress can still reach the terminal until Close +// closes the descriptor. +func (c *Console) Write(p []byte) (int, error) { + c.wmu.Lock() + defer c.wmu.Unlock() + c.mu.Lock() + closing := c.closing + c.mu.Unlock() + if closing { + return 0, os.ErrClosed + } + write := c.writeFn + if write == nil { + write = windows.WriteConsole + } + c.buf = c.enc.append(c.buf[:0], p) + for units := c.buf; len(units) > 0; { + chunk := units[:chunkLen(units, writeUnits)] + var n uint32 + if err := write(c.out, &chunk[0], uint32(len(chunk)), &n, nil); err != nil { + return 0, fmt.Errorf("console: write: %w", err) + } + if n == 0 { + return 0, io.ErrShortWrite + } + if int(n) > len(chunk) { + return 0, fmt.Errorf("console: write: console reported %d units written of %d", n, len(chunk)) + } + units = units[n:] + } + return len(p), nil +} + +// Size returns the visible window size in character cells. Windows does not +// report pixel sizes, so Width and Height are 0. Once Close has started it +// returns an error satisfying errors.Is(err, os.ErrClosed). Size holds the +// Console lock for the whole query, and Close takes that lock to start, so +// a Size that succeeds finished before Close started. +func (c *Console) Size() (Size, error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closing { + return Size{}, fmt.Errorf("console: size: %w", os.ErrClosed) + } + query := c.sizeFn + if query == nil { + query = windows.GetConsoleScreenBufferInfo + } + var info windows.ConsoleScreenBufferInfo + if err := query(c.out, &info); err != nil { + return Size{}, fmt.Errorf("console: size: %w", err) + } + w := info.Window + return Size{Rows: int(w.Bottom-w.Top) + 1, Cols: int(w.Right-w.Left) + 1}, nil +} + +// Resizes yields the window size each time it changes, until ctx ends, the +// Console is closed, or the loop body stops. Changes are measured against +// the size when Resizes is called, which is not yielded itself; callers +// read it with Size. The window is polled every 200 ms: window resize +// events are always filtered by ReadConsole (Microsoft docs, SetConsoleMode +// remarks). Nothing is yielded once ctx has ended, even a change made +// before it ended. Range over the result once. +// +// The sequence ends when ctx ends or the Console is closed. It notices a +// Close at the next poll. On Unix the next SIGWINCH notices it instead, or +// the start of ranging if the Console was closed before that. +func (c *Console) Resizes(ctx context.Context) iter.Seq[Size] { + last, _ := c.Size() + return func(yield func(Size) bool) { + tick := time.NewTicker(resizePoll) + defer tick.Stop() + for { + select { + case <-ctx.Done(): + return + case <-tick.C: + } + sz, err := c.Size() + if errors.Is(err, os.ErrClosed) { + return + } + if err != nil || sz == last { + continue + } + last = sz + // select picks at random when a tick and the cancellation are + // both ready, so check ctx again before yielding. + if ctx.Err() != nil || !yield(sz) { + return + } + } + } +} + +// Close returns the console to the modes recorded by Open whether or not +// MakeRaw was called, so a mode an interrupted prompt left behind is undone +// too. Like a restore func, it sets a baseline mode only where the current +// mode differs from it. Close unblocks a pending Read, which then returns +// os.ErrClosed, and flushes the input buffer. A Write in progress when +// Close starts finishes before the modes are restored. Close waits for it +// with no bound: whether WriteConsoleW can stall (for example while a +// classic conhost QuickEdit selection pauses output) is unmeasured. The +// handles are the process's standard handles and stay open. +// +// Close is idempotent. A second Close that runs while the first is still in +// progress waits for the first to finish, then returns nil; a Close after +// that returns nil at once. Only the first Close reports errors. A non-nil +// error can mean the reader is still blocked in ReadConsoleW; it returns +// os.ErrClosed once it wakes. +// +// A blocked ReadConsoleW is woken by injecting an Enter key-down record; +// the reader sees the closing flag and discards what it read. Enter is +// used because in line mode ReadConsole "returns only when a carriage +// return character is read" (Microsoft docs, SetConsoleMode, +// ENABLE_LINE_INPUT), which a Read meets after restore ran before Close or +// when MakeRaw was never called. MEASURED on win11-qa under ConPTY (ssh +// -tt), 2026-09-25: the record wakes a Read in raw mode and in line mode, +// and no input events remain after Close; a space record left a line-mode +// Read blocked. In line mode with echo on, the console echoes the carriage +// return as a line break on screen. +// +// The record is re-injected every 100 ms until the reader returns (or 1 s +// passes) because the console's single input buffer can be shared by any +// number of processes (Microsoft docs, Consoles), so another process +// attached to the console may consume a record before this reader does. +// The input buffer is flushed only after that, so no wake record or unread +// typeahead reaches the parent shell. +func (c *Console) Close() error { + c.mu.Lock() + if c.closing { + c.mu.Unlock() + <-c.done + return nil + } + c.closing = true + reading := c.reading + c.mu.Unlock() + defer close(c.done) + + var errs []error + if reading { + errs = append(errs, c.wakeAndWait()) + } + if err := windows.FlushConsoleInputBuffer(c.in); err != nil { + errs = append(errs, fmt.Errorf("console: flush input: %w", err)) + } + + // Let a Write in flight finish under the modes it started with; every + // later Write sees closing and writes nothing. + c.wmu.Lock() + c.mu.Lock() + errs = append(errs, c.restoreLocked()) + c.mu.Unlock() + c.wmu.Unlock() + return errors.Join(errs...) +} + +// wakeAndWait injects wake records until the reader reports that it +// returned, or closeWait passes. A failed injection does not end the wait: +// the next tick injects again, and if the reader returns anyway (another +// record woke it) the wait succeeds. +func (c *Console) wakeAndWait() error { + tick := time.NewTicker(wakeEvery) + defer tick.Stop() + deadline := time.After(closeWait) + var injectErr error + for { + if err := c.wakeReader(); err != nil { + injectErr = err + } + select { + case <-c.exited: + return nil + case <-tick.C: + case <-deadline: + return errors.Join(errors.New("console: reader did not return after close"), injectErr) + } + } +} + +// wakeReader injects an Enter key-down so a pending ReadConsoleW returns. A +// carriage return completes the read in raw mode and in line mode alike. +func (c *Console) wakeReader() error { + rec := inputRecord{ + EventType: windows.KEY_EVENT, + Event: keyEventRecord{ + KeyDown: 1, + RepeatCount: 1, + VirtualKeyCode: 0x0D, // VK_RETURN + // MapVirtualKeyW(VK_RETURN, MAPVK_VK_TO_VSC) (MEASURED on + // win11-qa, 2026-09-25). + VirtualScanCode: 0x1C, + UnicodeChar: '\r', + }, + } + inject := c.injectFn + if inject == nil { + inject = writeConsoleInput + } + if err := inject(c.in, []inputRecord{rec}); err != nil { + return fmt.Errorf("console: wake reader: %w", err) + } + return nil +} + +var ( + breakFunc atomic.Pointer[func()] + handlerOnce sync.Once +) + +// OnBreak registers f to be called when the user presses Ctrl+Break, which +// Windows still raises as a control event while processed input is off +// (Ctrl+C arrives as a byte instead; Microsoft docs, CTRL+C and CTRL+BREAK +// Signals: "CTRL+BREAK is always treated as a signal"). Only the most +// recent f is kept. OnBreak(nil) unregisters it, and Ctrl+Break then goes +// to the next handler: the Go runtime's, which delivers it as os.Interrupt +// if the program called signal.Notify for it, and otherwise passes it on to +// the default handler, which ends the process (runtime/os_windows.go, +// ctrlHandler, go1.27.0; Microsoft docs, HandlerRoutine remarks). If the +// handler cannot be installed, Ctrl+Break reaches those handlers directly. +func OnBreak(f func()) { + if f == nil { + breakFunc.Store(nil) + return + } + breakFunc.Store(&f) + handlerOnce.Do(func() { + _ = setConsoleCtrlHandler(windows.NewCallback(ctrlHandler), true) + }) +} + +// ctrlHandler handles CTRL_BREAK_EVENT when a function is registered, and +// leaves every other event (Ctrl+C, close, logoff, shutdown) to the next +// handler. It runs on a thread Windows creates for the call (Microsoft docs, +// HandlerRoutine: "the system creates a new thread in the process to +// execute the function"). +func ctrlHandler(ctrlType uint32) uintptr { + if ctrlType != windows.CTRL_BREAK_EVENT { + return 0 + } + f := breakFunc.Load() + if f == nil { + return 0 + } + (*f)() + return 1 +} diff --git a/internal/console/console_windows_test.go b/internal/console/console_windows_test.go new file mode 100644 index 0000000..f9981bf --- /dev/null +++ b/internal/console/console_windows_test.go @@ -0,0 +1,1057 @@ +package console + +import ( + "context" + "errors" + "io" + "os" + "slices" + "strings" + "sync/atomic" + "testing" + "time" + "unicode/utf16" + "unsafe" + + "golang.org/x/sys/windows" +) + +func TestInputRecordLayout(t *testing.T) { + if got := unsafe.Sizeof(keyEventRecord{}); got != 16 { + t.Fatalf("sizeof(KEY_EVENT_RECORD) = %d, want 16", got) + } + if got := unsafe.Sizeof(inputRecord{}); got != 20 { + t.Fatalf("sizeof(INPUT_RECORD) = %d, want 20", got) + } + if got := unsafe.Offsetof(inputRecord{}.Event); got != 4 { + t.Fatalf("offsetof(INPUT_RECORD.Event) = %d, want 4", got) + } + // KEY_EVENT_RECORD member offsets and sizes (Microsoft docs, + // KEY_EVENT_RECORD: BOOL, WORD, WORD, WORD, union uChar, DWORD). + var k keyEventRecord + for _, f := range []struct { + name string + off, size uintptr + wantOff, wantSize uintptr + }{ + {"bKeyDown", unsafe.Offsetof(k.KeyDown), unsafe.Sizeof(k.KeyDown), 0, 4}, + {"wRepeatCount", unsafe.Offsetof(k.RepeatCount), unsafe.Sizeof(k.RepeatCount), 4, 2}, + {"wVirtualKeyCode", unsafe.Offsetof(k.VirtualKeyCode), unsafe.Sizeof(k.VirtualKeyCode), 6, 2}, + {"wVirtualScanCode", unsafe.Offsetof(k.VirtualScanCode), unsafe.Sizeof(k.VirtualScanCode), 8, 2}, + {"uChar", unsafe.Offsetof(k.UnicodeChar), unsafe.Sizeof(k.UnicodeChar), 10, 2}, + {"dwControlKeyState", unsafe.Offsetof(k.ControlKeyState), unsafe.Sizeof(k.ControlKeyState), 12, 4}, + } { + if f.off != f.wantOff || f.size != f.wantSize { + t.Errorf("KEY_EVENT_RECORD.%s at offset %d, size %d; want offset %d, size %d", + f.name, f.off, f.size, f.wantOff, f.wantSize) + } + } +} + +func TestCtrlHandler(t *testing.T) { + t.Cleanup(func() { OnBreak(nil) }) + + called := 0 + OnBreak(func() { called++ }) + if got := ctrlHandler(windows.CTRL_BREAK_EVENT); got != 1 || called != 1 { + t.Fatalf("Ctrl+Break: handler returned %d and called f %d times, want 1 and 1", got, called) + } + if got := ctrlHandler(windows.CTRL_C_EVENT); got != 0 || called != 1 { + t.Fatalf("Ctrl+C: handler returned %d and called f %d times, want 0 and 1", got, called) + } + + OnBreak(nil) + if got := ctrlHandler(windows.CTRL_BREAK_EVENT); got != 0 || called != 1 { + t.Fatalf("after OnBreak(nil): handler returned %d and called f %d times, want 0 and 1", got, called) + } +} + +// TestReadAfterCloseIgnoresPending needs no console: a closed Console must +// not hand out bytes buffered from an earlier console read. +func TestReadAfterCloseIgnoresPending(t *testing.T) { + c := newConsole(0, 0) + c.pending = []byte("left over") + c.closing = true + n, err := c.Read(make([]byte, 16)) + if n != 0 || !errors.Is(err, os.ErrClosed) { + t.Fatalf("Read after Close = %d, %v; want 0, os.ErrClosed", n, err) + } +} + +// The tests below that need no console build it on zero handles: the +// console calls Close makes on them (flush, mode read) fail, so the first +// Close returns an error, which these tests do not inspect unless they say +// so. + +// waitTimeout bounds how long a test waits for a call or goroutine that +// must finish, with or without a console. +const waitTimeout = 5 * time.Second + +// stillRunning is how long a console-free test watches a call that must +// stay blocked. When the code under test is broken such a call makes only +// failing console calls and returns almost at once, so the bound only has +// to exceed scheduling noise. +const stillRunning = 200 * time.Millisecond + +// closeAsync runs c.Close on its own goroutine. +func closeAsync(c *Console) <-chan error { + ch := make(chan error, 1) + go func() { ch <- c.Close() }() + return ch +} + +func recv(t *testing.T, ch <-chan error, what string) error { + t.Helper() + select { + case err := <-ch: + return err + case <-time.After(waitTimeout): + t.Fatalf("%s never returned", what) + return nil + } +} + +// TestCloseIsIdempotent needs no console: once the first Close has +// finished, later calls return nil at once, without its error. +func TestCloseIsIdempotent(t *testing.T) { + c := newConsole(0, 0) + _ = recv(t, closeAsync(c), "first Close") + for i := range 2 { + if err := recv(t, closeAsync(c), "Close after Close"); err != nil { + t.Fatalf("Close %d after the first = %v, want nil", i+1, err) + } + } +} + +// TestConcurrentCloseWaits needs no console: a Close that starts while the +// first Close is still waking the reader must wait for it, then return nil. +func TestConcurrentCloseWaits(t *testing.T) { + c := newConsole(0, 0) + c.reading = true // the first Close takes the wake path + entered := make(chan struct{}) + release := make(chan struct{}) + c.injectFn = func(windows.Handle, []inputRecord) error { + select { + case <-entered: + default: + close(entered) + <-release + } + // The reader returns. A later tick may inject again before Close + // consumes this, so do not block on a full channel. + select { + case c.exited <- struct{}{}: + default: + } + return nil + } + + first := closeAsync(c) + select { + case <-entered: + case <-time.After(waitTimeout): + t.Fatal("the first Close never tried to wake the reader") + } + second := closeAsync(c) + select { + case err := <-second: + close(release) + t.Fatalf("a second Close returned (%v) while the first was still in progress", err) + case <-time.After(stillRunning): + } + close(release) + _ = recv(t, first, "first Close") + if err := recv(t, second, "second Close"); err != nil { + t.Fatalf("second Close = %v, want nil", err) + } +} + +// TestWakeRetriesAfterInjectError needs no console: a failed wake injection +// must not end Close's wait for the reader; the next tick injects again. +func TestWakeRetriesAfterInjectError(t *testing.T) { + c := newConsole(0, 0) + c.reading = true + errInject := errors.New("injected failure") + calls := 0 + c.injectFn = func(windows.Handle, []inputRecord) error { + calls++ + if calls == 1 { + return errInject + } + select { + case c.exited <- struct{}{}: + default: + } + return nil + } + err := c.Close() + if calls < 2 { + t.Fatalf("Close injected %d wake records, want a retry after the failed one", calls) + } + if errors.Is(err, errInject) { + t.Fatalf("Close = %v, want no wake error once the reader returned", err) + } +} + +func TestWriteAfterClose(t *testing.T) { + c := newConsole(0, 0) + c.writeFn = func(windows.Handle, *uint16, uint32, *uint32, *byte) error { + t.Error("Write after Close reached the console") + return nil + } + _ = c.Close() + if n, err := c.Write([]byte("x")); n != 0 || !errors.Is(err, os.ErrClosed) { + t.Fatalf("Write after Close = %d, %v; want 0, os.ErrClosed", n, err) + } +} + +func TestMakeRawAfterClose(t *testing.T) { + c := newConsole(0, 0) + _ = c.Close() + restore, err := c.MakeRaw() + if restore != nil || !errors.Is(err, os.ErrClosed) { + t.Fatalf("MakeRaw after Close = (restore set: %t), %v; want nil, os.ErrClosed", restore != nil, err) + } +} + +func TestSizeAfterClose(t *testing.T) { + c := newConsole(0, 0) + _ = c.Close() + if _, err := c.Size(); !errors.Is(err, os.ErrClosed) { + t.Fatalf("Size after Close = %v, want an error wrapping os.ErrClosed", err) + } +} + +// TestCloseWaitsForInFlightWrite needs no console: when Close starts while +// a Write is inside its console call, Close must not set a console mode +// until that call has returned, and must not return before it either. The +// fake mode calls report a mode that differs from the baseline, so Close +// really sets modes and the order can be observed. +func TestCloseWaitsForInFlightWrite(t *testing.T) { + c := newConsole(0, 0) + c.inBase, c.outBase = 0x1f7, 0x7 + var writing atomic.Bool // set while the fake console write is in progress + var sets atomic.Int32 + c.getModeFn = func(_ windows.Handle, mode *uint32) error { + *mode = 0 // differs from both baselines, so Close sets each + return nil + } + c.setModeFn = func(windows.Handle, uint32) error { + sets.Add(1) + if writing.Load() { + t.Error("Close set a console mode while a Write was still in its console call") + } + return nil + } + entered := make(chan struct{}) + release := make(chan struct{}) + c.writeFn = func(_ windows.Handle, _ *uint16, n uint32, written *uint32, _ *byte) error { + writing.Store(true) + close(entered) + <-release + writing.Store(false) + *written = n + return nil + } + wrote := make(chan error, 1) + go func() { + _, err := c.Write([]byte("x")) + wrote <- err + }() + select { + case <-entered: + case <-time.After(waitTimeout): + t.Fatal("Write never reached the console") + } + + closed := closeAsync(c) + select { + case <-closed: + close(release) + t.Fatal("Close returned while a Write was still in progress") + case <-time.After(stillRunning): + } + close(release) + if err := recv(t, wrote, "Write"); err != nil { + t.Fatalf("in-flight Write = %v, want nil", err) + } + _ = recv(t, closed, "Close") + if sets.Load() == 0 { + t.Fatal("Close set no console mode, so the test observed no ordering") + } +} + +// TestOpenWithoutConsoleKeepsCause runs where stdin is not a console (the +// test binary run with its standard handles redirected): Open must report +// ErrNotTerminal and keep the Windows error that caused it. +func TestOpenWithoutConsoleKeepsCause(t *testing.T) { + c, err := Open() + if err == nil { + if cerr := c.Close(); cerr != nil { + t.Errorf("Close: %v", cerr) + } + t.Skip("a console is attached; run the test binary with redirected standard handles") + } + if !errors.Is(err, ErrNotTerminal) { + t.Fatalf("Open error = %v, want ErrNotTerminal", err) + } + if _, ok := errors.AsType[windows.Errno](err); !ok { + t.Fatalf("Open error = %v, want it to wrap the Windows error that caused it", err) + } +} + +// rangeResizes ranges over c.Resizes(ctx) on its own goroutine and closes +// the returned channel when the range statement finishes. Any size yielded +// fails the test: these Consoles have no window to resize. +func rangeResizes(t *testing.T, c *Console, ctx context.Context) <-chan struct{} { + t.Helper() + done := make(chan struct{}) + resizes := c.Resizes(ctx) + go func() { + defer close(done) + for sz := range resizes { + t.Errorf("Resizes yielded %+v, want nothing", sz) + } + }() + return done +} + +// TestResizesEndsAfterClose needs no console: once the Console is closed, +// the next poll ends the range while ctx is still live. +func TestResizesEndsAfterClose(t *testing.T) { + c := newConsole(0, 0) + _ = c.Close() + done := rangeResizes(t, c, t.Context()) + select { + case <-done: + case <-time.After(waitTimeout): + t.Fatal("the range over Resizes did not end after Close") + } +} + +// TestResizesKeepsPollingOnOtherErrors needs no console: a Size error other +// than a closed Console (here the zero handles are invalid) does not end +// the range; only ctx does. +func TestResizesKeepsPollingOnOtherErrors(t *testing.T) { + c := newConsole(0, 0) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + done := rangeResizes(t, c, ctx) + select { + case <-done: + t.Fatal("the range over Resizes ended on a Size error that is not a Close") + case <-time.After(3 * resizePoll): + } + cancel() + select { + case <-done: + case <-time.After(waitTimeout): + t.Fatal("the range over Resizes did not end after ctx was cancelled") + } +} + +// openConsole returns the real console, or skips when the test binary has +// none (for example when its output is piped). The console is closed when +// the test ends, which returns it to the Open baseline for the next test; +// a Close the test made itself leaves that one returning nil. +func openConsole(t *testing.T) *Console { + t.Helper() + c, err := Open() + if errors.Is(err, ErrNotTerminal) { + t.Skip("no console attached; run the test binary under ssh -tt or in a console window") + } + if err != nil { + t.Fatalf("Open: %v", err) + } + t.Cleanup(func() { + if err := c.Close(); err != nil { + t.Errorf("Close at test end: %v", err) + } + }) + return c +} + +func TestCloseUnblocksRead(t *testing.T) { + c := openConsole(t) + if _, err := c.MakeRaw(); err != nil { + t.Fatalf("MakeRaw: %v", err) + } + readErr := make(chan error, 1) + go func() { + _, err := c.Read(make([]byte, 16)) + readErr <- err + }() + // Wait until the reader is inside ReadConsoleW, so Close takes the + // wake-up path rather than the not-yet-reading path. + waitReading(t, c) + + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + select { + case err := <-readErr: + if !errors.Is(err, os.ErrClosed) { + t.Fatalf("pending Read returned %v, want os.ErrClosed", err) + } + case <-time.After(waitTimeout): + t.Fatal("Close did not unblock a pending Read") + } + + if err := recv(t, closeAsync(c), "second Close"); err != nil { + t.Fatalf("second Close: %v, want nil", err) + } +} + +// TestCloseFlushesTypeahead checks that input nobody read is discarded by +// Close, so it does not reach the shell et returns to. +func TestCloseFlushesTypeahead(t *testing.T) { + c := openConsole(t) + key := inputRecord{ + EventType: windows.KEY_EVENT, + Event: keyEventRecord{KeyDown: 1, RepeatCount: 1, VirtualKeyCode: 'A', UnicodeChar: 'a'}, + } + if err := writeConsoleInput(c.in, []inputRecord{key, key}); err != nil { + t.Fatalf("inject typeahead: %v", err) + } + if n := inputEvents(t, c); n == 0 { + t.Fatal("injected typeahead is not in the input buffer") + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if n := inputEvents(t, c); n != 0 { + t.Fatalf("%d input events left after Close, want 0", n) + } +} + +func inputEvents(t *testing.T, c *Console) uint32 { + t.Helper() + var n uint32 + if err := windows.GetNumberOfConsoleInputEvents(c.in, &n); err != nil { + t.Fatalf("GetNumberOfConsoleInputEvents: %v", err) + } + return n +} + +// TestCloseRestoresOutputMode checks that Close returns the output handle, +// not only the input handle, to its Open baseline. +func TestCloseRestoresOutputMode(t *testing.T) { + c := openConsole(t) + if _, err := c.MakeRaw(); err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if mode := outputMode(t, c); mode == c.outBase { + t.Skipf("MakeRaw left the output mode at the baseline %#x", mode) + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if mode := outputMode(t, c); mode != c.outBase { + t.Fatalf("output mode after Close = %#x, want the Open baseline %#x", mode, c.outBase) + } +} + +func outputMode(t *testing.T, c *Console) uint32 { + t.Helper() + var mode uint32 + if err := windows.GetConsoleMode(c.out, &mode); err != nil { + t.Fatalf("GetConsoleMode(out): %v", err) + } + return mode +} + +// TestCloseUnblocksLineModeRead covers a Read blocked while the input is in +// line mode, where ReadConsoleW returns only on a carriage return: after +// restore ran before Close, or when MakeRaw was never called. +func TestCloseUnblocksLineModeRead(t *testing.T) { + for _, tc := range []struct { + name string + restore bool // call MakeRaw and its restore before reading + }{ + {name: "restored", restore: true}, + {name: "never_raw", restore: false}, + } { + t.Run(tc.name, func(t *testing.T) { + c := openConsole(t) + if c.inBase&windows.ENABLE_LINE_INPUT == 0 { + t.Skipf("Open-time input mode %#x is not line mode", c.inBase) + } + if tc.restore { + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := restore(); err != nil { + t.Fatalf("restore: %v", err) + } + } + readErr := make(chan error, 1) + go func() { + _, err := c.Read(make([]byte, 16)) + readErr <- err + }() + waitReading(t, c) + + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + select { + case err := <-readErr: + if !errors.Is(err, os.ErrClosed) { + t.Fatalf("pending Read returned %v, want os.ErrClosed", err) + } + case <-time.After(waitTimeout): + t.Fatal("Close did not unblock a Read in line mode") + } + var left uint32 + if err := windows.GetNumberOfConsoleInputEvents(c.in, &left); err != nil { + t.Fatalf("GetNumberOfConsoleInputEvents: %v", err) + } + if left != 0 { + t.Fatalf("%d input events left after Close, want 0", left) + } + }) + } +} + +// resetInputAtEnd sets c's input back to its Open baseline when the test +// ends, reporting a failure. +func resetInputAtEnd(t *testing.T, c *Console) { + t.Helper() + t.Cleanup(func() { + if err := windows.SetConsoleMode(c.in, c.inBase); err != nil { + t.Errorf("reset input mode to %#x: %v", c.inBase, err) + } + }) +} + +func TestRestoreReturnsToOpenBaseline(t *testing.T) { + c := openConsole(t) + resetInputAtEnd(t, c) + // Model ssh.exe interrupted at a password prompt: echo left off after Open. + if err := windows.SetConsoleMode(c.in, c.inBase&^windows.ENABLE_ECHO_INPUT); err != nil { + t.Fatalf("degrade input mode: %v", err) + } + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := restore(); err != nil { + t.Fatalf("restore: %v", err) + } + var mode uint32 + if err := windows.GetConsoleMode(c.in, &mode); err != nil { + t.Fatalf("GetConsoleMode: %v", err) + } + if mode != c.inBase { + t.Fatalf("input mode after restore = %#x, want the Open baseline %#x", mode, c.inBase) + } + if mode := outputMode(t, c); mode != c.outBase { + t.Fatalf("output mode after restore = %#x, want the Open baseline %#x", mode, c.outBase) + } +} + +// fakeModes makes c's restore read the given modes, in call order, and +// records the modes it sets. It needs no console. +func fakeModes(t *testing.T, c *Console, setErr error, reported ...uint32) *[]uint32 { + t.Helper() + var set []uint32 + c.getModeFn = func(_ windows.Handle, mode *uint32) error { + if len(reported) == 0 { + t.Errorf("restore read more console modes than the script reports") + return errors.New("script exhausted") + } + *mode = reported[0] + reported = reported[1:] + return nil + } + c.setModeFn = func(_ windows.Handle, mode uint32) error { + set = append(set, mode) + return setErr + } + return &set +} + +// TestRestoreSkipsUnchangedMode needs no console: Close sets a handle's +// baseline mode only when the current mode differs from it. +func TestRestoreSkipsUnchangedMode(t *testing.T) { + const inBase, outBase = 0x1f7, 0x7 + for _, tc := range []struct { + name string + reported []uint32 // input mode, then output mode + want []uint32 + }{ + {name: "at_baseline", reported: []uint32{inBase, outBase}, want: nil}, + {name: "input_changed", reported: []uint32{0x3f0, outBase}, want: []uint32{inBase}}, + {name: "both_changed", reported: []uint32{0x3f0, 0x1f}, want: []uint32{inBase, outBase}}, + } { + t.Run(tc.name, func(t *testing.T) { + c := newConsole(0, 0) + c.inBase, c.outBase = inBase, outBase + set := fakeModes(t, c, nil, tc.reported...) + _ = c.Close() // the flush on a zero handle fails; not inspected + if !slices.Equal(*set, tc.want) { + t.Fatalf("Close set modes %#x, want %#x", *set, tc.want) + } + }) + } + + // A failed set reaches Close's caller. + t.Run("set_fails", func(t *testing.T) { + c := newConsole(0, 0) + c.inBase, c.outBase = inBase, outBase + errSet := errors.New("set console mode failed") + set := fakeModes(t, c, errSet, 0x3f0, outBase) + err := c.Close() + if len(*set) != 1 { + t.Fatalf("Close made %d set calls, want 1", len(*set)) + } + if !errors.Is(err, errSet) { + t.Fatalf("Close = %v, want an error wrapping the failed set", err) + } + }) +} + +// TestRestoreSerializedWithClose needs no console: a Close that starts +// while a restore func is setting modes must wait for it, so the two never +// set console modes at the same time. +func TestRestoreSerializedWithClose(t *testing.T) { + c := newConsole(0, 0) + c.inBase, c.outBase = 0x1f7, 0x7 + c.getModeFn = func(_ windows.Handle, mode *uint32) error { + *mode = 0 // always differs, so every restore sets + return nil + } + entered := make(chan struct{}) + release := make(chan struct{}) + var inFlight atomic.Int32 + c.setModeFn = func(windows.Handle, uint32) error { + if inFlight.Add(1) > 1 { + t.Error("Close set a console mode while restore was setting one") + } + defer inFlight.Add(-1) + select { + case <-entered: + default: + close(entered) + <-release + } + return nil + } + + restored := make(chan error, 1) + go func() { restored <- c.restore() }() + select { + case <-entered: + case <-time.After(waitTimeout): + t.Fatal("restore never set a mode") + } + closed := closeAsync(c) + select { + case <-closed: + close(release) + t.Fatal("Close returned while restore was still setting modes") + case <-time.After(stillRunning): + } + close(release) + if err := recv(t, restored, "restore"); err != nil { + t.Fatalf("restore = %v, want nil", err) + } + _ = recv(t, closed, "Close") +} + +// echoOff turns echo off on c's input, as ssh.exe interrupted at a password +// prompt can leave it, and restores the Open baseline when the test ends. +func echoOff(t *testing.T, c *Console) uint32 { + t.Helper() + if c.inBase&windows.ENABLE_ECHO_INPUT == 0 { + t.Skipf("Open-time input mode %#x has echo off already", c.inBase) + } + resetInputAtEnd(t, c) + degraded := c.inBase &^ windows.ENABLE_ECHO_INPUT + if err := windows.SetConsoleMode(c.in, degraded); err != nil { + t.Fatalf("degrade input mode: %v", err) + } + return degraded +} + +func inputMode(t *testing.T, c *Console) uint32 { + t.Helper() + var mode uint32 + if err := windows.GetConsoleMode(c.in, &mode); err != nil { + t.Fatalf("GetConsoleMode: %v", err) + } + return mode +} + +// TestCloseRestoresBaselineWithoutMakeRaw covers et closing before MakeRaw +// ran, after a prompt left echo off: Close must still restore the baseline. +func TestCloseRestoresBaselineWithoutMakeRaw(t *testing.T) { + c := openConsole(t) + echoOff(t, c) + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if mode := inputMode(t, c); mode != c.inBase { + t.Fatalf("input mode after Close = %#x, want the Open baseline %#x", mode, c.inBase) + } +} + +// TestRestoreAfterClose checks that a restore func run after Close returns +// nil and leaves the console alone: Close already restored it. +func TestRestoreAfterClose(t *testing.T) { + c := openConsole(t) + restore, err := c.MakeRaw() + if err != nil { + t.Fatalf("MakeRaw: %v", err) + } + if err := c.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + // Another program now owns the console and turned echo off. + changed := echoOff(t, c) + if err := restore(); err != nil { + t.Fatalf("restore after Close = %v, want nil", err) + } + if mode := inputMode(t, c); mode != changed { + t.Fatalf("restore after Close set input mode %#x, want it left at %#x", mode, changed) + } +} + +// writeStep is one scripted WriteConsoleW result: the units it reports +// written (all of the chunk when all is set) or an error. +type writeStep struct { + n uint32 + all bool + err error +} + +// scriptWrite sets c.writeFn to a fake that records every chunk and answers +// with steps in order. It fails the test if Write calls it more often than +// the script allows. +func scriptWrite(t *testing.T, c *Console, steps ...writeStep) *[][]uint16 { + t.Helper() + var chunks [][]uint16 + c.writeFn = func(_ windows.Handle, buf *uint16, n uint32, written *uint32, _ *byte) error { + chunks = append(chunks, slices.Clone(unsafe.Slice(buf, n))) + if len(steps) == 0 { + t.Errorf("Write made console call %d, beyond its script", len(chunks)) + return errors.New("script exhausted") + } + s := steps[0] + steps = steps[1:] + *written = s.n + if s.all { + *written = n + } + return s.err + } + return &chunks +} + +// fullWrites answers every console call by reporting the whole chunk +// written, for up to n calls. +func fullWrites(n int) []writeStep { + return slices.Repeat([]writeStep{{all: true}}, n) +} + +func TestWriteLarge(t *testing.T) { + c := newConsole(0, 0) + chunksp := scriptWrite(t, c, fullWrites(8)...) + + // Several times writeUnits, with an emoji straddling the first chunk edge. + big := strings.Repeat("x", writeUnits-1) + "😀" + strings.Repeat("y", 3*writeUnits) + "\r\n" + if n, err := c.Write([]byte(big)); err != nil || n != len(big) { + t.Fatalf("Write(%d bytes) = %d, %v", len(big), n, err) + } + + want := utf16.Encode([]rune(big)) + chunks := *chunksp + if len(chunks) == 0 { + t.Fatal("Write made no console calls") + } + if len(chunks[0]) != writeUnits-1 { + t.Fatalf("first chunk has %d units, want %d (the pair moved to the next chunk)", len(chunks[0]), writeUnits-1) + } + var got []uint16 + for i, ch := range chunks { + if len(ch) > writeUnits { + t.Fatalf("chunk %d has %d units, want at most %d", i, len(ch), writeUnits) + } + if len(ch) > 0 && isHighSurrogate(ch[len(ch)-1]) { + t.Fatalf("chunk %d ends on a high surrogate", i) + } + got = append(got, ch...) + } + if !slices.Equal(got, want) { + t.Fatalf("chunks carry %d units, want the %d units of the input in order", len(got), len(want)) + } +} + +func waitReading(t *testing.T, c *Console) { + t.Helper() + tick := time.NewTicker(10 * time.Millisecond) + defer tick.Stop() + deadline := time.After(waitTimeout) + for { + c.mu.Lock() + reading := c.reading + c.mu.Unlock() + if reading { + return + } + select { + case <-tick.C: + case <-deadline: + t.Fatal("reader never entered ReadConsoleW") + } + } +} + +func TestWriteSplitUTF8(t *testing.T) { + c := newConsole(0, 0) + chunks := scriptWrite(t, c, fullWrites(1)...) + euro := []byte("€\r\n") + // The first two parts hold an incomplete sequence and reach no console + // call; the third completes it. + wantCalls := []int{0, 0, 1} + for i, part := range [][]byte{euro[:1], euro[1:2], euro[2:]} { + if n, err := c.Write(part); err != nil || n != len(part) { + t.Fatalf("Write(%q) = %d, %v", part, n, err) + } + if len(*chunks) != wantCalls[i] { + t.Fatalf("after Write %d (%q) the console had %d calls, want %d", i+1, part, len(*chunks), wantCalls[i]) + } + } + if want := []uint16{0x20AC, '\r', '\n'}; !slices.Equal((*chunks)[0], want) { + t.Fatalf("console got %#x, want %#x", (*chunks)[0], want) + } +} + +// TestWriteConsoleResults covers how Write reacts to what WriteConsoleW +// reports: all written, a partial write, nothing written, an error, and an +// impossible count larger than the chunk. +func TestWriteConsoleResults(t *testing.T) { + errConsole := errors.New("console failed") + twoChunks := strings.Repeat("z", writeUnits+10) // one full chunk and 10 units + for _, tc := range []struct { + name string + p string + steps []writeStep + wantN int + wantErr error // matched with errors.Is; nil means no error + anyErr bool // an error is wanted but no sentinel names it + wantCalls []int // units handed to each console call + }{ + {name: "full", p: "hello", steps: fullWrites(1), wantN: 5, wantCalls: []int{5}}, + {name: "partial", p: "hello", steps: []writeStep{{n: 2}, {all: true}}, wantN: 5, wantCalls: []int{5, 3}}, + {name: "nothing_written", p: "hello", steps: []writeStep{{n: 0}}, wantErr: io.ErrShortWrite, wantCalls: []int{5}}, + {name: "error", p: "hello", steps: []writeStep{{err: errConsole}}, wantErr: errConsole, wantCalls: []int{5}}, + { + name: "count_beyond_chunk", p: twoChunks, + steps: []writeStep{{n: writeUnits + 1}, {all: true}}, + anyErr: true, wantCalls: []int{writeUnits}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + c := newConsole(0, 0) + chunks := scriptWrite(t, c, tc.steps...) + n, err := c.Write([]byte(tc.p)) + switch { + case tc.wantErr != nil: + if n != 0 || !errors.Is(err, tc.wantErr) { + t.Fatalf("Write = %d, %v; want 0, %v", n, err, tc.wantErr) + } + case tc.anyErr: + if n != 0 || err == nil { + t.Fatalf("Write = %d, %v; want 0 and an error", n, err) + } + default: + if n != tc.wantN || err != nil { + t.Fatalf("Write = %d, %v; want %d, nil", n, err, tc.wantN) + } + } + calls := make([]int, 0, len(*chunks)) + for _, ch := range *chunks { + calls = append(calls, len(ch)) + } + if !slices.Equal(calls, tc.wantCalls) { + t.Fatalf("console calls carried %v units, want %v", calls, tc.wantCalls) + } + }) + } +} + +// scriptRead sets c.readFn to a fake that returns each element of reads +// from one ReadConsoleW call, in order. It fails the test if Read calls it +// more often than the script allows. +func scriptRead(t *testing.T, c *Console, reads ...[]uint16) { + t.Helper() + calls := 0 + c.readFn = func(_ windows.Handle, buf *uint16, toread uint32, read *uint32, _ *byte) error { + calls++ + if len(reads) == 0 { + t.Errorf("Read made console read %d, beyond its script", calls) + return errors.New("script exhausted") + } + *read = uint32(copy(unsafe.Slice(buf, toread), reads[0])) + reads = reads[1:] + return nil + } +} + +func readString(t *testing.T, c *Console, size int) string { + t.Helper() + p := make([]byte, size) + n, err := c.Read(p) + if err != nil { + t.Fatalf("Read: %v", err) + } + return string(p[:n]) +} + +// TestReadJoinsSplitSurrogatePair needs no console: a pair whose halves +// arrive from two console reads decodes to one UTF-8 sequence. +func TestReadJoinsSplitSurrogatePair(t *testing.T) { + c := newConsole(0, 0) + pair := utf16.Encode([]rune("😀")) + scriptRead(t, c, pair[:1], pair[1:]) + if got := readString(t, c, 16); got != "😀" { + t.Fatalf("Read = %q, want %q", got, "😀") + } +} + +// TestReadPassesEscapeSequences needs no console: terminal replies such as +// a device attributes reply and a cursor position report reach the reader +// byte for byte. +func TestReadPassesEscapeSequences(t *testing.T) { + c := newConsole(0, 0) + const replies = "\x1b[?1;2c\x1b[12;40R" + scriptRead(t, c, utf16.Encode([]rune(replies))) + if got := readString(t, c, 64); got != replies { + t.Fatalf("Read = %q, want %q", got, replies) + } +} + +// TestReadSmallBufferKeepsRest needs no console: bytes that do not fit p +// are returned by the next Read without another console read. +func TestReadSmallBufferKeepsRest(t *testing.T) { + c := newConsole(0, 0) + scriptRead(t, c, utf16.Encode([]rune("abcdef"))) + if got := readString(t, c, 3); got != "abc" { + t.Fatalf("first Read = %q, want %q", got, "abc") + } + if got := readString(t, c, 16); got != "def" { + t.Fatalf("second Read = %q, want %q", got, "def") + } +} + +// TestReadEmptyBufferReturnsAtOnce needs no console: a Read with an empty p +// returns 0, nil without reading the console, where it could block. +func TestReadEmptyBufferReturnsAtOnce(t *testing.T) { + c := newConsole(0, 0) + scriptRead(t, c) // any console read fails the test + if n, err := c.Read(nil); n != 0 || err != nil { + t.Fatalf("Read(nil) = %d, %v; want 0, nil", n, err) + } + if n, err := c.Read([]byte{}); n != 0 || err != nil { + t.Fatalf("Read(empty) = %d, %v; want 0, nil", n, err) + } + + // Once closed, even an empty Read reports the Close. + _ = c.Close() + if n, err := c.Read(nil); n != 0 || !errors.Is(err, os.ErrClosed) { + t.Fatalf("Read(nil) after Close = %d, %v; want 0 and an error wrapping os.ErrClosed", n, err) + } +} + +// TestSizeHoldsLockAgainstClose needs no console: a Size query in progress +// when Close starts finishes before Close sets any mode or returns, and +// Size after Close reports os.ErrClosed. +func TestSizeHoldsLockAgainstClose(t *testing.T) { + c := newConsole(0, 0) + c.inBase, c.outBase = 0x1f7, 0x7 + var querying atomic.Bool // set while the fake size query is in progress + var sets atomic.Int32 + c.getModeFn = func(_ windows.Handle, mode *uint32) error { + *mode = 0 // differs from both baselines, so Close sets each + return nil + } + c.setModeFn = func(windows.Handle, uint32) error { + sets.Add(1) + if querying.Load() { + t.Error("Close set a console mode while Size was still querying") + } + return nil + } + entered := make(chan struct{}) + release := make(chan struct{}) + c.sizeFn = func(_ windows.Handle, info *windows.ConsoleScreenBufferInfo) error { + querying.Store(true) + close(entered) + <-release + querying.Store(false) + info.Window = windows.SmallRect{Left: 0, Top: 0, Right: 79, Bottom: 23} + return nil + } + + type result struct { + sz Size + err error + } + sized := make(chan result, 1) + go func() { + sz, err := c.Size() + sized <- result{sz, err} + }() + select { + case <-entered: + case <-time.After(waitTimeout): + t.Fatal("Size never queried the console") + } + + closed := closeAsync(c) + select { + case <-closed: + close(release) + t.Fatal("Close returned while Size was still querying") + case <-time.After(stillRunning): + } + close(release) + select { + case r := <-sized: + if r.err != nil || r.sz != (Size{Rows: 24, Cols: 80}) { + t.Fatalf("Size = %+v, %v; want 24x80, nil", r.sz, r.err) + } + case <-time.After(waitTimeout): + t.Fatal("Size never returned") + } + _ = recv(t, closed, "Close") + if sets.Load() == 0 { + t.Fatal("Close set no console mode, so the test observed no ordering") + } + if _, err := c.Size(); !errors.Is(err, os.ErrClosed) { + t.Fatalf("Size after Close = %v, want an error wrapping os.ErrClosed", err) + } +} + +// TestReadRetriesEmptyRead needs no console: a console read that returns no +// units does not end Read with zero bytes; Read reads again. +func TestReadRetriesEmptyRead(t *testing.T) { + c := newConsole(0, 0) + scriptRead(t, c, nil, utf16.Encode([]rune("x"))) + if got := readString(t, c, 16); got != "x" { + t.Fatalf("Read = %q, want %q", got, "x") + } +} + +func TestSizeReportsCells(t *testing.T) { + c := openConsole(t) + sz, err := c.Size() + if err != nil { + t.Fatalf("Size: %v", err) + } + if sz.Rows < 1 || sz.Cols < 1 || sz.Width != 0 || sz.Height != 0 { + t.Fatalf("Size() = %+v, want positive cells and no pixels", sz) + } +} diff --git a/internal/console/pty_linux_test.go b/internal/console/pty_linux_test.go new file mode 100644 index 0000000..b658187 --- /dev/null +++ b/internal/console/pty_linux_test.go @@ -0,0 +1,69 @@ +package console + +import ( + "errors" + "fmt" + "os" + "testing" + + "golang.org/x/sys/unix" +) + +// openPTY returns a pseudo-terminal pair using only x/sys ioctls, so the +// console tests need no extra module and no controlling terminal. Both ends +// are closed when the test ends. +func openPTY(t *testing.T) (master, slave *os.File) { + t.Helper() + master, err := os.OpenFile("/dev/ptmx", os.O_RDWR|unix.O_NOCTTY, 0) + if err != nil { + t.Fatalf("open /dev/ptmx: %v", err) + } + closeAtEnd(t, master) + + var n int + err = controlFile(master, func(fd int) error { + if err := unix.IoctlSetPointerInt(fd, unix.TIOCSPTLCK, 0); err != nil { + return fmt.Errorf("unlockpt: %w", err) + } + var gerr error + n, gerr = unix.IoctlGetInt(fd, unix.TIOCGPTN) + return gerr + }) + if err != nil { + t.Fatalf("pty setup: %v", err) + } + slave, err = os.OpenFile(fmt.Sprintf("/dev/pts/%d", n), os.O_RDWR|unix.O_NOCTTY, 0) + if err != nil { + t.Fatalf("open pty slave: %v", err) + } + closeAtEnd(t, slave) + return master, slave +} + +// closeAtEnd closes f when the test ends and reports a close error. A +// Console built on f with newConsole owns it and closes it in Close, so +// the second close here is deliberate: os.ErrClosed is ignored, any other +// error is reported. +func closeAtEnd(t *testing.T, f *os.File) { + t.Helper() + t.Cleanup(func() { + if err := f.Close(); err != nil && !errors.Is(err, os.ErrClosed) { + t.Errorf("close %s: %v", f.Name(), err) + } + }) +} + +// termios reads the terminal attributes of f. +func termios(t *testing.T, f *os.File) *unix.Termios { + t.Helper() + var tio *unix.Termios + err := controlFile(f, func(fd int) error { + var gerr error + tio, gerr = unix.IoctlGetTermios(fd, unix.TCGETS) + return gerr + }) + if err != nil { + t.Fatalf("tcgets: %v", err) + } + return tio +} diff --git a/internal/console/syscall_windows.go b/internal/console/syscall_windows.go new file mode 100644 index 0000000..b093f86 --- /dev/null +++ b/internal/console/syscall_windows.go @@ -0,0 +1,66 @@ +package console + +import ( + "unsafe" + + "golang.org/x/sys/windows" +) + +// kernel32 functions that golang.org/x/sys/windows v0.48.0 does not wrap. +var ( + modkernel32 = windows.NewLazySystemDLL("kernel32.dll") + procWriteConsoleInputW = modkernel32.NewProc("WriteConsoleInputW") + procSetConsoleCtrlHandler = modkernel32.NewProc("SetConsoleCtrlHandler") +) + +// keyEventRecord is KEY_EVENT_RECORD (16 bytes). +type keyEventRecord struct { + KeyDown int32 // BOOL + RepeatCount uint16 + VirtualKeyCode uint16 + VirtualScanCode uint16 + UnicodeChar uint16 // union uChar; only the WCHAR member is used + ControlKeyState uint32 +} + +// inputRecord is INPUT_RECORD (20 bytes). Event is a union in C; KEY_EVENT +// is the only variant this package writes, and KEY_EVENT_RECORD is one of +// the union's largest members (MOUSE_EVENT_RECORD is as large), so the +// layout matches. +type inputRecord struct { + EventType uint16 + _ uint16 // padding before the union + Event keyEventRecord +} + +// writeConsoleInput appends records to the console input buffer. +func writeConsoleInput(h windows.Handle, recs []inputRecord) error { + if len(recs) == 0 { + return nil + } + var written uint32 + r1, _, err := procWriteConsoleInputW.Call( + uintptr(h), + uintptr(unsafe.Pointer(&recs[0])), + uintptr(len(recs)), + uintptr(unsafe.Pointer(&written)), + ) + if r1 == 0 { + return err + } + return nil +} + +// setConsoleCtrlHandler adds (or removes) a handler created with +// windows.NewCallback. +func setConsoleCtrlHandler(handler uintptr, add bool) error { + var a uintptr + if add { + a = 1 + } + r1, _, err := procSetConsoleCtrlHandler.Call(handler, a) + if r1 == 0 { + return err + } + return nil +} diff --git a/internal/console/utf16.go b/internal/console/utf16.go new file mode 100644 index 0000000..25c2a21 --- /dev/null +++ b/internal/console/utf16.go @@ -0,0 +1,98 @@ +package console + +import ( + "unicode/utf16" + "unicode/utf8" +) + +// utf16Decoder converts UTF-16 code units to UTF-8. A high surrogate at the +// end of one call is carried into the next, so a character split across two +// console reads is decoded once, correctly. Unpaired surrogates become U+FFFD. +// The zero value is ready to use. Not safe for concurrent use. +type utf16Decoder struct { + high uint16 // pending high surrogate, 0 if none +} + +// append decodes units and appends the UTF-8 result to dst. +func (d *utf16Decoder) append(dst []byte, units []uint16) []byte { + for _, u := range units { + if d.high != 0 { + high := d.high + d.high = 0 + if isLowSurrogate(u) { + dst = utf8.AppendRune(dst, utf16.DecodeRune(rune(high), rune(u))) + continue + } + dst = utf8.AppendRune(dst, utf8.RuneError) + } + switch { + case isHighSurrogate(u): + d.high = u + case isLowSurrogate(u): + dst = utf8.AppendRune(dst, utf8.RuneError) + default: + dst = utf8.AppendRune(dst, rune(u)) + } + } + return dst +} + +// utf8Encoder converts UTF-8 to UTF-16. An incomplete UTF-8 sequence at the +// end of one call is carried into the next, so a character split across two +// network packets is encoded once, correctly. Invalid bytes become U+FFFD, +// one per byte, as in a Go []rune conversion. The zero value is ready to use. +// Not safe for concurrent use. +type utf8Encoder struct { + carry [utf8.UTFMax]byte + n int // bytes held in carry, always < utf8.UTFMax +} + +// append encodes p and appends the UTF-16 result to dst. +func (e *utf8Encoder) append(dst []uint16, p []byte) []uint16 { + if e.n > 0 { + // Join the carried bytes with at most UTFMax bytes of p; that is + // always enough to finish or reject the carried sequence. + k := min(len(p), utf8.UTFMax) + joined := append(e.carry[:e.n:e.n], p[:k]...) + used := 0 + for used < e.n { + if !utf8.FullRune(joined[used:]) { + // Only reachable when all of p fit into joined: keep waiting. + e.n = copy(e.carry[:], joined[used:]) + return dst + } + r, size := utf8.DecodeRune(joined[used:]) + dst = utf16.AppendRune(dst, r) + used += size + } + p = p[used-e.n:] + e.n = 0 + } + for len(p) > 0 { + if !utf8.FullRune(p) { + e.n = copy(e.carry[:], p) + break + } + r, size := utf8.DecodeRune(p) + dst = utf16.AppendRune(dst, r) + p = p[size:] + } + return dst +} + +// chunkLen returns how many of units to pass to one console write of at most +// limit units. It never ends a chunk between a high and a low surrogate, so +// each write carries whole characters. limit must be at least 2. +func chunkLen(units []uint16, limit int) int { + if len(units) <= limit { + return len(units) + } + if isHighSurrogate(units[limit-1]) && isLowSurrogate(units[limit]) { + return limit - 1 + } + return limit +} + +func isHighSurrogate(u uint16) bool { return u >= 0xD800 && u < 0xDC00 } + +func isLowSurrogate(u uint16) bool { return u >= 0xDC00 && u < 0xE000 } diff --git a/internal/console/utf16_test.go b/internal/console/utf16_test.go new file mode 100644 index 0000000..0e0a480 --- /dev/null +++ b/internal/console/utf16_test.go @@ -0,0 +1,210 @@ +package console + +import ( + "bytes" + "slices" + "testing" + "unicode/utf16" + "unicode/utf8" +) + +func TestUTF16DecoderWhole(t *testing.T) { + tests := []struct { + name string + units []uint16 + want string + }{ + {"ascii", []uint16{'h', 'i'}, "hi"}, + {"bmp", utf16.Encode([]rune("äö€")), "äö€"}, + {"surrogate pair", utf16.Encode([]rune("😀")), "😀"}, + {"vt sequence", utf16.Encode([]rune("\x1b[A")), "\x1b[A"}, + {"lone low", []uint16{0xDC00, 'a'}, "�a"}, + {"high then ascii", []uint16{0xD83D, 'a'}, "�a"}, + {"high then high", []uint16{0xD83D, 0xD83D, 0xDE00}, "�😀"}, + {"empty", nil, ""}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var d utf16Decoder + got := string(d.append(nil, tt.units)) + if got != tt.want { + t.Fatalf("append(%U) = %q, want %q", tt.units, got, tt.want) + } + }) + } +} + +func TestUTF16DecoderCarriesHighSurrogate(t *testing.T) { + units := utf16.Encode([]rune("a😀b")) + var d utf16Decoder + first := d.append(nil, units[:2]) // 'a' plus the high surrogate + if string(first) != "a" { + t.Fatalf("first call = %q, want %q (high surrogate must be held back)", first, "a") + } + second := d.append(nil, units[2:]) + if string(second) != "😀b" { + t.Fatalf("second call = %q, want %q", second, "😀b") + } +} + +func TestUTF8EncoderWhole(t *testing.T) { + tests := []struct { + name string + in string + want []uint16 + }{ + {"ascii", "hi", []uint16{'h', 'i'}}, + {"multibyte", "äö€", utf16.Encode([]rune("äö€"))}, + {"astral", "😀", []uint16{0xD83D, 0xDE00}}, + {"invalid byte", "a\xffb", []uint16{'a', 0xFFFD, 'b'}}, + {"truncated inside", "\xe2\x82a", []uint16{0xFFFD, 0xFFFD, 'a'}}, + {"empty", "", nil}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var e utf8Encoder + got := e.append(nil, []byte(tt.in)) + if !slices.Equal(got, tt.want) { + t.Fatalf("append(%q) = %U, want %U", tt.in, got, tt.want) + } + }) + } +} + +func TestUTF8EncoderCarriesSplitSequence(t *testing.T) { + in := []byte("x€y") // € is 3 bytes: e2 82 ac + for split := range len(in) + 1 { + var e utf8Encoder + got := e.append(nil, in[:split]) + got = e.append(got, in[split:]) + want := utf16.Encode([]rune("x€y")) + if !slices.Equal(got, want) { + t.Fatalf("split at %d: got %U, want %U", split, got, want) + } + } +} + +func TestUTF8EncoderHoldsIncompleteTail(t *testing.T) { + var e utf8Encoder + got := e.append(nil, []byte("a\xf0\x9f")) // first half of 😀 + if !slices.Equal(got, []uint16{'a'}) { + t.Fatalf("got %U, want only 'a' while the tail is incomplete", got) + } + got = e.append(nil, []byte{0x98}) + if len(got) != 0 { + t.Fatalf("got %U, want nothing: 3 of 4 bytes still incomplete", got) + } + got = e.append(nil, []byte{0x80}) + if !slices.Equal(got, []uint16{0xD83D, 0xDE00}) { + t.Fatalf("got %U, want the surrogate pair for U+1F600", got) + } +} + +func TestChunkLen(t *testing.T) { + pair := utf16.Encode([]rune("😀")) // high, low + tests := []struct { + name string + units []uint16 + limit int + want int + }{ + {"fits", []uint16{'a', 'b'}, 4, 2}, + {"exact", []uint16{'a', 'b', 'c', 'd'}, 4, 4}, + {"split plain", []uint16{'a', 'b', 'c', 'd', 'e'}, 4, 4}, + {"pair at boundary", append([]uint16{'a', 'b', 'c'}, pair...), 4, 3}, + {"pair before boundary", append([]uint16{'a', 'b'}, append(pair, 'c')...), 4, 4}, + {"lone high at boundary", []uint16{'a', 'b', 'c', 0xD83D, 'd'}, 4, 4}, + {"empty", nil, 4, 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := chunkLen(tt.units, tt.limit); got != tt.want { + t.Fatalf("chunkLen(%U, %d) = %d, want %d", tt.units, tt.limit, got, tt.want) + } + }) + } +} + +// FuzzChunkLen checks that chunking any input covers it exactly, never +// exceeds the limit, and never separates a surrogate pair. +func FuzzChunkLen(f *testing.F) { + f.Add([]byte("h\x00i\x00=\xd8\x00\xde"), uint8(2)) + f.Add([]byte("\x30\x30\x30\x30\x30\xd8\x30\xd8"), uint8(2)) // two lone high surrogates, limit 4 + f.Fuzz(func(t *testing.T, raw []byte, lim uint8) { + limit := int(lim)%16 + 2 + units := make([]uint16, len(raw)/2) + for i := range units { + units[i] = uint16(raw[2*i]) | uint16(raw[2*i+1])<<8 + } + var joined []uint16 + for rest := units; len(rest) > 0; { + n := chunkLen(rest, limit) + if n < 1 || n > limit { + t.Fatalf("chunkLen = %d with limit %d", n, limit) + } + if n < len(rest) && isHighSurrogate(rest[n-1]) && isLowSurrogate(rest[n]) { + t.Fatalf("chunk separates a surrogate pair: %U | %U", rest[:n], rest[n:n+1]) + } + joined = append(joined, rest[:n]...) + rest = rest[n:] + } + if !slices.Equal(joined, units) { + t.Fatalf("chunks %U do not rejoin to %U", joined, units) + } + }) +} + +// FuzzUTF8EncoderSplit checks that splitting the input anywhere never changes +// the output, and that valid UTF-8 matches the standard library's encoding. +func FuzzUTF8EncoderSplit(f *testing.F) { + f.Add([]byte("hello, 世界 😀"), uint8(3)) + f.Add([]byte("\xe2\x82"), uint8(1)) + f.Add([]byte("a\xffb\xf0\x9f\x98\x80"), uint8(5)) + f.Fuzz(func(t *testing.T, in []byte, at uint8) { + var whole utf8Encoder + want := whole.append(nil, in) + + split := int(at) % (len(in) + 1) + var parts utf8Encoder + got := parts.append(nil, in[:split]) + got = parts.append(got, in[split:]) + + if !slices.Equal(got, want) { + t.Fatalf("split %d of %q: got %U, want %U", split, in, got, want) + } + if !slices.Equal(parts.carry[:parts.n], whole.carry[:whole.n]) { + t.Fatalf("split %d of %q: carry state differs", split, in) + } + if utf8.Valid(in) && !slices.Equal(want, utf16.Encode([]rune(string(in)))) { + t.Fatalf("valid input %q: got %U, want utf16.Encode", in, want) + } + }) +} + +// FuzzUTF16DecoderSplit is the decoder counterpart: any split gives the same +// output, and valid UTF-16 round-trips through the standard library. +func FuzzUTF16DecoderSplit(f *testing.F) { + f.Add([]byte("h\x00i\x00=\xd8\x00\xde"), uint8(2)) + f.Add([]byte("\x00\xdc"), uint8(0)) + f.Fuzz(func(t *testing.T, raw []byte, at uint8) { + units := make([]uint16, len(raw)/2) + for i := range units { + units[i] = uint16(raw[2*i]) | uint16(raw[2*i+1])<<8 + } + var whole utf16Decoder + want := whole.append(nil, units) + + split := int(at) % (len(units) + 1) + var parts utf16Decoder + got := parts.append(nil, units[:split]) + got = parts.append(got, units[split:]) + + if !bytes.Equal(got, want) || parts.high != whole.high { + t.Fatalf("split %d of %U: got %q, want %q", split, units, got, want) + } + runes := utf16.Decode(units) + if !slices.Contains(runes, utf8.RuneError) && whole.high == 0 && string(want) != string(runes) { + t.Fatalf("valid input %U: got %q, want %q", units, want, string(runes)) + } + }) +}