diff --git a/cmd/et/flags.go b/cmd/et/flags.go index 7f45236..9e55225 100644 --- a/cmd/et/flags.go +++ b/cmd/et/flags.go @@ -131,11 +131,19 @@ func parseArgs(args []string, stderr io.Writer) (*options, error) { // parseDestination parses [user@]host[:port]. IPv6 literals need brackets to // carry a port ("[::1]:2022"); a bare IPv6 literal ("::1") has no port. A -// user or host beginning with "-" is rejected, since it would otherwise be -// read as an option by ssh -G in resolveHost. A ":" with nothing after it is -// rejected rather than treated as "no port". +// user or host beginning with "-" is rejected as a usage error: bootstrap +// refuses it anyway, since ssh substitutes the host and user into a +// ProxyCommand or Match exec (MEASURED against OpenSSH_10.0p2 on 2026-09-26, +// see bootstrap's shellMeta). A ":" with nothing after it is rejected rather +// than treated as "no port", and so is an ssh:// URI. func parseDestination(s string) (destination, error) { var d destination + // An ssh:// URI would be split at its last '@' into a user such as + // "ssh://bob" or taken whole as a host, and its port names sshd's port + // where et's host:port names etserver's; refuse it rather than guess. + if strings.Contains(s, "://") { + return d, fmt.Errorf("ssh:// URIs are not supported, use [user@]host[:port]: %q", s) + } if before, after, found := strings.CutLast(s, "@"); found { if before == "" { return d, fmt.Errorf("empty user in destination %q", s) @@ -178,9 +186,8 @@ func parseDestination(s string) (destination, error) { if d.Host == "" { return d, fmt.Errorf("empty host in destination") } - // A host beginning with "-" would otherwise reach "ssh -G " (see - // resolveHost in resolve.go) as an option rather than as an argument; - // reject it here so the caller gets a usage error instead. + // A host beginning with "-" is refused by bootstrap's validation; reject + // it here so the caller gets a usage error instead. if strings.HasPrefix(d.Host, "-") { return d, fmt.Errorf("host %q looks like an option in destination %q", d.Host, s) } diff --git a/cmd/et/flags_test.go b/cmd/et/flags_test.go index 2657021..fc682b1 100644 --- a/cmd/et/flags_test.go +++ b/cmd/et/flags_test.go @@ -44,6 +44,10 @@ func TestParseArgsDestination(t *testing.T) { {name: "host looks like an option", args: []string{"--", "-evil"}, wantErr: true}, {name: "user looks like an option", args: []string{"--", "-evil@box"}, wantErr: true}, {name: "-u looks like an option", args: []string{"-u", "-oProxyCommand=x", "box"}, wantErr: true}, + {name: "user holding @", args: []string{"me@corp.example@box"}, want: destination{User: "me@corp.example", Host: "box", Port: 2022}}, + {name: "ssh URI with user", args: []string{"ssh://bob@h1"}, wantErr: true}, + {name: "user and ssh URI host", args: []string{"bob@ssh://h1:2222"}, wantErr: true}, + {name: "ssh URI", args: []string{"ssh://h1"}, wantErr: true}, {name: "trailing colon", args: []string{"box:"}, wantErr: true}, {name: "bracket trailing colon", args: []string{"[::1]:"}, wantErr: true}, } diff --git a/cmd/et/resolve.go b/cmd/et/resolve.go index e21b1ef..19514b3 100644 --- a/cmd/et/resolve.go +++ b/cmd/et/resolve.go @@ -20,20 +20,20 @@ const ( // resolveHost returns the host name OpenSSH would connect to for this // destination, so an ssh_config alias ("Host box / HostName 10.0.0.5") works // for the etserver TCP connection too. "ssh -G" prints the effective client -// configuration without connecting. Each of opts is passed as "-o", the -// same way bootstrap passes it to the ssh that starts etterminal, so an -// option such as HostName resolves here too. On any failure it returns host -// unchanged. +// configuration without connecting. The arguments have the shape bootstrap +// gives the ssh that starts etterminal: each of opts as "-o", the user +// as "-l ", and "--" before the host, so an option such as HostName +// and a user name holding '@' resolve as they do there. +// On any failure it returns host unchanged. func resolveHost(ctx context.Context, sshPath, user, host string, opts []string) string { - target := host - if user != "" { - target = user + "@" + host - } - args := make([]string, 0, len(opts)+2) + args := make([]string, 0, len(opts)+5) for _, opt := range opts { args = append(args, "-o"+opt) } - args = append(args, "-G", target) + if user != "" { + args = append(args, "-l", user) + } + args = append(args, "-G", "--", host) ctx, cancel := context.WithTimeout(ctx, resolveTimeout) defer cancel() cmd := exec.CommandContext(ctx, sshPath, args...) diff --git a/cmd/et/resolve_test.go b/cmd/et/resolve_test.go index c27899b..5d0e36d 100644 --- a/cmd/et/resolve_test.go +++ b/cmd/et/resolve_test.go @@ -31,10 +31,11 @@ func fakeSSH(t *testing.T, body string) string { return path } -// TestResolveHostPassesOptions pins the exact ssh -G argument list: each -// option is its own "-o" word before -G, as bootstrap passes it to the -// ssh that starts etterminal. The fake prints its arguments joined by "|" as -// the hostname, so an option containing a space stays one word. +// TestResolveHostPassesOptions pins the exact ssh -G argument list, which has +// the shape bootstrap gives the ssh that starts etterminal: each option as +// its own "-o" word, the user as "-l ", then "--" before the host. +// The fake prints its arguments joined by "|" as the hostname, so an option +// containing a space stays one word. func TestResolveHostPassesOptions(t *testing.T) { ssh := fakeSSH(t, `printf 'user x\nhostname '; printf '%s|' "$@"; printf '\n'`) tests := []struct { @@ -43,13 +44,14 @@ func TestResolveHostPassesOptions(t *testing.T) { opts []string want string }{ - {name: "no options", want: "-G|box|"}, - {name: "user", user: "me", want: "-G|me@box|"}, + {name: "no options", want: "-G|--|box|"}, + {name: "user", user: "me", want: "-l|me|-G|--|box|"}, + {name: "user with at", user: "me@corp.example", want: "-l|me@corp.example|-G|--|box|"}, { name: "options", user: "me", opts: []string{"HostName=10.0.0.5", "ProxyCommand=nc a b"}, - want: "-oHostName=10.0.0.5|-oProxyCommand=nc a b|-G|me@box|", + want: "-oHostName=10.0.0.5|-oProxyCommand=nc a b|-l|me|-G|--|box|", }, } for _, tt := range tests { diff --git a/internal/bootstrap/bootstrap.go b/internal/bootstrap/bootstrap.go index 25a9961..d93abfd 100644 --- a/internal/bootstrap/bootstrap.go +++ b/internal/bootstrap/bootstrap.go @@ -12,6 +12,7 @@ import ( "os/exec" "strings" "time" + "unicode/utf8" ) var ( @@ -24,7 +25,9 @@ var ( const ( // maxOutput caps how much ssh stdout is kept. Shell startup noise before // the marker is normally a few lines, far below the cap; a marker that - // arrives after the cap is dropped and the start fails. + // arrives after the cap is dropped and the start fails, with an error + // that says output was dropped. Keep it a whole number of MiB: that + // error states it in MiB. maxOutput = 1 << 20 // excerptLen caps the output excerpt quoted in an error. excerptLen = 512 @@ -38,7 +41,7 @@ const ( // usable: Destination is required. type Config struct { Destination string // ssh destination as typed; ssh_config aliases work - User string // optional; becomes user@destination + User string // optional; passed to ssh as -l, so it may contain '@' TerminalPath string // default "etterminal"; must match [A-Za-z0-9._/~-]+ Term string // default "xterm-256color"; must match [A-Za-z0-9.+-]+ (no underscore) SSHOptions []string // each passed as its own "-o" argument @@ -77,21 +80,7 @@ func Run(ctx context.Context, cfg Config) (Credentials, error) { } id, passkey := placeholder() - cmd := exec.CommandContext(ctx, sshPath, sshArgs(&cfg, remoteCommand(id, passkey, cfg.Term, cfg.TerminalPath))...) - cmd.Stdin = os.Stdin - cmd.Stderr = os.Stderr - out := &cappedBuffer{max: maxOutput} - cmd.Stdout = out - cmd.Cancel = func() error { - // Interrupt first so ssh can restore the terminal; Windows cannot - // deliver os.Interrupt to another process, so kill there. - if err := cmd.Process.Signal(os.Interrupt); err != nil { - return cmd.Process.Kill() - } - return nil - } - cmd.WaitDelay = waitDelay - + cmd, out := newCommand(ctx, sshPath, sshArgs(&cfg, remoteCommand(id, passkey, cfg.Term, cfg.TerminalPath))) runErr := cmd.Run() creds, parseErr := parseCredentials(out.Bytes()) // A cancelled ctx wins even over credentials that already arrived: the @@ -103,8 +92,15 @@ func Run(ctx context.Context, cfg Config) (Credentials, error) { } return creds, nil } - if ctx.Err() != nil { - return Credentials{}, fmt.Errorf("bootstrap: ssh interrupted: %w", context.Cause(ctx)) + if ctxErr := ctx.Err(); ctxErr != nil { + // A caller can match context.Canceled or DeadlineExceeded as well + // as a cause it set. A cause that already wraps the context error + // (or is it) is reported alone, so the text names it once. + cause := context.Cause(ctx) + if errors.Is(cause, ctxErr) { + return Credentials{}, fmt.Errorf("bootstrap: ssh interrupted: %w", cause) + } + return Credentials{}, fmt.Errorf("bootstrap: ssh interrupted: %w: %w", ctxErr, cause) } exitErr, isExit := errors.AsType[*exec.ExitError](runErr) // ErrWaitDelay means ssh exited 0 but a descendant kept its stdout open @@ -121,7 +117,28 @@ func Run(ctx context.Context, cfg Config) (Credentials, error) { if containsPiece(shown, passkey) { shown = []byte(withheldOutput) } - return Credentials{}, describeFailure(exitErr, parseErr, shown) + return Credentials{}, describeFailure(exitErr, parseErr, shown, out.dropped > 0) +} + +// newCommand prepares the ssh command. Its stdin and stderr are the process's +// own, so what ssh asks or reports on them reaches the user; its stdout is +// captured in the returned buffer. +func newCommand(ctx context.Context, sshPath string, args []string) (*exec.Cmd, *cappedBuffer) { + cmd := exec.CommandContext(ctx, sshPath, args...) + cmd.Stdin = os.Stdin + cmd.Stderr = os.Stderr + out := &cappedBuffer{max: maxOutput} + cmd.Stdout = out + cmd.Cancel = func() error { + // Interrupt first so ssh can restore the terminal; Windows cannot + // deliver os.Interrupt to another process, so kill there. + if err := cmd.Process.Signal(os.Interrupt); err != nil { + return cmd.Process.Kill() + } + return nil + } + cmd.WaitDelay = waitDelay + return cmd, out } const ( @@ -143,17 +160,21 @@ func containsPiece(out []byte, secret string) bool { } // describeFailure builds the single error the user sees when ssh finished -// without valid credentials: the parse problem, ssh's exit status, a hint for -// the common exit statuses, and a trimmed excerpt of stdout. -func describeFailure(exitErr *exec.ExitError, parseErr error, out []byte) error { +// without valid credentials: the parse problem, ssh's exit status or the +// signal that killed it, a hint for the common exit statuses, a note when +// output past maxOutput was dropped, and a trimmed excerpt of stdout. +func describeFailure(exitErr *exec.ExitError, parseErr error, out []byte, overflowed bool) error { var b strings.Builder b.WriteString("ssh ") - code := 0 + code, sig, killed := 0, "", false if exitErr != nil { - code = exitErr.ExitCode() - fmt.Fprintf(&b, "exited with status %d", code) + code = exitErr.ExitCode() // -1 when a signal ended ssh + sig, killed = killedBy(exitErr) + } + if killed { + fmt.Fprintf(&b, "was killed by signal %s", sig) } else { - b.WriteString("exited with status 0") + fmt.Fprintf(&b, "exited with status %d", code) } switch code { case 127: @@ -161,36 +182,45 @@ func describeFailure(exitErr *exec.ExitError, parseErr error, out []byte) error case 255: b.WriteString(": ssh could not connect or authenticate; check that plain ssh to this host works") } + if overflowed { + fmt.Fprintf(&b, "; output exceeded %d MiB; the credentials may have been cut off", maxOutput>>20) + } if ex := excerpt(out); ex != "" { fmt.Fprintf(&b, "; output: %q", ex) } return fmt.Errorf("%w: %s", parseErr, b.String()) } -// excerpt returns the last excerptLen bytes of the trimmed output before the -// first marker, prefixed with "..." when it was cut, so no part of a (possibly -// malformed) passkey is ever quoted. +// excerpt returns at most the last excerptLen bytes of the trimmed output +// before the first marker, prefixed with "..." when it was cut, so no part of +// a (possibly malformed) passkey is ever quoted. A cut never splits a UTF-8 +// sequence: the excerpt then starts at the next rune. func excerpt(out []byte) string { before, _, _ := bytes.Cut(out, []byte(marker)) - s := strings.TrimSpace(string(before)) - if len(s) > excerptLen { - s = "..." + s[len(s)-excerptLen:] + b := bytes.TrimSpace(before) + if len(b) <= excerptLen { + return string(b) } - return s + b = b[len(b)-excerptLen:] + for len(b) > 0 && !utf8.RuneStart(b[0]) { + b = b[1:] + } + return "..." + string(b) } -// cappedBuffer keeps the first max bytes written to it and silently drops the -// rest, so a chatty login shell cannot grow memory without bound. It never -// returns an error, so ssh never sees a broken pipe. +// cappedBuffer keeps the first max bytes written to it and drops the rest, +// counting them in dropped, so a chatty login shell cannot grow memory +// without bound. It never returns an error, so ssh never sees a broken pipe. type cappedBuffer struct { - buf bytes.Buffer - max int + buf bytes.Buffer + max int + dropped int // bytes written past max and discarded } func (c *cappedBuffer) Write(p []byte) (int, error) { - if room := c.max - c.buf.Len(); room > 0 { - c.buf.Write(p[:min(len(p), room)]) - } + kept := min(len(p), max(c.max-c.buf.Len(), 0)) + c.buf.Write(p[:kept]) + c.dropped += len(p) - kept return len(p), nil } diff --git a/internal/bootstrap/bootstrap_test.go b/internal/bootstrap/bootstrap_test.go index f012c66..0dcf51f 100644 --- a/internal/bootstrap/bootstrap_test.go +++ b/internal/bootstrap/bootstrap_test.go @@ -15,8 +15,10 @@ import ( "runtime" "slices" "strings" + "syscall" "testing" "time" + "unicode/utf8" ) // The test binary doubles as a fake ssh. When fakeSSHModeEnv is set, TestMain @@ -49,7 +51,7 @@ func runFakeSSH(mode string, args []string) int { return 91 } } - idpasskey := "IDPASSKEY:" + testID + "/" + testPasskey + "\n" + idpasskey := marker + testID + "/" + testPasskey + "\n" switch mode { case "ok": fmt.Print("Last login: Thu Sep 24 21:00:00 2026\nWelcome to the server\n" + idpasskey) @@ -64,7 +66,7 @@ func runFakeSSH(mode string, args []string) int { if m == nil { return 92 } - fmt.Print("IDPASSKEY:" + m[1] + "/" + m[2] + "\n") + fmt.Print(marker + m[1] + "/" + m[2] + "\n") return 0 case "notfound": // What bash prints when etterminal is not installed (it goes to stderr). @@ -76,8 +78,12 @@ func runFakeSSH(mode string, args []string) int { case "noise": fmt.Print("Welcome to the server\n") return 0 + case "flood": + // A login shell that prints more than maxOutput before the marker. + fmt.Print(strings.Repeat("x", maxOutput) + idpasskey) + return 0 case "malformed": - fmt.Print("IDPASSKEY:abcd\n") + fmt.Print(marker + "abcd\n") return 0 case "exit1": // A remote failure with some other status and no output. @@ -112,6 +118,18 @@ func runFakeSSH(mode string, args []string) int { case "hang": time.Sleep(time.Hour) return 0 + case "sigterm": + // Dies by SIGTERM, as an ssh killed by a signal would. Only the Unix + // tests use it: Windows cannot deliver the signal. + self, err := os.FindProcess(os.Getpid()) + if err != nil { + return 98 + } + if err := self.Signal(syscall.SIGTERM); err != nil { + return 99 + } + time.Sleep(10 * time.Second) + return 0 case "ok-hang": // Credentials arrive but ssh keeps running until it is cancelled. The // signaled file records that they were printed. @@ -177,6 +195,16 @@ func readArgv(t *testing.T, path string) []string { return args } +// readRemote returns the remote command the fake received: its last argument. +func readRemote(t *testing.T, path string) string { + t.Helper() + args := readArgv(t, path) + if len(args) == 0 { + t.Fatal("fake ssh recorded no arguments") + } + return args[len(args)-1] +} + func TestRunSuccess(t *testing.T) { cfg, argvPath := useFakeSSH(t, "ok") var logs bytes.Buffer @@ -195,12 +223,13 @@ func TestRunSuccess(t *testing.T) { } args := readArgv(t, argvPath) - if want := []string{"-oBatchMode=yes", "alice@example.test"}; len(args) != 3 || !slices.Equal(args[:2], want) { + want := []string{"-oBatchMode=yes", "-l", "alice", "--", "example.test"} + if len(args) != len(want)+1 || !slices.Equal(args[:len(want)], want) { t.Fatalf("ssh argv = %q, want %q followed by the remote command", args, want) } remote := regexp.MustCompile(`^echo 'XXX[A-Z2-7]{13}/[A-Z2-7]{32}_xterm-256color' \| etterminal --verbose=0$`) - if !remote.MatchString(args[2]) { - t.Fatalf("remote command = %q, want it to match %s", args[2], remote) + if last := args[len(args)-1]; !remote.MatchString(last) { + t.Fatalf("remote command = %q, want it to match %s", last, remote) } } @@ -222,9 +251,9 @@ func TestRunWarnsWhenServerDoesNotRegenerate(t *testing.T) { } // The expected values come from the command line the fake received, not // from the Credentials under test, so a broken accessor cannot hide a leak. - sent := regexp.MustCompile(`^echo '([A-Z2-7]{16})/([A-Z2-7]{32})_`).FindStringSubmatch(readArgv(t, argvPath)[2]) + sent := regexp.MustCompile(`^echo '([A-Z2-7]{16})/([A-Z2-7]{32})_`).FindStringSubmatch(readRemote(t, argvPath)) if len(sent) != 3 { - t.Fatalf("remote command %q does not carry a placeholder id and passkey", readArgv(t, argvPath)[2]) + t.Fatalf("remote command %q does not carry a placeholder id and passkey", readRemote(t, argvPath)) } if got.ID != sent[1] || got.Passkey() != sent[2] { t.Fatalf("Run() = %v with passkey match %v, want the placeholder sent in the remote command", got, got.Passkey() == sent[2]) @@ -311,9 +340,9 @@ func TestRunFailureRedactsPlaceholder(t *testing.T) { if !errors.Is(err, ErrNoCredentials) { t.Fatalf("Run() error = %v, want ErrNoCredentials", err) } - sent := regexp.MustCompile(`^echo '([A-Z2-7]{16})/([A-Z2-7]{32})_`).FindStringSubmatch(readArgv(t, argvPath)[2]) + sent := regexp.MustCompile(`^echo '([A-Z2-7]{16})/([A-Z2-7]{32})_`).FindStringSubmatch(readRemote(t, argvPath)) if len(sent) != 3 { - t.Fatalf("remote command %q does not carry a placeholder id and passkey", readArgv(t, argvPath)[2]) + t.Fatalf("remote command %q does not carry a placeholder id and passkey", readRemote(t, argvPath)) } if containsPiece([]byte(err.Error()), sent[2]) { t.Fatalf("error quotes passkey material: %v", err) @@ -368,11 +397,48 @@ func TestRunCancel(t *testing.T) { if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("Run() error = %v, want context.DeadlineExceeded", err) } + // Without a separate cause, the context error is wrapped once, not + // twice as ctx.Err() and a cause equal to it. + if n := strings.Count(err.Error(), context.DeadlineExceeded.Error()); n != 1 { + t.Fatalf("Run() error = %q names the context error %d times, want once", err, n) + } if elapsed := time.Since(start); elapsed > 10*time.Second { t.Fatalf("Run() took %v after cancel, want it bounded by waitDelay", elapsed) } } +// TestRunCancelWrapsErrAndCause: a caller that sets a cancellation cause can +// match both it and the context error it came with. +func TestRunCancelWrapsErrAndCause(t *testing.T) { + cfg, _ := useFakeSSH(t, "hang") + cause := errors.New("user gave up") + ctx, cancel := context.WithTimeoutCause(t.Context(), 200*time.Millisecond, cause) + defer cancel() + + _, err := Run(ctx, cfg) + if !errors.Is(err, context.DeadlineExceeded) || !errors.Is(err, cause) { + t.Fatalf("Run() error = %v, want it to match both context.DeadlineExceeded and the cause", err) + } +} + +// TestRunCancelCauseWrapsCtxErr: a cause that already wraps the context +// error is reported alone, so the context error is named once, and both +// still match. +func TestRunCancelCauseWrapsCtxErr(t *testing.T) { + cfg, _ := useFakeSSH(t, "hang") + ctx, cancel := context.WithCancelCause(t.Context()) + cause := fmt.Errorf("gave up: %w", context.Canceled) + cancel(cause) + + _, err := Run(ctx, cfg) + if !errors.Is(err, context.Canceled) || !errors.Is(err, cause) { + t.Fatalf("Run() error = %v, want it to match context.Canceled and the cause", err) + } + if n := strings.Count(err.Error(), context.Canceled.Error()); n != 1 { + t.Fatalf("Run() error = %q names the context error %d times, want once", err, n) + } +} + // TestRunCancelWinsOverCredentials: when the caller cancels while ssh is still // running, Run reports the cancellation even if the credentials already // arrived, so a Ctrl+C during bootstrap never goes on to connect. @@ -425,7 +491,7 @@ func TestRunSSHNotOnPath(t *testing.T) { } func TestExcerptCutsAtMarker(t *testing.T) { - out := []byte("banner\nIDPASSKEY:" + testID + "/" + testPasskey[:10]) + out := []byte("banner\n" + marker + testID + "/" + testPasskey[:10]) if got := excerpt(out); got != "banner" { t.Fatalf("excerpt() = %q, want %q", got, "banner") } @@ -460,6 +526,24 @@ func TestRunFindsSSHOnPath(t *testing.T) { } } +// TestNewCommandInheritsStdinAndStderr: ssh must share the process's stdin +// and stderr, or password and host key prompts never reach the user. +func TestNewCommandInheritsStdinAndStderr(t *testing.T) { + cmd, out := newCommand(t.Context(), "ssh", []string{"--", "host", "REMOTE"}) + if cmd.Stdin != os.Stdin { + t.Errorf("Stdin = %v, want os.Stdin", cmd.Stdin) + } + if cmd.Stderr != os.Stderr { + t.Errorf("Stderr = %v, want os.Stderr", cmd.Stderr) + } + if cmd.Stdout != out || out.max != maxOutput { + t.Errorf("Stdout is not the returned buffer capped at maxOutput") + } + if cmd.WaitDelay != waitDelay || cmd.Cancel == nil { + t.Errorf("WaitDelay = %v, Cancel set %v; want %v and a Cancel func", cmd.WaitDelay, cmd.Cancel != nil, waitDelay) + } +} + func TestCappedBuffer(t *testing.T) { c := &cappedBuffer{max: 4} for _, s := range []string{"ab", "cdef", "gh"} { @@ -470,4 +554,37 @@ func TestCappedBuffer(t *testing.T) { if got := string(c.Bytes()); got != "abcd" { t.Fatalf("Bytes() = %q, want %q", got, "abcd") } + if c.dropped != 4 { + t.Fatalf("dropped = %d, want 4 (\"ef\" and \"gh\")", c.dropped) + } +} + +// TestRunReportsOverflow: output past maxOutput is dropped, and when that +// leaves no credentials the error says the cap may have cut them off. +func TestRunReportsOverflow(t *testing.T) { + cfg, _ := useFakeSSH(t, "flood") + _, err := Run(t.Context(), cfg) + if !errors.Is(err, ErrNoCredentials) { + t.Fatalf("Run() error = %v, want ErrNoCredentials", err) + } + if want := "output exceeded 1 MiB; the credentials may have been cut off"; !strings.Contains(err.Error(), want) { + t.Fatalf("Run() error = %q, want it to contain %q", err, want) + } +} + +// TestExcerptKeepsRunesWhole: when the excerptLen cut lands inside a +// multi-byte rune, the excerpt starts at the next rune instead of quoting a +// stray continuation byte. +func TestExcerptKeepsRunesWhole(t *testing.T) { + const euro = "€" // three bytes in UTF-8 + tail := strings.Repeat("x", excerptLen-2) + // The last excerptLen bytes are the euro sign's two continuation bytes + // followed by tail. + got := excerpt([]byte("HEAD" + euro + tail)) + if !utf8.ValidString(got) { + t.Fatalf("excerpt() = %q, not valid UTF-8", got) + } + if want := "..." + tail; got != want { + t.Fatalf("excerpt() = %q, want %q", got, want) + } } diff --git a/internal/bootstrap/bootstrap_unix_test.go b/internal/bootstrap/bootstrap_unix_test.go index 76f5f78..32c201a 100644 --- a/internal/bootstrap/bootstrap_unix_test.go +++ b/internal/bootstrap/bootstrap_unix_test.go @@ -7,11 +7,30 @@ import ( "errors" "os" "path/filepath" + "strconv" + "strings" "syscall" "testing" "time" ) +// TestRunReportsKillingSignal: an ssh that dies by a signal has no exit +// status, so the error names the signal instead of "exited with status -1". +func TestRunReportsKillingSignal(t *testing.T) { + cfg, _ := useFakeSSH(t, "sigterm") + _, err := Run(t.Context(), cfg) + if !errors.Is(err, ErrNoCredentials) { + t.Fatalf("Run() error = %v, want ErrNoCredentials", err) + } + want := "ssh was killed by signal " + strconv.Itoa(int(syscall.SIGTERM)) + " (" + syscall.SIGTERM.String() + ")" + if !strings.Contains(err.Error(), want) { + t.Fatalf("Run() error = %q, want it to contain %q", err, want) + } + if strings.Contains(err.Error(), "status -1") { + t.Fatalf("Run() error = %q reports a status for a signalled ssh", err) + } +} + // TestRunCancelInterruptsThenKills pins both halves of Run's cancellation: ssh // is sent os.Interrupt first (so it can restore the terminal), and a child that // survives the interrupt is still killed once waitDelay has passed. Windows diff --git a/internal/bootstrap/command.go b/internal/bootstrap/command.go index 4ec9d0c..e804b51 100644 --- a/internal/bootstrap/command.go +++ b/internal/bootstrap/command.go @@ -2,7 +2,6 @@ package bootstrap import ( "fmt" - "regexp" "strings" "unicode" ) @@ -12,18 +11,46 @@ const ( defaultTerminalPath = "etterminal" ) -// termPattern and terminalPathPattern bound the two configurable values -// interpolated into the remote shell command. Neither admits a quote, space, -// glob, '$' or command separator; terminalPathPattern admits '~', which the -// remote shell expands on purpose so "~/bin/etterminal" works. +// termExtra and terminalPathExtra are the bytes besides ASCII letters and +// digits allowed in the two configurable values interpolated into the remote +// shell command. Neither admits a quote, space, glob, '$' or command +// separator; terminalPathExtra admits '~', which the remote shell expands on +// purpose so "~/bin/etterminal" works. // -// termPattern also excludes '_': etterminal splits its stdin line on '_' and +// termExtra also excludes '_': etterminal splits its stdin line on '_' and // aborts unless there are exactly two tokens, so a TERM containing '_' kills // the session start (upstream src/terminal/TerminalMain.cpp:111-124; MEASURED // against etserver 7.0.0: "Invalid number of tokens: 3", exit 134). -var ( - termPattern = regexp.MustCompile(`^[A-Za-z0-9.+-]+$`) - terminalPathPattern = regexp.MustCompile(`^[A-Za-z0-9._/~-]+$`) +const ( + termExtra = ".+-" + terminalPathExtra = "._/~-" +) + +// shellMeta holds the characters a Destination may not contain, and userMeta +// those a User may not contain, so neither value carries POSIX shell syntax +// into a ProxyCommand or Match exec that ssh expands %h or %r into (MEASURED +// against OpenSSH_10.0p2 on 2026-09-26: both commands received the host and +// the -l user substituted for %h and %r). OpenSSH 9.6 added its own hostname +// and user checks with the same aim (the CVE-2023-51385 fix). MEASURED +// against OpenSSH_10.0p2 on 2026-09-26 with +// ssh -G: a host with '$' and a user with ';', '(' or '"' or ending in '\' +// are refused, while the users CORP\alice, host$ and a$b are accepted. +// userMeta therefore leaves out '$' and '\', which winbind DOMAIN\user names +// and Samba machine accounts use, and validate refuses a trailing '\' +// instead. MEASURED against OpenSSH_for_Windows_9.5p2 on 2026-09-26: ssh -G +// accepts the user a;b and the host h$x, so on that client these sets are +// the only check. They model a POSIX shell; which interpreter the Windows +// client runs a ProxyCommand with is not measured. +// +// shellMeta also holds the glob characters, which a shell would expand +// against local file names: no host name contains them, and ssh keeps the +// brackets of a bracketed address as part of the host name (MEASURED against +// OpenSSH_10.0p2 on 2026-09-26: "ssh -G -- [::1]" gives hostname "[::1]"), so +// an IPv6 destination is passed bare. userMeta leaves them out, as OpenSSH +// accepts them in a user name. +const ( + shellMeta = "'`\"$\\;&<>|(){}*?[]" + userMeta = "'`\";&<>|(){}" ) // applyDefaults fills in empty optional fields. Run calls it on its own copy @@ -44,19 +71,33 @@ func (cfg *Config) validate() error { case cfg.Destination == "": return fmt.Errorf("%w: destination is empty", ErrInvalidConfig) case strings.HasPrefix(cfg.Destination, "-"): - // ssh would parse it as an option. + // sshArgs puts "--" before the destination, so ssh no longer reads it + // as an option, but ssh still substitutes it for %h in a ProxyCommand + // or Match exec (measured, see shellMeta), where a command such as + // "nc %h %p" would take it as an option. return fmt.Errorf("%w: destination %q starts with '-'", ErrInvalidConfig, cfg.Destination) case hasSpaceOrControl(cfg.Destination): return fmt.Errorf("%w: destination %q contains whitespace or control characters", ErrInvalidConfig, cfg.Destination) - case strings.ContainsRune(cfg.User, '@'), strings.HasPrefix(cfg.User, "-"), hasSpaceOrControl(cfg.User): + case strings.ContainsAny(cfg.Destination, shellMeta): + return fmt.Errorf("%w: destination %q contains a shell metacharacter", ErrInvalidConfig, cfg.Destination) + case strings.HasPrefix(cfg.User, "-"), hasSpaceOrControl(cfg.User): + // '@' is allowed: the user goes to ssh as its own -l argument, so + // "alice@corp.example" stays one user name. return fmt.Errorf("%w: user %q is not a valid user name", ErrInvalidConfig, cfg.User) + case strings.ContainsAny(cfg.User, userMeta): + return fmt.Errorf("%w: user %q contains a shell metacharacter", ErrInvalidConfig, cfg.User) + case strings.HasSuffix(cfg.User, `\`): + return fmt.Errorf("%w: user %q ends in a backslash", ErrInvalidConfig, cfg.User) case cfg.User != "" && strings.ContainsRune(cfg.Destination, '@'): - // ssh would read "user@alice@host" as user "user@alice", not what either value meant. + // ssh would silently let -l win over the user in the destination + // (MEASURED against OpenSSH_10.0p2 on 2026-09-26: "ssh -G -l alice + // bob@h2.example" resolves user alice), so one of the two values + // would be ignored. return fmt.Errorf("%w: user %q given but destination %q already names a user", ErrInvalidConfig, cfg.User, cfg.Destination) - case !termPattern.MatchString(cfg.Term): - return fmt.Errorf("%w: TERM %q must match %s", ErrInvalidConfig, cfg.Term, termPattern) - case !terminalPathPattern.MatchString(cfg.TerminalPath): - return fmt.Errorf("%w: terminal path %q must match %s", ErrInvalidConfig, cfg.TerminalPath, terminalPathPattern) + case !onlyAlnumOr(cfg.Term, termExtra): + return fmt.Errorf("%w: TERM %q must be letters, digits and %s only", ErrInvalidConfig, cfg.Term, termExtra) + case !onlyAlnumOr(cfg.TerminalPath, terminalPathExtra): + return fmt.Errorf("%w: terminal path %q must be letters, digits and %s only", ErrInvalidConfig, cfg.TerminalPath, terminalPathExtra) } for _, opt := range cfg.SSHOptions { if opt == "" || hasControl(opt) { @@ -75,17 +116,29 @@ func remoteCommand(id, passkey, term, terminalPath string) string { } // sshArgs builds ssh's argument list: every option as its own "-o" -// argument, then [user@]destination, then the remote command as one argument. +// argument, then "-l " when User is set, then "--", the destination and +// the remote command as one argument. +// +// The user goes through -l rather than as user@destination so a user name +// holding '@' stays whole and a URI destination is passed untouched +// (MEASURED against OpenSSH_10.0p2 on 2026-09-26 with ssh -G: +// "-l alice@corp.example -- h3.example" resolves user alice@corp.example and +// "-l alice ssh://h1.example:2222" user alice, host h1.example, port 2222). +// +// "--" ends option parsing, so nothing from the destination onwards is read +// as an option, while the -o options before it still apply (MEASURED against +// OpenSSH_10.0p2 on 2026-09-26: "ssh -G -oPort=2200 -- h4.example" keeps port +// 2200, and "ssh -G -- h5.example -p 2201" keeps port 22 where +// "ssh -G h5.example -p 2201" gives 2201). func sshArgs(cfg *Config, remote string) []string { - dest := cfg.Destination - if cfg.User != "" { - dest = cfg.User + "@" + dest - } - args := make([]string, 0, len(cfg.SSHOptions)+2) + args := make([]string, 0, len(cfg.SSHOptions)+5) for _, opt := range cfg.SSHOptions { args = append(args, "-o"+opt) } - return append(args, dest, remote) + if cfg.User != "" { + args = append(args, "-l", cfg.User) + } + return append(args, "--", cfg.Destination, remote) } func hasSpaceOrControl(s string) bool { diff --git a/internal/bootstrap/command_test.go b/internal/bootstrap/command_test.go index 00bc3e4..a623660 100644 --- a/internal/bootstrap/command_test.go +++ b/internal/bootstrap/command_test.go @@ -3,6 +3,7 @@ package bootstrap import ( "errors" "slices" + "strings" "testing" ) @@ -10,11 +11,12 @@ func TestValidate(t *testing.T) { valid := Config{Destination: "example.test"} valid.applyDefaults() - tests := []struct { + type validateCase struct { name string mutate func(*Config) wantErr bool - }{ + } + tests := []validateCase{ {"defaults", func(*Config) {}, false}, {"alias with dots and dashes", func(c *Config) { c.Destination = "my-host.lan" }, false}, {"user", func(c *Config) { c.User = "alice" }, false}, @@ -26,7 +28,13 @@ func TestValidate(t *testing.T) { {"destination is an option", func(c *Config) { c.Destination = "-oProxyCommand=evil" }, true}, {"destination with space", func(c *Config) { c.Destination = "a b" }, true}, {"destination with control", func(c *Config) { c.Destination = "host\x01" }, true}, - {"user with at", func(c *Config) { c.User = "a@b" }, true}, + {"IPv6", func(c *Config) { c.Destination = "::1" }, false}, + {"IPv6 with zone", func(c *Config) { c.Destination = "fe80::1%eth0" }, false}, + {"ssh URI with port", func(c *Config) { c.Destination = "ssh://host:2222" }, false}, + {"dotted alias", func(c *Config) { c.Destination = "my-host.example" }, false}, + {"underscore host", func(c *Config) { c.Destination = "host_1" }, false}, + {"user with at", func(c *Config) { c.User = "alice@corp.example" }, false}, + {"user with dot dash underscore", func(c *Config) { c.User = "a.b-c_d" }, false}, {"user with space", func(c *Config) { c.User = "a b" }, true}, {"user with control", func(c *Config) { c.User = "a\x01" }, true}, {"user is an option", func(c *Config) { c.User = "-x" }, true}, @@ -41,7 +49,32 @@ func TestValidate(t *testing.T) { {"empty ssh option", func(c *Config) { c.SSHOptions = []string{""} }, true}, {"ssh option with newline", func(c *Config) { c.SSHOptions = []string{"A=b\nc"} }, true}, } - for _, tt := range tests { + // Spelled out rather than taken from shellMeta and userMeta, so dropping + // a character from a production set turns a row red. + const ( + wantDestMeta = "'`\"$\\;&<>|(){}*?[]" + wantUserMeta = "'`\";&<>|(){}" + ) + metaCases := make([]validateCase, 0, len(wantDestMeta)+len(wantUserMeta)) + for _, r := range wantDestMeta { + metaCases = append(metaCases, + validateCase{"destination with " + string(r), func(c *Config) { c.Destination = "h" + string(r) + "x.example" }, true}) + } + for _, r := range wantUserMeta { + metaCases = append(metaCases, + validateCase{"user with " + string(r), func(c *Config) { c.User = "a" + string(r) + "b" }, true}) + } + // User names OpenSSH accepts and main accepted: a winbind DOMAIN\user and + // a Samba machine account. A trailing backslash is refused, as OpenSSH + // refuses it. + metaCases = append(metaCases, + validateCase{"winbind user", func(c *Config) { c.User = `CORP\alice` }, false}, + validateCase{"machine account user", func(c *Config) { c.User = "host$" }, false}, + validateCase{"user with dollar inside", func(c *Config) { c.User = "a$b" }, false}, + validateCase{"user ending in backslash", func(c *Config) { c.User = `alice\` }, true}, + ) + for i, tt := range slices.Concat(tests, metaCases) { + isMeta := i >= len(tests) t.Run(tt.name, func(t *testing.T) { cfg := valid cfg.SSHOptions = slices.Clone(valid.SSHOptions) @@ -53,6 +86,12 @@ func TestValidate(t *testing.T) { if err != nil && !errors.Is(err, ErrInvalidConfig) { t.Fatalf("validate() = %v, want it to wrap ErrInvalidConfig", err) } + // A metacharacter row must fail for the metacharacter, not for + // some other rule that happens to reject the same value. + if isMeta && tt.wantErr && + !strings.Contains(err.Error(), "shell metacharacter") && !strings.Contains(err.Error(), "backslash") { + t.Fatalf("validate() = %v, want the metacharacter reason", err) + } }) } } @@ -87,12 +126,22 @@ func TestSSHArgs(t *testing.T) { { name: "destination only", cfg: Config{Destination: "host"}, - want: []string{"host", "REMOTE"}, + want: []string{"--", "host", "REMOTE"}, }, { name: "user and options", cfg: Config{Destination: "host", User: "alice", SSHOptions: []string{"BatchMode=yes", "Port=2222"}}, - want: []string{"-oBatchMode=yes", "-oPort=2222", "alice@host", "REMOTE"}, + want: []string{"-oBatchMode=yes", "-oPort=2222", "-l", "alice", "--", "host", "REMOTE"}, + }, + { + name: "user with at", + cfg: Config{Destination: "host", User: "alice@corp.example"}, + want: []string{"-l", "alice@corp.example", "--", "host", "REMOTE"}, + }, + { + name: "user and ssh URI", + cfg: Config{Destination: "ssh://host:2222", User: "bob"}, + want: []string{"-l", "bob", "--", "ssh://host:2222", "REMOTE"}, }, } for _, tt := range tests { diff --git a/internal/bootstrap/parse.go b/internal/bootstrap/parse.go index 6af37b1..1566253 100644 --- a/internal/bootstrap/parse.go +++ b/internal/bootstrap/parse.go @@ -3,6 +3,7 @@ package bootstrap import ( "bytes" "fmt" + "strings" ) // marker precedes the credentials in etterminal's output @@ -21,24 +22,40 @@ func parseCredentials(out []byte) (Credentials, error) { if !found { return Credentials{}, ErrNoCredentials } - id, n := alnumRun(rest) + id := alnumRun(rest) + n := len(id) if n != idLen || len(rest) == n || rest[n] != '/' { return Credentials{}, fmt.Errorf("%w: malformed id after %q", ErrNoCredentials, marker) } - passkey, m := alnumRun(rest[n+1:]) - if m != passkeyLen { + passkey := alnumRun(rest[n+1:]) + if len(passkey) != passkeyLen { // Do not echo the value: a malformed passkey may still be a real one. return Credentials{}, fmt.Errorf("%w: malformed passkey after %q", ErrNoCredentials, marker) } return NewCredentials(string(id), string(passkey)), nil } -// alnumRun returns the leading run of ASCII letters and digits in b and its length. -func alnumRun(b []byte) (run []byte, n int) { +// alnumRun returns the leading run of ASCII letters and digits in b. +func alnumRun(b []byte) []byte { + n := 0 for n < len(b) && isAlnum(b[n]) { n++ } - return b[:n], n + return b[:n] +} + +// onlyAlnumOr reports whether s is non-empty and every byte of it is an +// ASCII letter, an ASCII digit or one of the bytes in extra. +func onlyAlnumOr(s, extra string) bool { + if s == "" { + return false + } + for i := range len(s) { + if !isAlnum(s[i]) && strings.IndexByte(extra, s[i]) < 0 { + return false + } + } + return true } func isAlnum(c byte) bool { diff --git a/internal/bootstrap/parse_test.go b/internal/bootstrap/parse_test.go index 0fa7ece..c4ba438 100644 --- a/internal/bootstrap/parse_test.go +++ b/internal/bootstrap/parse_test.go @@ -1,34 +1,96 @@ package bootstrap import ( + "bytes" "errors" "strings" "testing" ) +func TestOnlyAlnumOr(t *testing.T) { + tests := []struct { + s, extra string + want bool + }{ + {"abcXYZ019", "", true}, + {"a.b-c", ".-", true}, + {"", ".-", false}, // empty is not a valid value + {"a_b", ".-", false}, // a byte outside extra + {"café", ".-", false}, // non-ASCII letters are not allowed + {"a b", ".-", false}, + } + for _, tt := range tests { + if got := onlyAlnumOr(tt.s, tt.extra); got != tt.want { + t.Errorf("onlyAlnumOr(%q, %q) = %v, want %v", tt.s, tt.extra, got, tt.want) + } + } +} + +// TestMarkerIsUpstream pins the marker to what etterminal prints (upstream +// src/terminal/TerminalMain.cpp:185 at et-v7.0.0); the other tests build +// their input from the constant. +func TestMarkerIsUpstream(t *testing.T) { + if marker != "IDPASSKEY:" { + t.Fatalf("marker = %q, want %q", marker, "IDPASSKEY:") + } +} + +// FuzzParseCredentials: parseCredentials never panics, and whatever it +// accepts is a well-formed id and passkey that directly follow the first +// marker in the input. +func FuzzParseCredentials(f *testing.F) { + for _, s := range []string{ + marker + testID + "/" + testPasskey + "\n", + "banner\n" + marker + testID + "/" + testPasskey, + marker + testID + "/" + testPasskey + "X", + marker + "abcd", + marker, + "", + } { + f.Add([]byte(s)) + } + f.Fuzz(func(t *testing.T, out []byte) { + c, err := parseCredentials(out) + if err != nil { + return + } + id, passkey := c.ID, c.Passkey() + if len(id) != idLen || !onlyAlnumOr(id, "") { + t.Fatalf("parseCredentials(%q) id = %q, want %d letters and digits", out, id, idLen) + } + if len(passkey) != passkeyLen || !onlyAlnumOr(passkey, "") { + t.Fatalf("parseCredentials(%q) passkey has length %d or a non-alphanumeric byte, want %d letters and digits", out, len(passkey), passkeyLen) + } + _, rest, _ := bytes.Cut(out, []byte(marker)) + if !bytes.HasPrefix(rest, []byte(id+"/"+passkey)) { + t.Fatalf("parseCredentials(%q) returned credentials that do not follow the first marker", out) + } + }) +} + func TestParseCredentials(t *testing.T) { tests := []struct { name string out string wantErr string // "" means success with testID and testPasskey; otherwise a substring the error must contain }{ - {"plain", "IDPASSKEY:" + testID + "/" + testPasskey + "\n", ""}, - {"crlf", "IDPASSKEY:" + testID + "/" + testPasskey + "\r\n", ""}, - {"no trailing newline", "IDPASSKEY:" + testID + "/" + testPasskey, ""}, - {"noise before", "Last login: Thu\nWelcome!\nIDPASSKEY:" + testID + "/" + testPasskey + "\n", ""}, - {"first marker wins", "IDPASSKEY:" + testID + "/" + testPasskey + "\nIDPASSKEY:zzzzzzzzzzzzzzzz/" + testPasskey, ""}, + {"plain", marker + testID + "/" + testPasskey + "\n", ""}, + {"crlf", marker + testID + "/" + testPasskey + "\r\n", ""}, + {"no trailing newline", marker + testID + "/" + testPasskey, ""}, + {"noise before", "Last login: Thu\nWelcome!\n" + marker + testID + "/" + testPasskey + "\n", ""}, + {"first marker wins", marker + testID + "/" + testPasskey + "\n" + marker + "zzzzzzzzzzzzzzzz/" + testPasskey, ""}, {"no marker", "Welcome!\n", "no IDPASSKEY"}, {"empty", "", "no IDPASSKEY"}, - {"truncated id", "IDPASSKEY:abcd", "malformed id"}, - {"id at end of output", "IDPASSKEY:" + testID, "malformed id"}, - {"colon instead of slash", "IDPASSKEY:" + testID + ":" + testPasskey, "malformed id"}, - {"truncated passkey", "IDPASSKEY:" + testID + "/0123", "malformed passkey"}, - {"id too long", "IDPASSKEY:" + testID + "X/" + testPasskey, "malformed id"}, - {"id too short", "IDPASSKEY:" + testID[1:] + "/" + testPasskey, "malformed id"}, - {"passkey too long", "IDPASSKEY:" + testID + "/" + testPasskey + "X\n", "malformed passkey"}, - {"missing slash", "IDPASSKEY:" + testID + testPasskey, "malformed id"}, - {"non alphanumeric in id", "IDPASSKEY:abcdEFGH1234567-/" + testPasskey, "malformed id"}, - {"non alphanumeric in passkey", "IDPASSKEY:" + testID + "/0123456789abcdef_BCDEF0123456789", "malformed passkey"}, + {"truncated id", marker + "abcd", "malformed id"}, + {"id at end of output", marker + testID, "malformed id"}, + {"colon instead of slash", marker + testID + ":" + testPasskey, "malformed id"}, + {"truncated passkey", marker + testID + "/0123", "malformed passkey"}, + {"id too long", marker + testID + "X/" + testPasskey, "malformed id"}, + {"id too short", marker + testID[1:] + "/" + testPasskey, "malformed id"}, + {"passkey too long", marker + testID + "/" + testPasskey + "X\n", "malformed passkey"}, + {"missing slash", marker + testID + testPasskey, "malformed id"}, + {"non alphanumeric in id", marker + "abcdEFGH1234567-/" + testPasskey, "malformed id"}, + {"non alphanumeric in passkey", marker + testID + "/0123456789abcdef_BCDEF0123456789", "malformed passkey"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { diff --git a/internal/bootstrap/signal_other.go b/internal/bootstrap/signal_other.go new file mode 100644 index 0000000..604472d --- /dev/null +++ b/internal/bootstrap/signal_other.go @@ -0,0 +1,9 @@ +//go:build !unix + +package bootstrap + +import "os/exec" + +// killedBy reports false: outside Unix a process has no terminating signal +// to name, so describeFailure falls back to the exit status. +func killedBy(*exec.ExitError) (string, bool) { return "", false } diff --git a/internal/bootstrap/signal_unix.go b/internal/bootstrap/signal_unix.go new file mode 100644 index 0000000..60f1195 --- /dev/null +++ b/internal/bootstrap/signal_unix.go @@ -0,0 +1,20 @@ +//go:build unix + +package bootstrap + +import ( + "os/exec" + "strconv" + "syscall" +) + +// killedBy reports the signal that ended the process behind exitErr, as +// " ()", and whether a signal ended it at all. +func killedBy(exitErr *exec.ExitError) (string, bool) { + ws, ok := exitErr.Sys().(syscall.WaitStatus) + if !ok || !ws.Signaled() { + return "", false + } + sig := ws.Signal() + return strconv.Itoa(int(sig)) + " (" + sig.String() + ")", true +} diff --git a/internal/console/console.go b/internal/console/console.go index eab80a7..aee99c1 100644 --- a/internal/console/console.go +++ b/internal/console/console.go @@ -8,7 +8,11 @@ // control (SIGTTIN, SIGTTOU) until it is brought to the foreground. package console -import "errors" +import ( + "context" + "errors" + "os" +) // Size is the terminal size in character cells, plus pixels when the // platform reports them (0 otherwise). @@ -20,3 +24,21 @@ type Size struct { // 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") + +// resizeCheck reads the size once and yields it if it changed since *last. +// It reports false when the Console is closed, when ctx has ended with a +// changed size pending, or when yield stops; an unchanged size or any other +// Size error reports true, even after ctx ended, and the caller's select on +// ctx ends the loop. Checking ctx before yielding keeps a change from being +// yielded after cancellation, whether or not a select raced it. +func resizeCheck(ctx context.Context, size func() (Size, error), last *Size, yield func(Size) bool) bool { + sz, err := size() + if errors.Is(err, os.ErrClosed) { + return false + } + if err != nil || sz == *last { + return true + } + *last = sz + return ctx.Err() == nil && yield(sz) +} diff --git a/internal/console/console_linux_test.go b/internal/console/console_linux_test.go index 1589c5d..18686dd 100644 --- a/internal/console/console_linux_test.go +++ b/internal/console/console_linux_test.go @@ -5,6 +5,7 @@ import ( "errors" "io" "io/fs" + "iter" "os" "path/filepath" "testing" @@ -101,6 +102,25 @@ func TestOpenWithoutControllingTerminal(t *testing.T) { } } +// TestOpenReportsOtherOpenErrorsPlainly covers a tty path that fails to +// open for a reason other than a missing controlling terminal (here +// EISDIR): the error keeps its cause and is not ErrNotTerminal, which +// would misreport it as "not a terminal". +func TestOpenReportsOtherOpenErrorsPlainly(t *testing.T) { + _, slave := openPTY(t) + c, err := open(t.TempDir(), slave, slave) + if err == nil { + _ = c.Close() + t.Fatal("open(directory) succeeded, want an error") + } + if errors.Is(err, ErrNotTerminal) { + t.Fatalf("open(directory) error = %v, want an error that is not ErrNotTerminal", err) + } + if !errors.Is(err, unix.EISDIR) { + t.Fatalf("open(directory) error = %v, want the underlying EISDIR kept", err) + } +} + func TestOpenAcceptsTerminal(t *testing.T) { _, slave := openPTY(t) c := openThroughPath(t, slave) @@ -578,15 +598,7 @@ func TestResizesYieldsChangesOnly(t *testing.T) { 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 - } - }() + sizes, done := drainResizes(c.Resizes(ctx)) // baseline 24x80 is taken here tick := time.NewTicker(20 * time.Millisecond) defer tick.Stop() @@ -660,15 +672,7 @@ func TestResizesYieldsChangeBeforeRangingStarts(t *testing.T) { // 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 - } - }() - + sizes, done := drainResizes(resizes) select { case sz := <-sizes: if sz != (Size{Rows: 50, Cols: 120}) { @@ -681,6 +685,21 @@ func TestResizesYieldsChangeBeforeRangingStarts(t *testing.T) { waitDone(t, done, "Resizes after cancel") } +// drainResizes ranges over resizes on its own goroutine, sending each size +// yielded to sizes, and closes done when the range statement finishes. +// sizes is buffered so a test can count yields after the range ended. +func drainResizes(resizes iter.Seq[Size]) (sizes <-chan Size, done <-chan struct{}) { + out := make(chan Size, 4) + finished := make(chan struct{}) + go func() { + defer close(finished) + for sz := range resizes { + out <- sz + } + }() + return out, finished +} + // waitDone waits, bounded, for a consumer goroutine to close done. func waitDone(t *testing.T, done <-chan struct{}, what string) { t.Helper() @@ -705,14 +724,7 @@ func TestResizesStopsAfterCancel(t *testing.T) { 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 - } - }() + sizes, done := drainResizes(resizes) waitDone(t, done, "cancelled Resizes") if len(sizes) != 0 { t.Fatalf("Resizes yielded %+v after its context ended", <-sizes) @@ -834,14 +846,7 @@ func TestResizesEndsAfterClose(t *testing.T) { } } - sizes := make(chan Size, 4) - done := make(chan struct{}) - go func() { - defer close(done) - for sz := range resizes { - sizes <- sz - } - }() + sizes, done := drainResizes(resizes) if tc.inLoop { // Signal an unchanged size for 300 ms so the iterator has diff --git a/internal/console/console_test.go b/internal/console/console_test.go new file mode 100644 index 0000000..ab779be --- /dev/null +++ b/internal/console/console_test.go @@ -0,0 +1,59 @@ +package console + +import ( + "context" + "errors" + "os" + "testing" +) + +// TestResizeCheck pins the shared step of both platforms' Resizes loops for +// the size results it can meet (unchanged, changed, closed, another error), +// a changed or unchanged size after ctx ended, and a loop body that stops, +// without a terminal. +func TestResizeCheck(t *testing.T) { + base := Size{Rows: 24, Cols: 80} + changed := Size{Rows: 40, Cols: 120} + live := t.Context() + ended, cancel := context.WithCancel(t.Context()) + cancel() + + tests := []struct { + name string + ctx context.Context + size Size + sizeErr error + yieldMore bool + want bool + wantYield bool + wantLast Size + }{ + {name: "unchanged", ctx: live, size: base, yieldMore: true, want: true, wantLast: base}, + {name: "changed", ctx: live, size: changed, yieldMore: true, want: true, wantYield: true, wantLast: changed}, + {name: "changed and the body stops", ctx: live, size: changed, want: false, wantYield: true, wantLast: changed}, + {name: "closed", ctx: live, sizeErr: os.ErrClosed, yieldMore: true, want: false, wantLast: base}, + {name: "other size error", ctx: live, sizeErr: errors.New("transient"), yieldMore: true, want: true, wantLast: base}, + {name: "changed after ctx ended", ctx: ended, size: changed, yieldMore: true, want: false, wantLast: changed}, + {name: "unchanged after ctx ended", ctx: ended, size: base, yieldMore: true, want: true, wantLast: base}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + last := base + yielded := false + got := resizeCheck(tt.ctx, + func() (Size, error) { return tt.size, tt.sizeErr }, + &last, + func(sz Size) bool { + if sz != tt.size { + t.Errorf("yielded %+v, want %+v", sz, tt.size) + } + yielded = true + return tt.yieldMore + }) + if got != tt.want || yielded != tt.wantYield || last != tt.wantLast { + t.Fatalf("resizeCheck = %v, yielded %v, last %+v; want %v, %v, %+v", + got, yielded, last, tt.want, tt.wantYield, tt.wantLast) + } + }) + } +} diff --git a/internal/console/console_unix.go b/internal/console/console_unix.go index 39ed90f..8e3f467 100644 --- a/internal/console/console_unix.go +++ b/internal/console/console_unix.go @@ -45,7 +45,9 @@ type Console struct { // 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. +// the process has no controlling terminal. Any other failure to open +// /dev/tty, such as permission denied, is returned as a plain error that +// keeps its cause. // // Open records the terminal state as the baseline that restore returns to. // The intended caller opens the console before running anything that may @@ -63,7 +65,12 @@ func open(ttyPath string, stdin, stdout *os.File) (*Console, error) { } tty, err := os.OpenFile(ttyPath, os.O_RDWR, 0) if err != nil { - return nil, fmt.Errorf("%w: open %s: %w", ErrNotTerminal, ttyPath, err) + // err is an *fs.PathError, whose text already names the operation + // and the path. + if noTerminalToOpen(err) { + return nil, fmt.Errorf("%w: %w", ErrNotTerminal, err) + } + return nil, fmt.Errorf("console: %w", err) } c, err := newConsole(tty) if err != nil { @@ -73,19 +80,22 @@ func open(ttyPath string, stdin, stdout *os.File) (*Console, error) { return c, nil } +// noTerminalToOpen reports whether a failed open of the controlling terminal +// means there is none: ENXIO (the process has no controlling terminal) or +// ENOENT (no such device node). Anything else, such as EACCES, is a plain +// failure. +func noTerminalToOpen(err error) bool { + return errors.Is(err, unix.ENXIO) || errors.Is(err, unix.ENOENT) +} + // 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 - }) + base, err := controlValue(tty, term.GetState) if err != nil { return nil, fmt.Errorf("%w: read terminal state: %w", ErrNotTerminal, err) } - return c, nil + return &Console{tty: tty, base: base, setState: term.Restore}, nil } // MakeRaw switches the terminal to raw mode. It may be called again after @@ -169,11 +179,8 @@ func (c *Console) Size() (Size, error) { 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 + ws, err := controlValue(c.tty, func(fd int) (*unix.Winsize, error) { + return unix.IoctlGetWinsize(fd, unix.TIOCGWINSZ) }) if err != nil { return Size{}, fmt.Errorf("console: size: %w", err) @@ -212,16 +219,9 @@ func (c *Console) Resizes(ctx context.Context) iter.Seq[Size] { // 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) { + if !resizeCheck(ctx, c.Size, &last, yield) { return } - if err == nil && sz != last { - last = sz - if ctx.Err() != nil || !yield(sz) { - return - } - } for { select { @@ -229,17 +229,7 @@ func (c *Console) Resizes(ctx context.Context) iter.Seq[Size] { 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) { + if !resizeCheck(ctx, c.Size, &last, yield) { return } } @@ -289,12 +279,21 @@ func controlFile(file *os.File, f func(fd int) error) error { return ferr } +// controlValue runs f with file's descriptor and returns its result. +func controlValue[T any](file *os.File, f func(fd int) (T, error)) (T, error) { + var out T + err := controlFile(file, func(fd int) error { + var ferr error + out, ferr = f(fd) + return ferr + }) + return out, err +} + // 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 + ok, err := controlValue(f, func(fd int) (bool, error) { + return term.IsTerminal(fd), nil }) return err == nil && ok } diff --git a/internal/console/console_unix_test.go b/internal/console/console_unix_test.go new file mode 100644 index 0000000..cf64a59 --- /dev/null +++ b/internal/console/console_unix_test.go @@ -0,0 +1,35 @@ +//go:build unix + +package console + +import ( + "io/fs" + "testing" + + "golang.org/x/sys/unix" +) + +// TestNoTerminalToOpen pins which open failures mean there is no terminal: +// ENXIO, what opening /dev/tty gives a process without a controlling +// terminal, which the package tests cannot reach from a process that has +// one, and ENOENT; the rest are plain failures. +func TestNoTerminalToOpen(t *testing.T) { + tests := []struct { + errno unix.Errno + want bool + }{ + {unix.ENXIO, true}, + {unix.ENOENT, true}, + {unix.EACCES, false}, + {unix.EISDIR, false}, + {unix.EMFILE, false}, + } + for _, tt := range tests { + t.Run(tt.errno.Error(), func(t *testing.T) { + err := &fs.PathError{Op: "open", Path: "/dev/tty", Err: tt.errno} + if got := noTerminalToOpen(err); got != tt.want { + t.Fatalf("noTerminalToOpen(%v) = %v, want %v", err, got, tt.want) + } + }) + } +} diff --git a/internal/console/console_windows.go b/internal/console/console_windows.go index 6dddd28..4e0e475 100644 --- a/internal/console/console_windows.go +++ b/internal/console/console_windows.go @@ -70,14 +70,19 @@ type Console struct { // (GetConsoleScreenBufferInfo). nil means the real call. sizeFn func(h windows.Handle, info *windows.ConsoleScreenBufferInfo) error - // Reader-owned state. - dec utf16Decoder - units []uint16 - pending []byte + // Reader-owned state. pending holds decoded bytes not yet returned, + // from pendingOff on; it is reset to empty, keeping its capacity, once + // Read has returned all of it. + dec utf16Decoder + units []uint16 + unitsRead uint32 // units the last console read returned + pending []byte + pendingOff int // Writer-owned state. - enc utf8Encoder - buf []uint16 + enc utf8Encoder + buf []uint16 + unitsWritten uint32 // units the last console write reported } // Open returns the console attached to stdin and stdout. It returns an error @@ -248,8 +253,11 @@ func (c *Console) Read(p []byte) (int, error) { c.reading = true c.mu.Unlock() - var n uint32 - err := read(c.in, &c.units[0], uint32(len(c.units)), &n, nil) + // The count goes into a field: a local whose address is passed to + // the read func value is moved to the heap, one allocation per + // console read (go build -gcflags=-m, go1.27). + c.unitsRead = 0 + err := read(c.in, &c.units[0], uint32(len(c.units)), &c.unitsRead, nil) c.mu.Lock() c.reading = false @@ -265,10 +273,14 @@ func (c *Console) Read(p []byte) (int, error) { if err != nil { return 0, fmt.Errorf("console: read: %w", err) } - c.pending = c.dec.append(c.pending[:0], c.units[:n]) + c.pending = c.dec.append(c.pending[:0], c.units[:c.unitsRead]) + } + n := copy(p, c.pending[c.pendingOff:]) + c.pendingOff += n + if c.pendingOff == len(c.pending) { + // Drained: keep the whole capacity for the next decode. + c.pending, c.pendingOff = c.pending[:0], 0 } - n := copy(p, c.pending) - c.pending = c.pending[n:] return n, nil } @@ -308,10 +320,14 @@ func (c *Console) Write(p []byte) (int, error) { 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 { + // The count goes into a field, as in Read: a local whose address + // is passed to the write func value is moved to the heap, one + // allocation per console write (go build -gcflags=-m, go1.27). + c.unitsWritten = 0 + if err := write(c.out, &chunk[0], uint32(len(chunk)), &c.unitsWritten, nil); err != nil { return 0, fmt.Errorf("console: write: %w", err) } + n := c.unitsWritten if n == 0 { return 0, io.ErrShortWrite } @@ -368,17 +384,7 @@ func (c *Console) Resizes(ctx context.Context) iter.Seq[Size] { 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) { + if !resizeCheck(ctx, c.Size, &last, yield) { return } } diff --git a/internal/console/console_windows_test.go b/internal/console/console_windows_test.go index f9981bf..4de8829 100644 --- a/internal/console/console_windows_test.go +++ b/internal/console/console_windows_test.go @@ -9,6 +9,7 @@ import ( "strings" "sync/atomic" "testing" + "testing/synctest" "time" "unicode/utf16" "unsafe" @@ -946,6 +947,64 @@ func TestReadSmallBufferKeepsRest(t *testing.T) { } } +// TestReadReusesPendingBuffer needs no console: draining the decoded bytes +// in small Reads keeps the buffer's capacity, so a steady stream of console +// reads decodes into the same buffer without allocating. The decode is 8 +// bytes, a whole allocation size class, so a buffer that lost the capacity +// in front of the bytes already handed out would have to grow on every +// decode. +func TestReadReusesPendingBuffer(t *testing.T) { + c := newConsole(0, 0) + units := utf16.Encode([]rune("abcdefgh")) + c.readFn = func(_ windows.Handle, buf *uint16, toread uint32, read *uint32, _ *byte) error { + *read = uint32(copy(unsafe.Slice(buf, toread), units)) + return nil + } + p := make([]byte, len(units)) + drain := func() { + // Two 1-byte Reads, then the rest, so the buffer is drained from + // an offset rather than in one copy. + n := 0 + for _, size := range []int{1, 1, len(p)} { + m, err := c.Read(p[n : n+min(size, len(p)-n)]) + if err != nil { + t.Fatalf("Read: %v", err) + } + n += m + } + if string(p[:n]) != "abcdefgh" { + t.Fatalf("Reads returned %q, want %q", p[:n], "abcdefgh") + } + } + drain() // the first decode allocates the buffer + if allocs := testing.AllocsPerRun(100, drain); allocs != 0 { + t.Fatalf("draining a decode allocated %v times per console read, want 0", allocs) + } +} + +// TestWriteSteadyStateAllocs needs no console: once the UTF-16 buffer has +// grown, a Write of the same size allocates nothing, so remote output does +// not cost a heap allocation per console write. The count the console +// reports back is the one thing that could escape: a local whose address +// goes to the write func value is moved to the heap. +func TestWriteSteadyStateAllocs(t *testing.T) { + c := newConsole(0, 0) + c.writeFn = func(_ windows.Handle, _ *uint16, n uint32, written *uint32, _ *byte) error { + *written = n + return nil + } + p := []byte("remote output line\r\n") + write := func() { + if n, err := c.Write(p); n != len(p) || err != nil { + t.Fatalf("Write = %d, %v; want %d, nil", n, err, len(p)) + } + } + write() // the first Write grows the UTF-16 buffer + if allocs := testing.AllocsPerRun(100, write); allocs != 0 { + t.Fatalf("a steady-state Write allocated %v times, want 0", allocs) + } +} + // 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) { @@ -1045,7 +1104,84 @@ func TestReadRetriesEmptyRead(t *testing.T) { } } +// TestSizeReportsCells needs no console: Size reports the visible window +// rectangle, not the screen buffer (dwSize), which the fake makes larger +// in both directions, as a console with scrollback reports it. func TestSizeReportsCells(t *testing.T) { + c := newConsole(0, 0) + c.sizeFn = func(_ windows.Handle, info *windows.ConsoleScreenBufferInfo) error { + info.Size = windows.Coord{X: 120, Y: 9001} + info.Window = windows.SmallRect{Left: 10, Top: 8977, Right: 89, Bottom: 9000} + return nil + } + sz, err := c.Size() + if err != nil { + t.Fatalf("Size: %v", err) + } + if sz != (Size{Rows: 24, Cols: 80}) { + t.Fatalf("Size() = %+v, want the 24x80 window, not the 9001x120 buffer", sz) + } +} + +// TestResizesWindows needs no console: a poll that finds the size unchanged +// yields nothing, a changed size is yielded once, and the range ends at the +// first poll after Close. The fake clock advances the 200 ms poll. +func TestResizesWindows(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + c := newConsole(0, 0) + var rows atomic.Int32 + rows.Store(24) + c.sizeFn = func(_ windows.Handle, info *windows.ConsoleScreenBufferInfo) error { + info.Window = windows.SmallRect{Right: 79, Bottom: int16(rows.Load() - 1)} + return nil + } + resizes := c.Resizes(t.Context()) // baseline 24x80 is taken here + sizes := make(chan Size, 4) + done := make(chan struct{}) + go func() { + defer close(done) + for sz := range resizes { + sizes <- sz + } + }() + poll := func(n int) { synctest.Sleep(time.Duration(n) * resizePoll) } + + poll(3) + if len(sizes) != 0 { + t.Fatalf("an unchanged size was yielded: %+v", <-sizes) + } + + rows.Store(50) + poll(1) + if len(sizes) != 1 { + t.Fatalf("%d sizes yielded after one poll of a changed size, want 1", len(sizes)) + } + if sz := <-sizes; sz != (Size{Rows: 50, Cols: 80}) { + t.Fatalf("Resizes yielded %+v, want 50x80", sz) + } + poll(3) + if len(sizes) != 0 { + t.Fatalf("the new size was yielded again: %+v", <-sizes) + } + + _ = c.Close() // zero handles: the flush and mode calls fail + select { + case <-done: + t.Fatal("the range over Resizes ended before the poll after Close") + default: + } + poll(1) + select { + case <-done: + default: + t.Fatal("the range over Resizes did not end at the poll after Close") + } + }) +} + +// TestSizeOnRealConsole reads the size of the console the test binary +// runs in, if it has one. +func TestSizeOnRealConsole(t *testing.T) { c := openConsole(t) sz, err := c.Size() if err != nil { diff --git a/internal/console/pty_linux_test.go b/internal/console/pty_linux_test.go index b658187..1ee2cb9 100644 --- a/internal/console/pty_linux_test.go +++ b/internal/console/pty_linux_test.go @@ -12,6 +12,11 @@ import ( // 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. +// +// A missing /dev/ptmx fails the test rather than skipping it: the Linux CI +// runner (ubuntu-latest in .github/workflows/ci.yml) has one, so a skip +// would only hide a broken environment and silently drop every test that +// needs a pty. func openPTY(t *testing.T) (master, slave *os.File) { t.Helper() master, err := os.OpenFile("/dev/ptmx", os.O_RDWR|unix.O_NOCTTY, 0) @@ -56,11 +61,8 @@ func closeAtEnd(t *testing.T, f *os.File) { // 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 + tio, err := controlValue(f, func(fd int) (*unix.Termios, error) { + return unix.IoctlGetTermios(fd, unix.TCGETS) }) if err != nil { t.Fatalf("tcgets: %v", err) diff --git a/internal/console/utf16.go b/internal/console/utf16.go index 25c2a21..8257e4c 100644 --- a/internal/console/utf16.go +++ b/internal/console/utf16.go @@ -1,6 +1,7 @@ package console import ( + "fmt" "unicode/utf16" "unicode/utf8" ) @@ -82,8 +83,13 @@ func (e *utf8Encoder) append(dst []uint16, p []byte) []uint16 { // 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. +// each write carries whole characters. limit must be at least 2, or a pair +// at the boundary would leave an empty chunk; chunkLen panics otherwise, +// whatever the input. func chunkLen(units []uint16, limit int) int { + if limit < 2 { + panic(fmt.Sprintf("console: chunkLen limit %d, want at least 2", limit)) + } if len(units) <= limit { return len(units) } diff --git a/internal/console/utf16_test.go b/internal/console/utf16_test.go index 0e0a480..cbb07b9 100644 --- a/internal/console/utf16_test.go +++ b/internal/console/utf16_test.go @@ -3,6 +3,7 @@ package console import ( "bytes" "slices" + "strings" "testing" "unicode/utf16" "unicode/utf8" @@ -125,6 +126,26 @@ func TestChunkLen(t *testing.T) { } } +// TestChunkLenRejectsSmallLimit checks the precondition: with a limit below +// 2 a pair at the boundary would leave a chunk of zero units, so chunkLen +// panics even for input short enough to take the fast path. +func TestChunkLenRejectsSmallLimit(t *testing.T) { + for _, limit := range []int{1, 0, -1} { + for _, units := range [][]uint16{nil, {'a'}, {'a', 'b', 'c'}} { + func() { + defer func() { + r := recover() + msg, ok := r.(string) + if !ok || !strings.Contains(msg, "limit") { + t.Errorf("chunkLen(%U, %d) panic = %v, want a string naming the limit", units, limit, r) + } + }() + chunkLen(units, limit) + }() + } + } +} + // FuzzChunkLen checks that chunking any input covers it exactly, never // exceeds the limit, and never separates a surrogate pair. func FuzzChunkLen(f *testing.F) { @@ -155,11 +176,14 @@ func FuzzChunkLen(f *testing.F) { } // FuzzUTF8EncoderSplit checks that splitting the input anywhere never changes -// the output, and that valid UTF-8 matches the standard library's encoding. +// the output, and that any input, valid or not, matches the standard +// library's encoding of the Go []rune conversion once the carried tail is +// flushed. 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.Add([]byte("\xf0\x9f\x98a\xed\xa0\x80"), uint8(2)) // truncated 4-byte, UTF-8 encoded surrogate f.Fuzz(func(t *testing.T, in []byte, at uint8) { var whole utf8Encoder want := whole.append(nil, in) @@ -175,17 +199,27 @@ func FuzzUTF8EncoderSplit(f *testing.F) { 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) + // The carry only ever holds one incomplete sequence, and the Go + // []rune conversion turns each byte of that into its own U+FFFD, so + // flushing it appends one U+FFFD per carried byte. + flushed := got + for range parts.n { + flushed = append(flushed, utf8.RuneError) + } + if ref := utf16.Encode([]rune(string(in))); !slices.Equal(flushed, ref) { + t.Fatalf("split %d of %q, flushed: got %U, want utf16.Encode %U", split, in, flushed, ref) } }) } // FuzzUTF16DecoderSplit is the decoder counterpart: any split gives the same -// output, and valid UTF-16 round-trips through the standard library. +// output, and any input, valid or not, matches the standard library's +// decoding (one U+FFFD per unpaired surrogate) once a carried high +// surrogate is flushed. func FuzzUTF16DecoderSplit(f *testing.F) { f.Add([]byte("h\x00i\x00=\xd8\x00\xde"), uint8(2)) f.Add([]byte("\x00\xdc"), uint8(0)) + f.Add([]byte("=\xd8a\x00=\xd8=\xd8\x00\xde=\xd8"), uint8(1)) // high then ASCII, high then pair, trailing high f.Fuzz(func(t *testing.T, raw []byte, at uint8) { units := make([]uint16, len(raw)/2) for i := range units { @@ -202,9 +236,14 @@ func FuzzUTF16DecoderSplit(f *testing.F) { 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)) + // A high surrogate still carried at the end is unpaired: flushing + // it appends one U+FFFD. + flushed := got + if parts.high != 0 { + flushed = utf8.AppendRune(flushed, utf8.RuneError) + } + if ref := string(utf16.Decode(units)); string(flushed) != ref { + t.Fatalf("split %d of %U, flushed: got %q, want utf16.Decode %q", split, units, flushed, ref) } }) } diff --git a/internal/etcp/backpressure_internal_test.go b/internal/etcp/backpressure_internal_test.go index f0774d0..d00c62a 100644 --- a/internal/etcp/backpressure_internal_test.go +++ b/internal/etcp/backpressure_internal_test.go @@ -17,7 +17,7 @@ func blockedConn(t *testing.T) *Conn { var d Dialer c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) c.mu.Lock() - c.unsent = c.limit + 1 + c.unsent = c.ring.limit + 1 c.mu.Unlock() return c } diff --git a/internal/etcp/conn.go b/internal/etcp/conn.go index 3c6c8ad..3918724 100644 --- a/internal/etcp/conn.go +++ b/internal/etcp/conn.go @@ -30,7 +30,6 @@ type Conn struct { id string keepAlive time.Duration probe protocol.Packet - limit int logger *slog.Logger ctx context.Context // lifetime of the Conn; its cause is what callers see @@ -61,7 +60,7 @@ type Conn struct { inbox chan protocol.Packet wake chan struct{} // cap 1: new outbound data for the link writer - space chan struct{} // cap 1: the unsent backlog fell to limit or below + space chan struct{} // cap 1: the unsent backlog fell to ReplayLimit (ring.limit) or below } // WritePacket seals p and queues it for delivery. It returns once p is queued, @@ -90,9 +89,9 @@ func (c *Conn) WritePacket(ctx context.Context, p protocol.Packet) error { c.mu.Unlock() return context.Cause(c.ctx) } - if c.unsent <= c.limit { + if c.unsent <= c.ring.limit { c.enqueueLocked(p) - room := c.unsent <= c.limit + room := c.unsent <= c.ring.limit c.mu.Unlock() signal(c.wake) if room { @@ -191,6 +190,14 @@ func (c *Conn) enqueueLocked(p protocol.Packet) { c.unsent += len(data) } +// releaseLocked trims the ring after flushed or unsent changed and reports +// whether the backlog is at or below the limit, so the caller signals space +// once it drops c.mu. The caller holds c.mu. +func (c *Conn) releaseLocked() bool { + c.ring.trim(c.flushed, c.unsent) + return c.unsent <= c.ring.limit +} + // fail ends the Conn with err and returns it. func (c *Conn) fail(err error) error { c.cancel(err) diff --git a/internal/etcp/dialer.go b/internal/etcp/dialer.go index a5eb0f6..734956a 100644 --- a/internal/etcp/dialer.go +++ b/internal/etcp/dialer.go @@ -52,10 +52,13 @@ var ( // that is oversized or does not decode: the stream is out of step or // tampered with and cannot be resumed. ErrIntegrity = errors.New("etcp: stream integrity failure") - // ErrReplayExceeded reports that the peer needs packets no longer - // retained, is ahead of what was sent, or needs a catchup too large to - // send in one message. - ErrReplayExceeded = errors.New("etcp: peer needs data beyond the replay window") + // ErrReplayExceeded reports that the session cannot be resumed: the + // peer needs packets no longer retained, is ahead of what was sent, or + // needs a catchup too large to send in one message, or more packets + // were received than a SequenceHeader's int32 sequence number can + // state. A peer position past that range wraps outside the retained + // window, so the send direction ends the same way. + ErrReplayExceeded = errors.New("etcp: session cannot be resumed") // ErrRejected reports that the server refused the session: INVALID_KEY // on the first connect, NEW_CLIENT on a redial, or an unknown status. ErrRejected = errors.New("etcp: server rejected the session") @@ -177,8 +180,7 @@ func (d *Dialer) newConn(addr, id, passkey string) *Conn { space: make(chan struct{}, 1), } c.lastProbe = -1 - c.limit = cmp.Or(d.ReplayLimit, defaultReplayLimit) - c.ring.limit = c.limit + c.ring.limit = cmp.Or(d.ReplayLimit, defaultReplayLimit) var key [32]byte copy(key[:], passkey) c.out = seal.New(&key, seal.ClientToServer) diff --git a/internal/etcp/helpers_test.go b/internal/etcp/helpers_test.go index ffaa1fe..ec084e8 100644 --- a/internal/etcp/helpers_test.go +++ b/internal/etcp/helpers_test.go @@ -25,6 +25,11 @@ const ( testAddr = "et.example:2022" ) +// contextDialer is the dialer a wrapping test NetDialer delegates to. +type contextDialer interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) +} + // harness is one session against the fake server over the in-memory network. type harness struct { srv *etservertest.Server diff --git a/internal/etcp/link.go b/internal/etcp/link.go index 39ef6cb..863fc9e 100644 --- a/internal/etcp/link.go +++ b/internal/etcp/link.go @@ -238,8 +238,7 @@ func (c *Conn) writeLoop(ctx context.Context, l *link, w io.Writer) error { c.mu.Lock() c.flushed += int64(sent) c.unsent -= n - c.ring.trim(c.flushed, c.unsent) - room := c.unsent <= c.limit + room := c.releaseLocked() c.mu.Unlock() if room { signal(c.space) diff --git a/internal/etcp/outage_test.go b/internal/etcp/outage_test.go index 9cb0bc6..babda3b 100644 --- a/internal/etcp/outage_test.go +++ b/internal/etcp/outage_test.go @@ -14,9 +14,7 @@ import ( // dialClock records when each dial happens. type dialClock struct { - inner interface { - DialContext(ctx context.Context, network, address string) (net.Conn, error) - } + inner contextDialer mu sync.Mutex times []time.Time } diff --git a/internal/etcp/recover.go b/internal/etcp/recover.go index bae0478..9a59936 100644 --- a/internal/etcp/recover.go +++ b/internal/etcp/recover.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "math" "net" "sync" "time" @@ -117,8 +118,7 @@ func (c *Conn) recover(conn net.Conn) ([][]byte, error) { // Trim as writeLoop does after a Write: a link that dies before its // first Write would otherwise leave the catchup in the ring while // WritePacket admits another ReplayLimit. - c.ring.trim(c.flushed, c.unsent) - room := c.unsent <= c.limit + room := c.releaseLocked() c.mu.Unlock() if room { signal(c.space) @@ -150,8 +150,18 @@ var maxCatchupSize = wire.MaxMessageSize // writeRecover writes our half of the recover exchange and returns the // sequence number the new link starts sending from. It fails with // ErrReplayExceeded when the peer's position is outside the retained window -// or our catchup is too large for one message. +// or our catchup is too large for one message, and also when we have +// received more packets than SequenceHeader's int32 sequence_number can +// state (internal/protocol ET.pb.go), rather than send a wrapped negative +// count. The peer's position needs no such guard: a position truncated to +// int32 lies at least 2^31 packets below its true value, while the retained +// window holds at most a few ReplayLimits of bytes, far fewer packets, so +// ring.since refuses it. func (c *Conn) writeRecover(conn net.Conn, gotSeq <-chan error, peer *protocol.SequenceHeader) (int64, error) { + if c.recvSeq > math.MaxInt32 { + return 0, fmt.Errorf("%w: received %d packets, more than the protocol's int32 sequence number can express", + ErrReplayExceeded, c.recvSeq) + } mine := &protocol.SequenceHeader{} mine.SetSequenceNumber(int32(c.recvSeq)) if err := wire.WriteMessage(conn, mine); err != nil { diff --git a/internal/etcp/recover_internal_test.go b/internal/etcp/recover_internal_test.go new file mode 100644 index 0000000..c534145 --- /dev/null +++ b/internal/etcp/recover_internal_test.go @@ -0,0 +1,86 @@ +package etcp + +import ( + "errors" + "math" + "net" + "strings" + "testing" + "testing/synctest" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" +) + +// A receive count past the int32 range cannot be stated in a SequenceHeader: +// writeRecover must return ErrReplayExceeded, which isFatal treats as the end +// of the Conn, rather than send a wrapped, negative sequence number. +func TestWriteRecoverRefusesSequenceBeyondInt32(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + defer c.cancel(nil) + c.recvSeq = math.MaxInt32 + 1 + + client, server := net.Pipe() + got := make(chan error, 1) + go func() { + var sh protocol.SequenceHeader + err := wire.ReadMessage(server, &sh) + if err == nil { + err = errors.New("server received a SequenceHeader") + } else { + err = nil + } + _ = server.Close() + got <- err + }() + + // A reader failure is preloaded so a writeRecover that did send + // returns instead of waiting for the peer's sequence. + gotSeq := make(chan error, 1) + gotSeq <- errors.New("no peer sequence") + _, err := c.writeRecover(client, gotSeq, &protocol.SequenceHeader{}) + _ = client.Close() + if !errors.Is(err, ErrReplayExceeded) { + t.Fatalf("writeRecover = %v, want ErrReplayExceeded", err) + } + if serr := <-got; serr != nil { + t.Fatal(serr) + } + }) +} + +// A receive count of exactly MaxInt32 still fits a SequenceHeader, so +// writeRecover sends it unchanged: the guard refuses only past the range. +func TestWriteRecoverSendsSequenceAtInt32Limit(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + defer c.cancel(nil) + c.recvSeq = math.MaxInt32 + + client, server := net.Pipe() + sent := make(chan int32, 1) + go func() { + var sh protocol.SequenceHeader + if err := wire.ReadMessage(server, &sh); err != nil { + sent <- -1 + } else { + sent <- sh.GetSequenceNumber() + } + _ = server.Close() + }() + + gotSeq := make(chan error, 1) + gotSeq <- errors.New("no peer sequence") + _, err := c.writeRecover(client, gotSeq, &protocol.SequenceHeader{}) + _ = client.Close() + if errors.Is(err, ErrReplayExceeded) { + t.Fatalf("writeRecover = %v, want the header sent at the int32 limit", err) + } + if seq := <-sent; seq != math.MaxInt32 { + t.Fatalf("server read sequence %d, want %d", seq, int32(math.MaxInt32)) + } + }) +} diff --git a/internal/etcp/ring.go b/internal/etcp/ring.go index 251756b..a9e7baa 100644 --- a/internal/etcp/ring.go +++ b/internal/etcp/ring.go @@ -1,7 +1,5 @@ package etcp -import "slices" - // ring holds sealed, serialized outbound packets by sequence number, for the // link writer and for replay after a reconnect. Sequence numbers are // contiguous: entries[i] has sequence number first+i. @@ -52,12 +50,14 @@ func (r *ring) appendRange(dst [][]byte, from, to int64) [][]byte { } // since returns a copy of entries [from, to), or false when from is outside -// the retained window [first, to] or to is past next(). +// the retained window [first, to] or to is past next(). The copy (appended +// to a nil slice) lets the caller use it after releasing the lock that +// guards the ring. func (r *ring) since(from, to int64) ([][]byte, bool) { if from < r.first || from > to || to > r.next() { return nil, false } - return slices.Clone(r.entries[from-r.first : to-r.first]), true + return r.appendRange(nil, from, to), true } // bytesBetween sums the sizes of entries [from, to). The caller guarantees diff --git a/internal/etcp/ring_test.go b/internal/etcp/ring_test.go index 6b1b90b..07924ff 100644 --- a/internal/etcp/ring_test.go +++ b/internal/etcp/ring_test.go @@ -80,6 +80,26 @@ func TestRingSince(t *testing.T) { } } +// since returns a copy: trim, which clears the slots it drops, leaves the +// entries since returned intact, as writeRecover relies on after releasing +// the lock. +func TestRingSinceIsACopy(t *testing.T) { + r := filledRing(20, 5, 10) + got, ok := r.since(r.first, r.next()) + if !ok || len(got) != 5 { + t.Fatalf("since(first, next) = %d entries, %v; want 5, true", len(got), ok) + } + r.trim(r.next(), 0) + if r.first == 0 { + t.Fatal("trim dropped nothing; the test needs it to drop entries") + } + for i, e := range got { + if want := bytes.Repeat([]byte{byte(i)}, 10); !bytes.Equal(e, want) { + t.Fatalf("entry %d after trim = %v, want %v", i, e, want) + } + } +} + func TestRingBytesBetweenAndRange(t *testing.T) { r := filledRing(1<<20, 6, 10) if got := r.bytesBetween(2, 5); got != 30 { diff --git a/internal/etcp/throttle_test.go b/internal/etcp/throttle_test.go index e0d74d1..2a1dd9a 100644 --- a/internal/etcp/throttle_test.go +++ b/internal/etcp/throttle_test.go @@ -15,9 +15,7 @@ import ( // time (zero means 125 ms, 8 KiB/s), like a congested uplink, or with // downlink set every client-side read instead, like a congested downlink. type throttledDialer struct { - inner interface { - DialContext(ctx context.Context, network, address string) (net.Conn, error) - } + inner contextDialer downlink bool perKiB time.Duration } diff --git a/internal/etservertest/server.go b/internal/etservertest/server.go index 234c6a6..9c04b9b 100644 --- a/internal/etservertest/server.go +++ b/internal/etservertest/server.go @@ -77,7 +77,11 @@ func NewServer(id, passkey string) *Server { // connects twice waits for the first link to settle. func (s *Server) Serve(ctx context.Context, c net.Conn) error { defer func() { _ = c.Close() }() - stop := context.AfterFunc(ctx, func() { _ = c.Close() }) + // lctx ends with ctx or when either stream loop fails; closing c then + // unblocks the handshake, the recover exchange or the other loop. + lctx, cancel := context.WithCancelCause(ctx) + defer cancel(nil) + stop := context.AfterFunc(lctx, func() { _ = c.Close() }) defer stop() var req protocol.ConnectRequest @@ -109,11 +113,6 @@ func (s *Server) Serve(ctx context.Context, c net.Conn) error { flushed = f } - lctx, cancel := context.WithCancelCause(ctx) - defer cancel(nil) - stopClose := context.AfterFunc(lctx, func() { _ = c.Close() }) - defer stopClose() - var wg sync.WaitGroup wg.Go(func() { cancel(s.readLoop(c)) }) wg.Go(func() { cancel(s.writeLoop(lctx, c, flushed)) }) diff --git a/internal/protocol/redact_test.go b/internal/protocol/redact_test.go index dd3a62e..7d78747 100644 --- a/internal/protocol/redact_test.go +++ b/internal/protocol/redact_test.go @@ -2,6 +2,7 @@ package protocol import ( "bytes" + "encoding/json" "fmt" "log/slog" "strings" @@ -11,7 +12,8 @@ import ( // TestTerminalUserInfoRedacted checks that no fmt verb and no slog handler // renders a TerminalUserInfo's passkey, that each rendering says REDACTED, // and that the id still appears, for two different ids so a hardcoded id -// cannot pass. +// cannot pass. encoding/json, which bypasses Format, is checked only for +// leaving the passkey out. func TestTerminalUserInfoRedacted(t *testing.T) { const passkey = "SECRETPASSKEY0123456789abcdefXYZ" for _, id := range []string{"XXXabcdefghijklm", "YYYnopqrstuvwxyz"} { @@ -46,6 +48,22 @@ func TestTerminalUserInfoRedacted(t *testing.T) { }) } + // encoding/json does not go through Format, and opaque messages hide + // their fields from it (MEASURED 2026-09-26: the output is only + // XXX_ bookkeeping fields). Pin that the passkey stays out. + t.Run("encoding/json", func(t *testing.T) { + u := &TerminalUserInfo{} + u.SetId("XXXabcdefghijklm") + u.SetPasskey(passkey) + b, err := json.Marshal(u) + if err != nil { + t.Fatalf("json.Marshal: %v", err) + } + if strings.Contains(string(b), passkey) { + t.Errorf("json.Marshal output %s contains the passkey", b) + } + }) + t.Run("nil", func(t *testing.T) { var u *TerminalUserInfo // Through fmt, not u.String(): the generated String bypasses Format. diff --git a/internal/seal/seal.go b/internal/seal/seal.go index 916f131..64f0de6 100644 --- a/internal/seal/seal.go +++ b/internal/seal/seal.go @@ -54,10 +54,13 @@ type state struct { // New returns a Stream for direction d with the nonce at its initial value. // It copies *key, so later changes to the caller's array do not affect it. -// key must not be nil, and d must be ClientToServer or ServerToClient: any -// other direction produces a nonce stream the peer never uses, so every Open -// fails as if the key were wrong. +// Both misuses are programming errors and panic: a nil key, and a d other +// than ClientToServer or ServerToClient, which would produce a nonce stream +// the peer never uses, so every Open would fail as if the key were wrong. func New(key *[32]byte, d Direction) *Stream { + if d != ClientToServer && d != ServerToClient { + panic(fmt.Sprintf("seal: invalid direction %d", d)) + } k := *key st := &state{key: &k} st.nonce[len(st.nonce)-1] = byte(d) diff --git a/internal/seal/seal_test.go b/internal/seal/seal_test.go index a0c46c9..b9ebab3 100644 --- a/internal/seal/seal_test.go +++ b/internal/seal/seal_test.go @@ -86,6 +86,17 @@ func readGolden(t *testing.T) []goldenRow { return rows } +// goldenFor returns the rows for direction dir, keyed by operation index. +func goldenFor(rows []goldenRow, dir Direction) map[int]goldenRow { + want := map[int]goldenRow{} + for _, r := range rows { + if r.dir == dir { + want[r.index] = r + } + } + return want +} + // TestGolden seals the same sequence as testdata/gen/golden.c and compares // nonce and box at the recorded indices, including the carry from byte 0 to // byte 1 at operation 256. @@ -93,12 +104,7 @@ func TestGolden(t *testing.T) { rows := readGolden(t) for _, dir := range []Direction{ClientToServer, ServerToClient} { t.Run(fmt.Sprintf("dir%d", dir), func(t *testing.T) { - want := map[int]goldenRow{} - for _, r := range rows { - if r.dir == dir { - want[r.index] = r - } - } + want := goldenFor(rows, dir) s := New(&testKey, dir) checked := 0 for i := 1; i <= 257; i++ { @@ -134,12 +140,7 @@ func TestGoldenOpen(t *testing.T) { rows := readGolden(t) for _, dir := range []Direction{ClientToServer, ServerToClient} { t.Run(fmt.Sprintf("dir%d", dir), func(t *testing.T) { - want := map[int]goldenRow{} - for _, r := range rows { - if r.dir == dir { - want[r.index] = r - } - } + want := goldenFor(rows, dir) opener := New(&testKey, dir) sealer := New(&testKey, dir) checked := 0 @@ -169,6 +170,19 @@ func TestGoldenOpen(t *testing.T) { } } +// TestNewInvalidDirection pins that a direction other than the two upstream +// defines is refused at construction, not left to fail every Open later. +func TestNewInvalidDirection(t *testing.T) { + defer func() { + r := recover() + msg, _ := r.(string) + if !strings.Contains(msg, "invalid direction") { + t.Fatalf("New(Direction(2)) recovered %v, want a panic containing %q", r, "invalid direction") + } + }() + New(&testKey, Direction(2)) +} + func TestRoundTrip(t *testing.T) { sealer := New(&testKey, ServerToClient) opener := New(&testKey, ServerToClient) diff --git a/internal/wire/frame.go b/internal/wire/frame.go index e26fc76..511a87f 100644 --- a/internal/wire/frame.go +++ b/internal/wire/frame.go @@ -35,7 +35,7 @@ func WriteFrame(w io.Writer, frame []byte) error { if err != nil { return err } - if err := writeAll(w, buf); err != nil { + if err := writeChecked(w, buf); err != nil { return fmt.Errorf("wire: write frame: %w", err) } return nil @@ -73,9 +73,9 @@ func ReadFrame(r io.Reader, buf []byte) ([]byte, error) { return buf, nil } -// writeAll issues one Write of buf and turns a short count without an error -// into io.ErrShortWrite, as io.Copy does. -func writeAll(w io.Writer, buf []byte) error { +// writeChecked issues exactly one Write of buf and turns a short count without +// an error into io.ErrShortWrite, as io.Copy does. +func writeChecked(w io.Writer, buf []byte) error { n, err := w.Write(buf) if err == nil && n != len(buf) { err = io.ErrShortWrite diff --git a/internal/wire/frame_test.go b/internal/wire/frame_test.go index 22c1e34..47b62c0 100644 --- a/internal/wire/frame_test.go +++ b/internal/wire/frame_test.go @@ -4,6 +4,7 @@ import ( "bytes" "encoding/binary" "errors" + "fmt" "io" "testing" ) @@ -287,6 +288,62 @@ func TestReadFrameGrowsWithData(t *testing.T) { } } +// A body that outgrows the caller's buffer ends in a buffer of exactly the +// declared length: the last growth step is capped at the length, not rounded +// up by append's growth policy. +func TestReadFrameGrowsExactly(t *testing.T) { + const n = 300_000 + var enc bytes.Buffer + if err := WriteFrame(&enc, bytes.Repeat([]byte{'x'}, n)); err != nil { + t.Fatalf("WriteFrame: %v", err) + } + got, err := ReadFrame(&enc, nil) + if err != nil { + t.Fatalf("ReadFrame: %v", err) + } + if len(got) != n || cap(got) != n { + t.Fatalf("ReadFrame body len %d cap %d, want both %d", len(got), cap(got), n) + } + if !bytes.Equal(got, bytes.Repeat([]byte{'x'}, n)) { + t.Fatal("ReadFrame body differs from the frame written") + } +} + +// A reader that reuses the previous frame as its buffer, as etcp's link +// readLoop does, grows it geometrically: frames that each exceed every +// earlier one reallocate a handful of times, not once per frame. +func TestReadFrameReusedBufferGrowsGeometrically(t *testing.T) { + const frames = 4000 + var enc bytes.Buffer + for i := 1; i <= frames; i++ { + if err := WriteFrame(&enc, bytes.Repeat([]byte{'x'}, i)); err != nil { + t.Fatalf("WriteFrame: %v", err) + } + } + stream := enc.Bytes() + var failed error + allocs := testing.AllocsPerRun(3, func() { + r := bytes.NewReader(stream) + var buf []byte + for i := 1; i <= frames; i++ { + frame, err := ReadFrame(r, buf) + if err != nil || len(frame) != i { + failed = fmt.Errorf("frame %d: len %d: %w", i, len(frame), err) + return + } + buf = frame + } + }) + if failed != nil { + t.Fatal(failed) + } + // The bytes.Reader costs one allocation per run; growth from 1 to 4000 + // bytes by doubling takes about a dozen more. + if allocs > 16 { + t.Fatalf("reading %d growing frames through one reused buffer allocated %v times, want at most 16", frames, allocs) + } +} + // ReadFrame into a buffer that is large enough, and AppendFrame into one, // allocate nothing: etcp's read and write loops run them for every packet. func TestFrameSteadyStateAllocs(t *testing.T) { diff --git a/internal/wire/message.go b/internal/wire/message.go index b0f9cb9..04a27d4 100644 --- a/internal/wire/message.go +++ b/internal/wire/message.go @@ -10,19 +10,25 @@ import ( // WriteMessage writes m as an 8-byte little-endian length followed by its // protobuf encoding, in a single Write call. A writer that reports fewer -// bytes than it was given yields io.ErrShortWrite. +// bytes than it was given yields io.ErrShortWrite. The message is marshaled +// straight into the output buffer, after room for the length, so the size +// limit is checked before anything is allocated. func WriteMessage(w io.Writer, m proto.Message) error { - body, err := proto.Marshal(m) + size := proto.Size(m) + if size > MaxMessageSize { + return fmt.Errorf("wire: write %T of %d bytes: %w", m, size, ErrTooLarge) + } + buf, err := proto.MarshalOptions{}.MarshalAppend(make([]byte, 8, 8+size), m) if err != nil { return fmt.Errorf("wire: marshal %T: %w", m, err) } - if len(body) > MaxMessageSize { - return fmt.Errorf("wire: write %T of %d bytes: %w", m, len(body), ErrTooLarge) + // proto.Size is exact for a message nothing else is changing; the + // prefix and the limit still use the length actually marshaled. + if len(buf)-8 > MaxMessageSize { + return fmt.Errorf("wire: write %T of %d bytes: %w", m, len(buf)-8, ErrTooLarge) } - buf := make([]byte, 8, 8+len(body)) - binary.LittleEndian.PutUint64(buf, uint64(len(body))) - buf = append(buf, body...) - if err := writeAll(w, buf); err != nil { + binary.LittleEndian.PutUint64(buf, uint64(len(buf)-8)) + if err := writeChecked(w, buf); err != nil { return fmt.Errorf("wire: write %T: %w", m, err) } return nil diff --git a/internal/wire/message_limit_test.go b/internal/wire/message_limit_test.go index ce6310c..32c4ac9 100644 --- a/internal/wire/message_limit_test.go +++ b/internal/wire/message_limit_test.go @@ -4,6 +4,7 @@ package wire import ( "errors" + "runtime" "testing" "github.com/tphakala/et-go/internal/protocol" @@ -12,14 +13,14 @@ import ( // TestWriteMessageLimit pins WriteMessage's bound from both sides: a body of // exactly MaxMessageSize is written, one byte more is refused with -// ErrTooLarge and nothing reaches the writer. Each case holds a few hundred -// MiB (the fixture, the marshaled body and the framed buffer), so it is -// skipped in -short mode and excluded from -race builds, where the race -// detector's shadow memory pushes it past 1 GiB. CI runs it in the Windows -// leg and in a dedicated non-race step on Ubuntu. +// ErrTooLarge and nothing reaches the writer. The at-limit case holds about +// 256 MiB (the fixture and the framed buffer) and the one-over case the +// fixture alone, so it is skipped in -short mode and excluded from -race +// builds, where the race detector's shadow memory pushes it past 1 GiB. CI +// runs it in the Windows leg and in a dedicated non-race step on Ubuntu. func TestWriteMessageLimit(t *testing.T) { if testing.Short() { - t.Skip("allocates a few hundred MiB per case") + t.Skip("allocates up to about 256 MiB per case") } // A CatchupBuffer with one entry encodes as tag 0x0a, a 4-byte varint // length, then the entry: 5 bytes of overhead at this size. @@ -40,10 +41,18 @@ func TestWriteMessageLimit(t *testing.T) { t.Fatalf("fixture encodes to %d bytes, want %d", got, tt.size) } w := &lenWriter{} + var before, after runtime.MemStats + runtime.ReadMemStats(&before) err := WriteMessage(w, cb) + runtime.ReadMemStats(&after) if !errors.Is(err, tt.wantErr) { t.Fatalf("WriteMessage(%d-byte body) = %v, want %v", tt.size, err, tt.wantErr) } + // A refused message is refused before the output buffer is + // allocated: the check runs on proto.Size, not after marshaling. + if d := after.TotalAlloc - before.TotalAlloc; tt.wantErr != nil && d > 1<<20 { + t.Errorf("refusing the message allocated %d bytes, want the refusal before any buffer", d) + } wantWritten := 0 if tt.wantErr == nil { wantWritten = 8 + tt.size diff --git a/internal/wire/message_test.go b/internal/wire/message_test.go index d3f27b0..ca6b230 100644 --- a/internal/wire/message_test.go +++ b/internal/wire/message_test.go @@ -170,6 +170,27 @@ func TestWriteMessageSingleWrite(t *testing.T) { } } +// TestWriteMessageAllocs pins that WriteMessage marshals straight into its +// output buffer: the only allocation is that buffer, not a separate +// marshaled body copied into it. +func TestWriteMessageAllocs(t *testing.T) { + req := &protocol.ConnectRequest{} + req.SetClientId(string(bytes.Repeat([]byte{'x'}, 1<<10))) + req.SetVersion(protocol.Version) + var werr error + allocs := testing.AllocsPerRun(100, func() { + if err := WriteMessage(io.Discard, req); err != nil { + werr = err + } + }) + if werr != nil { + t.Fatalf("WriteMessage: %v", werr) + } + if allocs > 1 { + t.Fatalf("WriteMessage allocated %v times per call, want at most 1", allocs) + } +} + func TestReadMessageGarbage(t *testing.T) { // Length 2, then bytes that are not a valid protobuf (field 0 is illegal). in := []byte{2, 0, 0, 0, 0, 0, 0, 0, 0x00, 0x00} diff --git a/internal/wire/wire.go b/internal/wire/wire.go index 2c6a9cf..99ffa3e 100644 --- a/internal/wire/wire.go +++ b/internal/wire/wire.go @@ -16,7 +16,6 @@ package wire import ( "errors" "io" - "slices" ) // Size limits for lengths read from the network. @@ -42,23 +41,20 @@ var ErrMalformed = errors.New("wire: malformed message") // ErrShortPacket reports a serialized packet without its 2-byte header. var ErrShortPacket = errors.New("wire: packet shorter than its 2-byte header") -// growChunk is the first allocation for a body that does not fit the -// caller's buffer; the buffer then doubles as bytes arrive. +// growChunk is the first read step for a body that does not fit the +// caller's buffer; the read steps then double as bytes arrive. const growChunk = 64 << 10 -// bodyErr maps an io.EOF (possibly wrapped) from a body read to -// io.ErrUnexpectedEOF. A body read happens only after a complete length -// prefix, so an end of stream there, whether before the first body byte or -// between two growth chunks, is a broken link, not a clean end of stream. +// bodyErr classifies a failed body read as headerErr does after one byte: a +// body read always follows a complete length prefix, so any end of stream +// there is io.ErrUnexpectedEOF. func bodyErr(err error) error { - if errors.Is(err, io.EOF) { - return io.ErrUnexpectedEOF - } - return err + return headerErr(1, err) } -// headerErr classifies a failed length-prefix read that got n bytes. Only a -// read that got no byte at all is a clean end, reported as a plain io.EOF. +// headerErr classifies a failed read that got n bytes, of a length prefix +// or (through bodyErr) of a body. Only a read that got no byte at all is a +// clean end, reported as a plain io.EOF. // io.ReadFull maps only an unwrapped io.EOF after a partial read to // io.ErrUnexpectedEOF, so a reader that returns bytes together with a // wrapped io.EOF is mapped here. @@ -73,10 +69,13 @@ func headerErr(n int, err error) error { } // readBody reads exactly n bytes into buf's storage and returns buf[:n], -// reusing buf's capacity when it is large enough. Otherwise the buffer grows -// as bytes arrive, starting at growChunk and doubling, so a peer that -// declares a large length and sends little costs memory in proportion to -// what it sent, not to what it declared. +// reusing buf's capacity when it is large enough. Otherwise it reads in +// steps, the first of min(n, growChunk) bytes and each later one doubling +// what has arrived, and grows the buffer only when a step does not fit: to +// twice its old capacity or the step, whichever is larger, but never past +// max(n, growChunk). A peer that declares a large length and sends little +// therefore costs memory in proportion to what it sent, not to what it +// declared. func readBody(r io.Reader, buf []byte, n int) ([]byte, error) { if cap(buf) >= n { buf = buf[:n] @@ -88,7 +87,19 @@ func readBody(r io.Reader, buf []byte, n int) ([]byte, error) { buf = buf[:0] for len(buf) < n { next := min(n, max(2*len(buf), growChunk)) - buf = slices.Grow(buf, next-len(buf)) + if cap(buf) < next { + // A reused buffer at least doubles, so a reader that passes + // its previous frame back grows it a few times rather than + // once per larger frame. The capacity is capped at n above + // growChunk, where the body ends exactly at n rather than at + // append's rounded size; a reused buffer growing past + // growChunk still reallocates for each larger body, which + // frames carrying one pty read (16 KiB, see MaxFrameSize, plus + // packet overhead) never reach. + nb := make([]byte, len(buf), min(max(next, 2*cap(buf)), max(n, growChunk))) + copy(nb, buf) + buf = nb + } m, err := io.ReadFull(r, buf[len(buf):next]) buf = buf[:len(buf)+m] if err != nil {