diff --git a/internal/nix/command.go b/internal/nix/command.go index 902f2a33aca..f46cbbeca84 100644 --- a/internal/nix/command.go +++ b/internal/nix/command.go @@ -8,6 +8,10 @@ func init() { "--option", "experimental-features", "nix-command flakes fetch-closure", } + // Retry commands that fail because of a flaky network, such as a + // truncated nixpkgs tarball download. + Default.MaxAttempts = 3 + // Add GitHub access token if available to avoid rate limiting // This is a backup in case the config file isn't picked up properly if token := os.Getenv("GITHUB_TOKEN"); token != "" { diff --git a/nix/command.go b/nix/command.go index ab7cd4793f2..a64827bab71 100644 --- a/nix/command.go +++ b/nix/command.go @@ -15,6 +15,8 @@ import ( "strings" "syscall" "time" + + "github.com/mattn/go-isatty" ) // Cmd is an external command that invokes a [*Nix] executable. It provides @@ -43,6 +45,15 @@ type Cmd struct { // defaults to [slog.Default]. Logger *slog.Logger + // MaxAttempts is the maximum number of times to run the command when + // it fails with a transient network error, such as a truncated + // download. Values less than 2 disable retries. See [Nix.MaxAttempts]. + // + // Stdout may receive output from failed attempts before a retry. In + // practice the retried errors happen while fetching, before Nix writes + // any output. + MaxAttempts int + execCmd *exec.Cmd err error dur time.Duration @@ -52,8 +63,9 @@ type Cmd struct { // Logger and other defaults from n. func (n *Nix) Command(args ...any) *Cmd { cmd := &Cmd{ - Args: make(Args, 1, 1+len(n.ExtraArgs)+len(args)), - Logger: n.logger(), + Args: make(Args, 1, 1+len(n.ExtraArgs)+len(args)), + Logger: n.logger(), + MaxAttempts: n.MaxAttempts, } cmd.Path, cmd.err = n.resolvePath() @@ -69,35 +81,157 @@ func (n *Nix) Command(args ...any) *Cmd { func (c *Cmd) CombinedOutput(ctx context.Context) ([]byte, error) { defer c.logRunFunc(ctx)() - - start := time.Now() - out, err := c.initExecCommand(ctx).CombinedOutput() - c.dur = time.Since(start) - - c.err = c.error(ctx, err) - return out, c.err + return c.run(ctx, true, (*exec.Cmd).CombinedOutput) } func (c *Cmd) Output(ctx context.Context) ([]byte, error) { defer c.logRunFunc(ctx)() - - start := time.Now() - out, err := c.initExecCommand(ctx).Output() - c.dur = time.Since(start) - - c.err = c.error(ctx, err) - return out, c.err + return c.run(ctx, false, (*exec.Cmd).Output) } func (c *Cmd) Run(ctx context.Context) error { defer c.logRunFunc(ctx)() + _, err := c.run(ctx, false, func(cmd *exec.Cmd) ([]byte, error) { + return nil, cmd.Run() + }) + return err +} - start := time.Now() - err := c.initExecCommand(ctx).Run() - c.dur = time.Since(start) +// run calls runFunc with a new [exec.Cmd] for each attempt, retrying up to +// c.MaxAttempts times when Nix fails with a transient network error. combined +// indicates that runFunc returns stderr interleaved with stdout. +func (c *Cmd) run(ctx context.Context, combined bool, runFunc func(*exec.Cmd) ([]byte, error)) ([]byte, error) { + for attempt := 1; ; attempt++ { + c.execCmd = nil + execCmd := c.initExecCommand(ctx) + + // When the caller provides its own stderr, keep a copy of the + // end of it so we can check for transient errors. Terminals + // are left alone so that Nix still renders its progress bar, + // which means those commands aren't retried. + var stderrTail *tailWriter + if c.canRetry() && c.Stderr != nil && !isTerminal(c.Stderr) { + stderrTail = &tailWriter{} + execCmd.Stderr = io.MultiWriter(c.Stderr, stderrTail) + } - c.err = c.error(ctx, err) - return c.err + start := time.Now() + out, err := runFunc(execCmd) + c.dur = time.Since(start) + c.err = c.error(ctx, err) + if c.err == nil || attempt >= c.MaxAttempts || !c.canRetry() || ctx.Err() != nil { + return out, c.err + } + + // Nix's stderr is in stderrTail for a caller-provided stderr, + // in the exit error for Output, or in out for CombinedOutput. + // Never check stdout alone, which might contain one of the + // error strings. + var stderr []byte + var exitErr *exec.ExitError + switch { + case stderrTail != nil: + stderr = stderrTail.buf + case errors.As(err, &exitErr) && len(exitErr.Stderr) != 0: + stderr = exitErr.Stderr + case combined: + stderr = out + } + if !isTransientError(stderr) { + return out, c.err + } + + delay := time.Duration(attempt) * retryDelay + c.logger().DebugContext(ctx, "retrying nix command after transient error", + "attempt", attempt, "delay", delay, "cmd", c) + w := c.Stderr + if w == nil { + w = os.Stderr + } + fmt.Fprintf(w, "Nix failed with a transient error, retrying in %s (attempt %d of %d): %s\n", + delay, attempt+1, c.MaxAttempts, c.stderrExcerpt(stderr)) + + timer := time.NewTimer(delay) + select { + case <-ctx.Done(): + timer.Stop() + return out, c.err + case <-timer.C: + } + } +} + +// retryDelay is how long to wait before the first retry. Each subsequent retry +// waits an additional retryDelay. +var retryDelay = 2 * time.Second + +// canRetry reports if c can safely be run more than once. Stdin that isn't a +// file might have been consumed by the previous attempt. File stdin (usually +// os.Stdin) isn't replayable either, but the transient errors that trigger a +// retry happen while fetching, before Nix reads any input. +func (c *Cmd) canRetry() bool { + if c.MaxAttempts < 2 { + return false + } + if c.Stdin == nil { + return true + } + _, isFile := c.Stdin.(*os.File) + return isFile +} + +// transientErrors are substrings of Nix error messages caused by flaky +// network connections or servers. Nix already retries failed downloads, but +// not ones that fail partway through unpacking a tarball, and it gives up on +// server errors after a few quick attempts. +// +// Errors that won't go away within a few seconds, such as GitHub rate limits +// (HTTP 403/429) or DNS failures when offline, are deliberately excluded. +var transientErrors = []string{ + "Truncated tar archive", + "Damaged tar archive", + "Failure when receiving data from the peer", + "Connection reset by peer", + "Timeout was reached", + "HTTP error 500", + "HTTP error 502", + "HTTP error 503", + "HTTP error 504", +} + +func isTransientError(stderr []byte) bool { + for _, msg := range transientErrors { + if bytes.Contains(stderr, []byte(msg)) { + return true + } + } + return false +} + +func isTerminal(w io.Writer) bool { + f, ok := w.(*os.File) + return ok && (isatty.IsTerminal(f.Fd()) || isatty.IsCygwinTerminal(f.Fd())) +} + +// tailWriter keeps the last few KiB written to it, which is enough to hold +// Nix's error message. +type tailWriter struct { + buf []byte +} + +func (t *tailWriter) Write(data []byte) (int, error) { + const maxLen = 8 << 10 + n := len(data) + if len(data) > maxLen { + data = data[len(data)-maxLen:] + } + // Shift out old bytes in place so the buffer never grows beyond + // maxLen. + if drop := len(t.buf) + len(data) - maxLen; drop > 0 { + t.buf = t.buf[:copy(t.buf, t.buf[drop:])] + } + t.buf = append(t.buf, data...) + return n, nil } func (c *Cmd) LogValue() slog.Value { diff --git a/nix/command_test.go b/nix/command_test.go new file mode 100644 index 00000000000..06ae2afccda --- /dev/null +++ b/nix/command_test.go @@ -0,0 +1,204 @@ +package nix + +import ( + "bytes" + "context" + "os" + "path/filepath" + "strconv" + "strings" + "testing" + "time" +) + +// fakeNix writes a script that pretends to be Nix. Each invocation prints +// the next entry in stderrs to stderr and exits with an error, until it runs +// out of entries and succeeds. It returns the script path and a function that +// reports how many times the script ran. +func fakeNix(t *testing.T, stderrs ...string) (path string, runs func() int) { + t.Helper() + retryDelay = 0 + t.Cleanup(func() { retryDelay = 2 * time.Second }) + + dir := t.TempDir() + countFile := filepath.Join(dir, "count") + script := "#!/bin/sh\n" + + "n=$(cat " + countFile + " 2>/dev/null || echo 0)\n" + + "echo $((n + 1)) > " + countFile + "\n" + + "case $n in\n" + for i, stderr := range stderrs { + script += strconv.Itoa(i) + ") echo '" + stderr + "' >&2; exit 1 ;;\n" + } + script += "esac\necho ok\n" + + path = filepath.Join(dir, "nix") + if err := os.WriteFile(path, []byte(script), 0o755); err != nil { + t.Fatal(err) + } + return path, func() int { + // Don't fail the test here, it can be called from another + // goroutine and before the script has run. + b, _ := os.ReadFile(countFile) + n, _ := strconv.Atoi(strings.TrimSpace(string(b))) + return n + } +} + +const truncatedTarErr = "error: cannot read file from tarball: Truncated tar archive detected while reading data" + +func TestCmdOutputRetriesTransientError(t *testing.T) { + path, runs := fakeNix(t, truncatedTarErr) + cmd := &Cmd{Path: path, Args: Args{path}, MaxAttempts: 3} + + out, err := cmd.Output(t.Context()) + if err != nil { + t.Fatalf("got error: %v", err) + } + if got, want := strings.TrimSpace(string(out)), "ok"; got != want { + t.Errorf("got output %q, want %q", got, want) + } + if got, want := runs(), 2; got != want { + t.Errorf("got %d runs, want %d", got, want) + } +} + +func TestCmdRunRetriesTransientErrorWithStderr(t *testing.T) { + path, runs := fakeNix(t, "error: unable to download 'https://github.com/x': HTTP error 502") + stderr := &bytes.Buffer{} + cmd := &Cmd{Path: path, Args: Args{path}, MaxAttempts: 3, Stderr: stderr} + + if err := cmd.Run(t.Context()); err != nil { + t.Fatalf("got error: %v", err) + } + if got, want := runs(), 2; got != want { + t.Errorf("got %d runs, want %d", got, want) + } + if !strings.Contains(stderr.String(), "HTTP error 502") { + t.Errorf("stderr doesn't contain Nix's original error:\n%s", stderr) + } + if !strings.Contains(stderr.String(), "retrying") { + t.Errorf("stderr doesn't contain a retry message:\n%s", stderr) + } +} + +func TestCmdCombinedOutputRetriesTransientError(t *testing.T) { + path, runs := fakeNix(t, truncatedTarErr) + cmd := &Cmd{Path: path, Args: Args{path}, MaxAttempts: 3} + + if _, err := cmd.CombinedOutput(t.Context()); err != nil { + t.Fatalf("got error: %v", err) + } + if got, want := runs(), 2; got != want { + t.Errorf("got %d runs, want %d", got, want) + } +} + +func TestCmdGivesUpAfterMaxAttempts(t *testing.T) { + path, runs := fakeNix(t, truncatedTarErr, truncatedTarErr, truncatedTarErr) + cmd := &Cmd{Path: path, Args: Args{path}, MaxAttempts: 3} + + _, err := cmd.Output(t.Context()) + if err == nil { + t.Fatal("got nil error, want error after max attempts") + } + if !strings.Contains(err.Error(), "Truncated tar archive") { + t.Errorf("got error %q, want it to contain Nix's error", err) + } + if got, want := runs(), 3; got != want { + t.Errorf("got %d runs, want %d", got, want) + } +} + +func TestCmdDoesNotRetry(t *testing.T) { + tests := []struct { + name string + stderr string + modify func(*Cmd) + }{ + { + name: "NonTransientError", + stderr: "error: flake 'path:/x' does not provide attribute 'packages.aarch64-darwin.foo'", + }, + { + name: "RateLimit", + stderr: "error: unable to download 'https://api.github.com/x': HTTP error 403", + }, + { + name: "RetriesDisabled", + stderr: truncatedTarErr, + modify: func(c *Cmd) { c.MaxAttempts = 0 }, + }, + { + name: "NonFileStdin", + stderr: truncatedTarErr, + modify: func(c *Cmd) { c.Stdin = strings.NewReader("input") }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + path, runs := fakeNix(t, test.stderr) + cmd := &Cmd{Path: path, Args: Args{path}, MaxAttempts: 3} + if test.modify != nil { + test.modify(cmd) + } + + if _, err := cmd.Output(t.Context()); err == nil { + t.Fatal("got nil error, want error") + } + if got, want := runs(), 1; got != want { + t.Errorf("got %d runs, want %d", got, want) + } + }) + } +} + +func TestCmdStopsRetryingWhenCanceled(t *testing.T) { + path, runs := fakeNix(t, truncatedTarErr) + retryDelay = time.Hour + + ctx, cancel := context.WithCancel(t.Context()) + cmd := &Cmd{Path: path, Args: Args{path}, MaxAttempts: 3, Stderr: &bytes.Buffer{}} + go func() { + for runs() == 0 { + time.Sleep(10 * time.Millisecond) + } + cancel() + }() + + if err := cmd.Run(ctx); err == nil { + t.Fatal("got nil error, want error") + } + if got, want := runs(), 1; got != want { + t.Errorf("got %d runs, want %d", got, want) + } +} + +func TestTailWriter(t *testing.T) { + w := &tailWriter{} + _, _ = w.Write(bytes.Repeat([]byte("x"), 10<<10)) + _, _ = w.Write([]byte(truncatedTarErr)) + if got, want := len(w.buf), 8<<10; got != want { + t.Errorf("got len %d, want %d", got, want) + } + if !bytes.HasSuffix(w.buf, []byte(truncatedTarErr)) { + t.Error("tail doesn't end with the last write") + } +} + +func TestTailWriterSmallWrites(t *testing.T) { + tail := &tailWriter{} + line := []byte("copying path '/nix/store/xxx' from 'https://cache.nixos.org'\n") + for range 1000 { + _, _ = tail.Write(line) + } + _, _ = tail.Write([]byte(truncatedTarErr)) + if got, want := len(tail.buf), 8<<10; got != want { + t.Errorf("got len %d, want %d", got, want) + } + if got, limit := cap(tail.buf), 16<<10; got > limit { + t.Errorf("got cap %d, want <= %d", got, limit) + } + if !bytes.HasSuffix(tail.buf, []byte(truncatedTarErr)) { + t.Error("tail doesn't end with the last write") + } +} diff --git a/nix/nix.go b/nix/nix.go index f793b606f0d..7c6758f3bae 100644 --- a/nix/nix.go +++ b/nix/nix.go @@ -70,6 +70,10 @@ type Nix struct { // Logger logs information at [slog.LevelDebug] about Nix command // starts and exits. If nil, it defaults to [slog.Default]. Logger *slog.Logger + + // MaxAttempts is the default [Cmd.MaxAttempts] for commands created + // with [Nix.Command]. The zero value disables retries. + MaxAttempts int } // resolvePath resolves the path to the Nix executable. It returns n.Path if it