Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions internal/nix/command.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 != "" {
Expand Down
176 changes: 155 additions & 21 deletions nix/command.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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()

Expand All @@ -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 {
Expand Down
Loading
Loading