From 3099a410160a2cf924457e033f1a6ec279796364 Mon Sep 17 00:00:00 2001 From: sneak Date: Sun, 4 Oct 2026 12:18:02 +0000 Subject: [PATCH] fetch: client timeout, retry with backoff, url.JoinPath (closes #63) fetch now makes every request through an http.Client with a time limit, ten minutes by default and set with --timeout, which must be greater than zero. A connection error, a timeout, or a 5xx or 429 response is retried, up to five tries in all, after a random wait whose limit doubles from one second, or after the wait the server's Retry-After asks for, up to one minute. Each try of a file starts a new temp file, so a retry never keeps a partial file; the size and hash checks are unchanged. Manifest and file URLs are built with URL.JoinPath, so a base URL with or without a trailing slash or with a query string, and names that need escaping, all work. Model: opus-5-5 --- internal/cli/errmsg_test.go | 16 +- internal/cli/fetch.go | 226 ++++++++++++++++----- internal/cli/fetch_test.go | 392 +++++++++++++++++++++++++++++++++++- internal/cli/mfer.go | 10 +- 4 files changed, 581 insertions(+), 63 deletions(-) diff --git a/internal/cli/errmsg_test.go b/internal/cli/errmsg_test.go index 69c8585..1f37556 100644 --- a/internal/cli/errmsg_test.go +++ b/internal/cli/errmsg_test.go @@ -272,10 +272,15 @@ func TestFetchManifestHTTPStatusMessage(t *testing.T) { })) defer server.Close() - set := flag.NewFlagSet("fetch", flag.ContinueOnError) + mfa := &CLIApp{Fs: afero.NewMemMapFs()} + + set := flag.NewFlagSet(cmdFetch, flag.ContinueOnError) + for _, f := range mfa.fetchCommand().Flags { + require.NoError(t, f.Apply(set)) + } + require.NoError(t, set.Parse([]string{server.URL})) - mfa := &CLIApp{Fs: afero.NewMemMapFs()} ctx := urfcli.NewContext(nil, set, nil) // fetchManifestOperation logs to the process-global logger. @@ -293,8 +298,11 @@ func TestFetchFileHTTPStatusMessage(t *testing.T) { })) defer server.Close() - err := downloadFile(context.Background(), server.URL+"/x", "x", - &mfer.MFFilePath{}, nil) + // downloadFile logs each retry of the 500 to the process-global logger. + err := runLocked(func() error { + return downloadFile(context.Background(), testClient(), server.URL+"/x", "x", + &mfer.MFFilePath{}, nil) + }) require.ErrorIs(t, err, errHTTPStatus) assert.EqualError(t, err, "HTTP 500") } diff --git a/internal/cli/fetch.go b/internal/cli/fetch.go index d6d9917..48020ea 100644 --- a/internal/cli/fetch.go +++ b/internal/cli/fetch.go @@ -7,11 +7,13 @@ import ( "errors" "fmt" "io" + "math/rand/v2" + "net" "net/http" "net/url" "os" - "path" "path/filepath" + "strconv" "strings" "time" @@ -45,12 +47,33 @@ const ( bpsPerGbps = 1e9 bpsPerMbps = 1e6 bpsPerKbps = 1e3 + + // httpTimeout is the default time limit for one HTTP request, from + // connecting to reading the last byte of the body. It also bounds how + // long one file may take to download; fetch's --timeout changes it. + httpTimeout = 10 * time.Minute + + // fetchAttempts is how many times fetch tries a request before it + // gives up on a transient failure. + fetchAttempts = 5 + + // firstRetryDelay is the longest fetch waits before its first retry. + // The limit doubles for each retry after that; the wait itself is + // random up to the limit. + firstRetryDelay = time.Second + + // maxRetryAfter is the longest wait a server's Retry-After header can + // ask for. A server asking for longer fails the request at once. + maxRetryAfter = time.Minute ) var ( // errURLRequired indicates the fetch command was run without a URL // argument. errURLRequired = errors.New("URL argument required") + // errInvalidTimeout indicates a fetch --timeout of zero or less, which + // http.Client would take as no time limit at all. + errInvalidTimeout = errors.New("--timeout must be greater than zero") // errEmptyPath indicates an empty file path in the manifest. errEmptyPath = errors.New("empty path") // errAbsolutePath indicates an absolute file path in the manifest. @@ -78,24 +101,109 @@ type DownloadProgress struct { ETA time.Duration // Estimated time to completion } -// httpGet issues a GET request for the given URL using the provided -// context and returns the response. The caller must close the body. +// retryingClient is the HTTP client fetch makes every request with. It +// waits up to firstDelay, which must be positive, before its first retry +// of a transient failure. Tests shorten both the client's timeout and +// firstDelay. +type retryingClient struct { + client *http.Client + firstDelay time.Duration +} + +// get issues a GET for rawURL and passes a 200 OK response to use. // -// Errors are returned unwrapped: this helper replaced direct http.Get -// calls, and each caller already supplies its own context string, so -// adding one here would change user-visible messages. -func httpGet(ctx context.Context, fileURL string) (*http.Response, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, fileURL, nil) +// A connection error, a timeout, or a 5xx or 429 status is retried, up to +// fetchAttempts tries in all. Before each retry it waits a random time up +// to a limit that doubles each time, unless the server's Retry-After +// header says how long to wait. Any other failure, an error from use +// included, is returned at once and unwrapped; a non-OK status is +// returned as errHTTPStatus. use must start over each time it is called, +// since a retry calls it again with a new response. +func (c retryingClient) get( + ctx context.Context, rawURL string, use func(*http.Response) error, +) error { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil) if err != nil { - return nil, err + return err } - resp, err := http.DefaultClient.Do(req) - if err != nil { - return nil, err + delay := c.firstDelay + + for attempt := 1; ; attempt++ { + // Wait a random time up to delay, so that clients that failed + // together do not all retry together. + wait := rand.N(delay) //nolint:gosec // G404: jitter, not a secret + + var retry bool + + resp, err := c.client.Do(req) + + switch { + case err != nil: + retry = isConnectionError(err) + case resp.StatusCode == http.StatusOK: + err = use(resp) + retry = isConnectionError(err) + + _ = resp.Body.Close() + default: + _ = resp.Body.Close() + + err = fmt.Errorf("%w %d", errHTTPStatus, resp.StatusCode) + retry = resp.StatusCode >= http.StatusInternalServerError || + resp.StatusCode == http.StatusTooManyRequests + + after, ok := retryAfter(resp.Header.Get("Retry-After")) + if ok { + wait = after + retry = retry && after <= maxRetryAfter + } + } + + if !retry || attempt == fetchAttempts || ctx.Err() != nil { + return err + } + + log.Warnf("%s: %s, retrying in %s", rawURL, err, wait.Round(time.Millisecond)) + + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(wait): + } + + delay *= 2 + } +} + +// isConnectionError reports whether err is a connection error or a +// timeout: a connection that could not be made, was reset, or closed +// early, or a request that ran out of time. +func isConnectionError(err error) bool { + var ( + opErr *net.OpError + netErr net.Error + ) + + return errors.As(err, &opErr) || + (errors.As(err, &netErr) && netErr.Timeout()) || + errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) +} + +// retryAfter returns the wait a Retry-After header value asks for, given +// either in seconds or as an HTTP date. ok is false if there is none. +func retryAfter(value string) (time.Duration, bool) { + seconds, err := strconv.Atoi(value) + if err == nil { + return time.Duration(seconds) * time.Second, true } - return resp, nil + when, err := http.ParseTime(value) + if err == nil { + return time.Until(when), true + } + + return 0, false } // reportDownloadProgress renders download progress until the channel @@ -119,25 +227,23 @@ func reportDownloadProgress(progress <-chan DownloadProgress, done chan<- struct } // manifestBaseURL returns the URL of the directory containing the -// manifest, with a trailing slash. +// manifest. func manifestBaseURL(manifestURL string) (*url.URL, error) { - baseURL, err := url.Parse(manifestURL) + parsed, err := url.Parse(manifestURL) if err != nil { return nil, fmt.Errorf("fetch: invalid manifest URL: %w", err) } - baseURL.Path = path.Dir(baseURL.Path) - if !strings.HasSuffix(baseURL.Path, "/") { - baseURL.Path += "/" - } - - return baseURL, nil + // JoinPath cleans the path it builds, so ".." drops the manifest's + // file name. + return parsed.JoinPath(".."), nil } // downloadManifestFiles downloads every file in the manifest, reporting // progress on the progress channel. func downloadManifestFiles( ctx context.Context, + client retryingClient, baseURL *url.URL, files []*mfer.MFFilePath, progress chan<- DownloadProgress, @@ -149,10 +255,12 @@ func downloadManifestFiles( return fmt.Errorf("invalid path in manifest: %w", err) } - fileURL := baseURL.String() + encodeFilePath(f.GetPath()) + // JoinPath takes escaped path text, so a name such as "100%.txt" + // must be escaped first. + fileURL := baseURL.JoinPath(encodeFilePath(f.GetPath())).String() log.Infof("fetching %s", f.GetPath()) - err = downloadFile(ctx, fileURL, localPath, f, progress) + err = downloadFile(ctx, client, fileURL, localPath, f, progress) if err != nil { return fmt.Errorf("failed to download %s: %w", f.GetPath(), err) } @@ -168,6 +276,11 @@ func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error { return errURLRequired } + timeout := ctx.Duration(flagTimeout) + if timeout <= 0 { + return errInvalidTimeout + } + inputURL := ctx.Args().Get(0) manifestURL, err := resolveManifestURL(inputURL) @@ -175,23 +288,31 @@ func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error { return fmt.Errorf("invalid URL: %w", err) } + client := retryingClient{ + client: &http.Client{Timeout: timeout}, + firstDelay: firstRetryDelay, + } + log.Infof("fetching manifest from %s", manifestURL) - // Fetch manifest - resp, err := httpGet(ctx.Context, manifestURL) + // Read the whole manifest before parsing it, so that a connection + // lost partway through is retried rather than reported as a bad + // manifest. + var manifestData []byte + + err = client.get(ctx.Context, manifestURL, func(resp *http.Response) error { + var readErr error + + manifestData, readErr = io.ReadAll(resp.Body) + + return readErr + }) if err != nil { return fmt.Errorf("failed to fetch manifest: %w", err) } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("failed to fetch manifest: %w %d", - errHTTPStatus, resp.StatusCode) - } - // Parse manifest - manifest, err := mfer.NewManifestFromReader(resp.Body) + manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData)) if err != nil { return fmt.Errorf("failed to parse manifest: %w", err) } @@ -221,7 +342,7 @@ func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error { startTime := time.Now() // Download each file - dlErr := downloadManifestFiles(ctx.Context, baseURL, files, progress) + dlErr := downloadManifestFiles(ctx.Context, client, baseURL, files, progress) close(progress) <-done @@ -325,14 +446,7 @@ func resolveManifestURL(inputURL string) (string, error) { return inputURL, nil } - // Ensure path ends with / - if !strings.HasSuffix(parsed.Path, "/") { - parsed.Path += "/" - } - - parsed.Path += defaultManifestName - - return parsed.String(), nil + return parsed.JoinPath(defaultManifestName).String(), nil } // progressWriter wraps an io.Writer and reports progress to a channel. @@ -441,6 +555,7 @@ func verifyDownloadedHash(digest []byte, entry *mfer.MFFilePath) error { // Progress is reported via the progress channel. func downloadFile( ctx context.Context, + client retryingClient, fileURL, localPath string, entry *mfer.MFFilePath, progress chan<- DownloadProgress, @@ -468,18 +583,21 @@ func downloadFile( tmpPath := tempPathFor(localPath) - // Fetch file - resp, err := httpGet(ctx, fileURL) - if err != nil { - return fmt.Errorf("HTTP request failed: %w", err) - } - - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("%w %d", errHTTPStatus, resp.StatusCode) - } + return client.get(ctx, fileURL, func(resp *http.Response) error { + return saveResponse(resp, tmpPath, localPath, entry, progress) + }) +} +// saveResponse writes resp's body to tmpPath, verifies it against entry, +// and renames it to localPath. It starts a new temp file each time and +// removes it on failure, so a retry after a failed try never appends to +// or keeps a partial file. +func saveResponse( + resp *http.Response, + tmpPath, localPath string, + entry *mfer.MFFilePath, + progress chan<- DownloadProgress, +) error { // Determine expected size expectedSize := entry.GetSize() @@ -488,7 +606,7 @@ func downloadFile( totalBytes = expectedSize } - err = checkNoSymlinks(tmpPath) + err := checkNoSymlinks(tmpPath) if err != nil { return err } diff --git a/internal/cli/fetch_test.go b/internal/cli/fetch_test.go index 3095475..82ad178 100644 --- a/internal/cli/fetch_test.go +++ b/internal/cli/fetch_test.go @@ -4,16 +4,24 @@ package cli import ( "bytes" "context" + "flag" + "fmt" "io" + "net" "net/http" "net/http/httptest" "os" "path/filepath" + "strconv" + "sync" + "sync/atomic" "testing" + "time" "github.com/spf13/afero" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + urfcli "github.com/urfave/cli/v2" "sneak.berlin/go/mfer/mfer" ) @@ -214,6 +222,41 @@ func fetchTestHandler( } } +// manifestOf scans a tree holding files and returns its manifest bytes. +func manifestOf(t *testing.T, files map[string][]byte) []byte { + t.Helper() + + sourceFs := afero.NewMemMapFs() + for p, content := range files { + require.NoError(t, sourceFs.MkdirAll(filepath.Dir("/"+p), 0o755)) + require.NoError(t, afero.WriteFile(sourceFs, "/"+p, content, 0o644)) + } + + return scanToManifest(t, sourceFs) +} + +// testClient returns the client tests download with. Its retries wait +// milliseconds rather than seconds. +func testClient() retryingClient { + return retryingClient{ + client: &http.Client{Timeout: 10 * time.Second}, + firstDelay: time.Millisecond, + } +} + +// getNothing calls client.get for rawURL with a use that reads nothing, +// under a context that expires after timeout. It holds runMu, since get +// logs each retry to the process-global logger, and starts the timeout +// only once it has the lock. +func getNothing(client retryingClient, rawURL string, timeout time.Duration) error { + return runLocked(func() error { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + return client.get(ctx, rawURL, func(*http.Response) error { return nil }) + }) +} + //nolint:paralleltest // changes the process-global working directory func TestFetchFromHTTP(t *testing.T) { // Create source filesystem with test files @@ -266,7 +309,8 @@ func TestFetchFromHTTP(t *testing.T) { require.NoError(t, err) fileURL := baseURL + f.GetPath() - err = downloadFile(context.Background(), fileURL, localPath, f, progress) + err = downloadFile(context.Background(), testClient(), + fileURL, localPath, f, progress) require.NoError(t, err, "failed to download %s", f.GetPath()) } @@ -313,7 +357,7 @@ func TestFetchHashMismatch(t *testing.T) { chdirTemp(t) // Try to download - should fail with hash mismatch - err = downloadFile(context.Background(), + err = downloadFile(context.Background(), testClient(), server.URL+"/file.txt", testFileTxt, files[0], nil) require.Error(t, err) assert.Contains(t, err.Error(), "mismatch") @@ -359,7 +403,7 @@ func TestFetchSizeMismatch(t *testing.T) { chdirTemp(t) // Try to download - should fail with size mismatch - err = downloadFile(context.Background(), + err = downloadFile(context.Background(), testClient(), server.URL+"/file.txt", testFileTxt, files[0], nil) require.Error(t, err) assert.Contains(t, err.Error(), "size mismatch") @@ -417,7 +461,7 @@ func TestFetchProgress(t *testing.T) { }() // Download - err = downloadFile(context.Background(), + err = downloadFile(context.Background(), testClient(), server.URL+"/large.txt", "large.txt", files[0], progress) close(progress) <-done @@ -523,3 +567,343 @@ func TestFetchReplacesHardLinkAtTempName(t *testing.T) { require.NoError(t, err) assert.Equal(t, "outside", string(outside), "fetch wrote outside the destination") } + +// TestGetRetriesTransientStatusesOnly answers every request with one +// status and counts the requests: a 5xx or 429 is retried until +// fetchAttempts runs out, and any other status fails at once. +func TestGetRetriesTransientStatusesOnly(t *testing.T) { + t.Parallel() + + tests := []struct { + status int + requests int32 + }{ + {http.StatusNotFound, 1}, + {http.StatusForbidden, 1}, + {http.StatusInternalServerError, fetchAttempts}, + {http.StatusServiceUnavailable, fetchAttempts}, + {http.StatusTooManyRequests, fetchAttempts}, + } + + for _, tt := range tests { + t.Run(http.StatusText(tt.status), func(t *testing.T) { + t.Parallel() + + var requests atomic.Int32 + + server := httptest.NewServer( + http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + w.WriteHeader(tt.status) + })) + defer server.Close() + + err := getNothing(testClient(), server.URL, 10*time.Second) + require.ErrorIs(t, err, errHTTPStatus) + require.EqualError(t, err, fmt.Sprintf("HTTP %d", tt.status)) + assert.Equal(t, tt.requests, requests.Load()) + }) + } +} + +// TestGetTimesOutOnStalledServer points get at a server that takes a +// request and never answers it. The try must give up at the client's +// timeout and be retried; when every try stalls, get must return a +// timeout error rather than hang. +func TestGetTimesOutOnStalledServer(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + stallEvery bool + }{ + {"first request stalls", false}, + {"every request stalls", true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + var requests atomic.Int32 + + stop := make(chan struct{}) + server := httptest.NewServer( + http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + if requests.Add(1) == 1 || tt.stallEvery { + select { + case <-r.Context().Done(): + case <-stop: + } + } + })) + + defer server.Close() + defer close(stop) + + client := testClient() + client.client.Timeout = 50 * time.Millisecond + + err := getNothing(client, server.URL, 10*time.Second) + if !tt.stallEvery { + require.NoError(t, err) + + return + } + + var netErr net.Error + require.ErrorAs(t, err, &netErr) + assert.True(t, netErr.Timeout(), "want a timeout, got %v", err) + }) + } +} + +// TestGetHonorsRetryAfter answers the first request with a 503 and a +// Retry-After header, and any later one with 200 OK. The client's own +// wait before a retry is an hour, so get finishes in time only if it +// waits as long as Retry-After says instead. A Retry-After longer than +// maxRetryAfter must fail at once. +func TestGetHonorsRetryAfter(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value string + requests int32 + wantErr bool + }{ + {"seconds", "0", 2, false}, + {"HTTP date", time.Now().Add(-time.Minute).UTC().Format(http.TimeFormat), 2, false}, + {"too long", strconv.Itoa(int((maxRetryAfter + time.Second).Seconds())), 1, true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + var requests atomic.Int32 + + server := httptest.NewServer( + http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if requests.Add(1) == 1 { + w.Header().Set("Retry-After", tt.value) + w.WriteHeader(http.StatusServiceUnavailable) + } + })) + defer server.Close() + + client := testClient() + client.firstDelay = time.Hour + + err := getNothing(client, server.URL, 10*time.Second) + if tt.wantErr { + require.ErrorIs(t, err, errHTTPStatus) + } else { + require.NoError(t, err) + } + + assert.Equal(t, tt.requests, requests.Load()) + }) + } +} + +// TestGetStopsWaitingWhenCanceled gives get a server that always answers +// 503 and an hour's wait before each retry, then lets the context expire +// during that wait. get returning at all shows it stopped waiting. +func TestGetStopsWaitingWhenCanceled(t *testing.T) { + t.Parallel() + + server := httptest.NewServer( + http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusServiceUnavailable) + })) + defer server.Close() + + client := testClient() + client.firstDelay = time.Hour + + err := getNothing(client, server.URL, 50*time.Millisecond) + require.ErrorIs(t, err, context.DeadlineExceeded) +} + +// TestDownloadFileRetriesToSuccess serves a file that fails twice before +// it succeeds: first with a 503, then with a body cut off halfway. The +// download must end with the whole, verified file and no temp file. +// +//nolint:paralleltest // changes the process-global working directory +func TestDownloadFileRetriesToSuccess(t *testing.T) { + content := []byte("the whole file, every byte of it") + + manifest, err := mfer.NewManifestFromReader(bytes.NewReader( + manifestOf(t, map[string][]byte{testFileTxt: content}))) + require.NoError(t, err) + + var requests atomic.Int32 + + server := httptest.NewServer( + http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + switch requests.Add(1) { + case 1: + w.WriteHeader(http.StatusServiceUnavailable) + case 2: + // Promise the whole file but send half of it; the server + // then closes the connection. + w.Header().Set("Content-Length", strconv.Itoa(len(content))) + _, _ = w.Write(content[:len(content)/2]) + default: + _, _ = w.Write(content) + } + })) + defer server.Close() + + chdirTemp(t) + + err = downloadFile(context.Background(), testClient(), + server.URL+"/"+testFileTxt, testFileTxt, manifest.Files()[0], nil) + require.NoError(t, err) + assert.Equal(t, int32(3), requests.Load()) + + fetched, err := os.ReadFile(testFileTxt) + require.NoError(t, err) + assert.Equal(t, content, fetched) + + _, err = os.Stat(".file.txt.tmp") + assert.True(t, os.IsNotExist(err), "temp file left behind") +} + +// TestFetchBaseURLForms fetches one tree by its directory URL with and +// without a trailing slash, with a query string, and by its manifest URL. +// Each must request the same manifest and file paths. +// +//nolint:paralleltest // changes the process-global working directory +func TestFetchBaseURLForms(t *testing.T) { + files := map[string][]byte{"one.txt": []byte("1"), "sub/two.txt": []byte("2")} + tree := http.StripPrefix("/tree", fetchTestHandler(manifestOf(t, files), files)) + + var ( + mu sync.Mutex + requested []string + ) + + server := httptest.NewServer( + http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + + requested = append(requested, r.URL.Path) + + mu.Unlock() + + tree.ServeHTTP(w, r) + })) + defer server.Close() + + inputs := []string{"/tree", "/tree/", "/tree?key=value", "/tree/index.mf"} + for _, input := range inputs { + t.Run(input, func(t *testing.T) { + mu.Lock() + requested = nil + mu.Unlock() + + chdirTemp(t) + + opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL + input}, + afero.NewOsFs()) + require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) + + mu.Lock() + defer mu.Unlock() + + assert.ElementsMatch(t, + []string{"/tree/index.mf", "/tree/one.txt", "/tree/sub/two.txt"}, requested) + }) + } +} + +// TestFetchEscapedPaths fetches files whose names need escaping in a URL, +// one of them a name that is itself an escape sequence. Each must arrive +// under its own name with its own content. +// +//nolint:paralleltest // changes the process-global working directory +func TestFetchEscapedPaths(t *testing.T) { + files := map[string][]byte{ + "a b.txt": []byte("space"), + "a%20b.txt": []byte("percent sign, two, zero"), + "100%.txt": []byte("percent sign"), + "q?x=1#frag.txt": []byte("question mark and hash"), + "dir with space/sub.txt": []byte("directory with a space"), + } + + server := httptest.NewServer(fetchTestHandler(manifestOf(t, files), files)) + defer server.Close() + + chdirTemp(t) + + opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) + require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) + + for name, content := range files { + fetched, err := os.ReadFile(name) //nolint:gosec // test-controlled path + require.NoError(t, err) + assert.Equal(t, content, fetched, name) + } +} + +// TestFetchTimeoutFlag runs fetch with --timeout against a server that +// never answers. Without the flag's limit the request would wait forever; +// once fetch gives up on it, the server cancels fetch's context so that +// fetch returns instead of retrying. +func TestFetchTimeoutFlag(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + server := httptest.NewServer( + http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + <-r.Context().Done() // fetch gave up on this request + cancel() + })) + defer server.Close() + + mfa := &CLIApp{Fs: afero.NewMemMapFs()} + + set := flag.NewFlagSet(cmdFetch, flag.ContinueOnError) + for _, f := range mfa.fetchCommand().Flags { + require.NoError(t, f.Apply(set)) + } + + require.NoError(t, set.Parse([]string{"--" + flagTimeout, "100ms", server.URL})) + + cliCtx := urfcli.NewContext(nil, set, nil) + cliCtx.Context = ctx + + // fetchManifestOperation logs to the process-global logger. + err := runLocked(func() error { return mfa.fetchManifestOperation(cliCtx) }) + require.Error(t, err) +} + +// TestFetchRejectsTimeoutOfZeroOrLess runs fetch with a --timeout of zero +// and of less than zero. http.Client takes either as no time limit at all, +// so fetch must refuse it before making any request. +func TestFetchRejectsTimeoutOfZeroOrLess(t *testing.T) { + t.Parallel() + + var requests atomic.Int32 + + server := httptest.NewServer( + http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + requests.Add(1) + })) + defer server.Close() + + for _, timeout := range []string{"0", "-1s"} { + opts := testOpts( + []string{testApp, cmdFetch, "--" + flagTimeout + "=" + timeout, server.URL}, + afero.NewMemMapFs()) + + assert.Equal(t, 1, runCLI(opts), timeout) + assert.Contains(t, testStderr(t, opts), errInvalidTimeout.Error(), timeout) + } + + assert.Zero(t, requests.Load(), "fetch made a request") +} diff --git a/internal/cli/mfer.go b/internal/cli/mfer.go index 171a1da..0e12299 100644 --- a/internal/cli/mfer.go +++ b/internal/cli/mfer.go @@ -23,6 +23,7 @@ const ( cmdVersion = "version" flagProgress = "progress" + flagTimeout = "timeout" manifestArgsUsage = "[manifest file]" @@ -346,7 +347,14 @@ func (mfa *CLIApp) fetchCommand() *cli.Command { return mfa.fetchManifestOperation(c) }, - Flags: commonFlags(), + Flags: append(commonFlags(), + &cli.DurationFlag{ + Name: flagTimeout, + Value: httpTimeout, + Usage: "Time limit for each HTTP request, including the download " + + "of its body", + }, + ), } }