//nolint:testpackage // white-box tests exercise unexported internals package cli import ( "bytes" "context" "flag" "fmt" "io" "maps" "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" ) const ( testFileTxt = "file.txt" testDirFile = "dir/file.txt" testIndexMF = "https://example.com/path/index.mf" // Exactly what url.Parse renders, with no wrapper of our own. urlParseControlCharErr = `parse "http://example.com/\x7f": ` + `net/url: invalid control character in URL` ) func TestEncodeFilePath(t *testing.T) { t.Parallel() tests := []struct { input string expected string }{ {testFileTxt, testFileTxt}, {testDirFile, testDirFile}, {"my file.txt", "my%20file.txt"}, {"dir/my file.txt", "dir/my%20file.txt"}, {"file#1.txt", "file%231.txt"}, {"file?v=1.txt", "file%3Fv=1.txt"}, {"path/to/file with spaces.txt", "path/to/file%20with%20spaces.txt"}, {"100%done.txt", "100%25done.txt"}, } for _, tt := range tests { t.Run(tt.input, func(t *testing.T) { t.Parallel() result := encodeFilePath(tt.input) assert.Equal(t, tt.expected, result) }) } } func TestSanitizePath(t *testing.T) { t.Parallel() // Valid paths that should be accepted validTests := []struct { input string expected string }{ {testFileTxt, testFileTxt}, {testDirFile, testDirFile}, {"dir/subdir/file.txt", "dir/subdir/file.txt"}, {"./file.txt", testFileTxt}, {"./dir/file.txt", testDirFile}, {"dir/./file.txt", testDirFile}, } for _, tt := range validTests { t.Run("valid:"+tt.input, func(t *testing.T) { t.Parallel() result, err := sanitizePath(tt.input) require.NoError(t, err) assert.Equal(t, tt.expected, result) }) } // Invalid paths that should be rejected invalidTests := []struct { input string desc string }{ {"", "empty path"}, {"..", "parent directory"}, {"../file.txt", "parent traversal"}, {"../../file.txt", "double parent traversal"}, {"dir/../../../file.txt", "traversal escaping base"}, {"/etc/passwd", "absolute path"}, {"/file.txt", "absolute path with single component"}, {"dir/../../etc/passwd", "traversal to system file"}, } for _, tt := range invalidTests { t.Run("invalid:"+tt.desc, func(t *testing.T) { t.Parallel() _, err := sanitizePath(tt.input) assert.Error(t, err, "expected error for path: %s", tt.input) }) } } func TestResolveManifestURL(t *testing.T) { t.Parallel() tests := []struct { input string expected string }{ // Already ends with .mf - use as-is {testIndexMF, testIndexMF}, {"https://example.com/path/custom.mf", "https://example.com/path/custom.mf"}, {"https://example.com/foo.mf", "https://example.com/foo.mf"}, // Directory with trailing slash - append index.mf {"https://example.com/path/", testIndexMF}, {"https://example.com/", "https://example.com/index.mf"}, // Directory without trailing slash - add slash and index.mf {"https://example.com/path", testIndexMF}, {"https://example.com", "https://example.com/index.mf"}, // With query strings { "https://example.com/path?foo=bar", "https://example.com/path/index.mf?foo=bar", }, } for _, tt := range tests { t.Run(tt.input, func(t *testing.T) { t.Parallel() result, err := resolveManifestURL(tt.input) require.NoError(t, err) assert.Equal(t, tt.expected, result) }) } // The sole caller wraps this error as "invalid URL: %w", so // resolveManifestURL must return url.Parse's error unadorned. t.Run("invalid:control character", func(t *testing.T) { t.Parallel() _, err := resolveManifestURL("http://example.com/\x7f") require.ErrorContains(t, err, urlParseControlCharErr) assert.NotContains(t, err.Error(), "failed to parse URL") }) } // scanToManifest scans sourceFs and returns the serialized manifest bytes. func scanToManifest(t *testing.T, sourceFs afero.Fs) []byte { t.Helper() s := mfer.NewScannerWithOptions(&mfer.ScannerOptions{Fs: sourceFs}) require.NoError(t, s.EnumerateFS(sourceFs, "/", nil)) var manifestBuf bytes.Buffer require.NoError(t, s.ToManifest(context.Background(), &manifestBuf, nil)) return manifestBuf.Bytes() } // chdirTemp switches the working directory to a fresh temp dir for the // duration of the test and returns its path. func chdirTemp(t *testing.T) string { t.Helper() destDir := t.TempDir() origDir, err := os.Getwd() require.NoError(t, err) require.NoError(t, os.Chdir(destDir)) t.Cleanup(func() { _ = os.Chdir(origDir) }) return destDir } // fetchTestHandler serves the manifest at /index.mf and the given files // at their paths. func fetchTestHandler( manifestData []byte, testFiles map[string][]byte, ) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { path := r.URL.Path if path == "/index.mf" { w.Header().Set("Content-Type", "application/octet-stream") _, _ = w.Write(manifestData) return } // Strip leading slash if len(path) > 0 && path[0] == '/' { path = path[1:] } content, exists := testFiles[path] if !exists { http.NotFound(w, r) return } w.Header().Set("Content-Type", "application/octet-stream") _, _ = w.Write(content) } } // 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) } // filesUnder returns the content of every file under dir, by its path // relative to dir. func filesUnder(t *testing.T, dir string) map[string][]byte { t.Helper() files := map[string][]byte{} err := filepath.WalkDir(dir, func(path string, entry os.DirEntry, err error) error { if err != nil || entry.IsDir() { return err } rel, err := filepath.Rel(dir, path) if err != nil { return err } content, err := os.ReadFile(path) //nolint:gosec // test-controlled path files[filepath.ToSlash(rel)] = content return err }) require.NoError(t, err) return files } // 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 sourceFs := afero.NewMemMapFs() testFiles := map[string][]byte{ "file1.txt": []byte("Hello, World!"), "file2.txt": []byte("This is file 2 with more content."), "subdir/file3.txt": []byte("Nested file content here."), "subdir/deep/f.txt": []byte("Deeply nested file."), } for path, content := range testFiles { fullPath := "/" + path // MemMapFs needs absolute paths dir := filepath.Dir(fullPath) require.NoError(t, sourceFs.MkdirAll(dir, 0o755)) require.NoError(t, afero.WriteFile(sourceFs, fullPath, content, 0o644)) } // Generate manifest using scanner manifestData := scanToManifest(t, sourceFs) // Create HTTP server that serves the source filesystem server := httptest.NewServer(fetchTestHandler(manifestData, testFiles)) defer server.Close() // Change to a fresh destination directory for the test destDir := chdirTemp(t) // Parse the manifest to get file entries manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData)) require.NoError(t, err) files := manifest.Files() require.Len(t, files, len(testFiles)) // Download each file using downloadFile progress := make(chan DownloadProgress, 10) go func() { for p := range progress { _ = p // drain progress channel } }() baseURL := server.URL + "/" for _, f := range files { localPath, err := sanitizePath(f.GetPath()) require.NoError(t, err) fileURL := baseURL + f.GetPath() err = downloadFile(context.Background(), testClient(), fileURL, ".", localPath, f, progress) require.NoError(t, err, "failed to download %s", f.GetPath()) } close(progress) // Verify downloaded files match originals for path, expectedContent := range testFiles { downloadedPath := filepath.Join(destDir, path) //nolint:gosec // test-controlled path downloadedContent, err := os.ReadFile(downloadedPath) require.NoError(t, err, "failed to read downloaded file %s", path) assert.Equal(t, expectedContent, downloadedContent, "content mismatch for %s", path) } } //nolint:paralleltest // changes the process-global working directory func TestFetchHashMismatch(t *testing.T) { // Create source filesystem with a test file sourceFs := afero.NewMemMapFs() originalContent := []byte("Original content") require.NoError(t, afero.WriteFile(sourceFs, "/file.txt", originalContent, 0o644)) // Generate and parse manifest manifestData := scanToManifest(t, sourceFs) manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData)) require.NoError(t, err) files := manifest.Files() require.Len(t, files, 1) // Create server that serves DIFFERENT content (to trigger hash mismatch) tamperedContent := []byte("Tampered content!") server := httptest.NewServer( http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/octet-stream") _, _ = w.Write(tamperedContent) })) defer server.Close() // Work in a fresh temp directory chdirTemp(t) // Try to download - should fail with hash mismatch err = downloadFile(context.Background(), testClient(), server.URL+"/file.txt", ".", testFileTxt, files[0], nil) require.Error(t, err) assert.Contains(t, err.Error(), "mismatch") // Verify temp file was cleaned up _, err = os.Stat(".file.txt.tmp") assert.True(t, os.IsNotExist(err), "temp file should be cleaned up on hash mismatch") // Verify final file was not created _, err = os.Stat(testFileTxt) assert.True(t, os.IsNotExist(err), "final file should not exist on hash mismatch") } //nolint:paralleltest // changes the process-global working directory func TestFetchSizeMismatch(t *testing.T) { // Create source filesystem with a test file sourceFs := afero.NewMemMapFs() originalContent := []byte("Original content with specific size") require.NoError(t, afero.WriteFile(sourceFs, "/file.txt", originalContent, 0o644)) // Generate and parse manifest manifestData := scanToManifest(t, sourceFs) manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData)) require.NoError(t, err) files := manifest.Files() require.Len(t, files, 1) // Create server that serves content with wrong size wrongSizeContent := []byte("Short") server := httptest.NewServer( http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/octet-stream") _, _ = w.Write(wrongSizeContent) })) defer server.Close() // Work in a fresh temp directory chdirTemp(t) // Try to download - should fail with size mismatch err = downloadFile(context.Background(), testClient(), server.URL+"/file.txt", ".", testFileTxt, files[0], nil) require.Error(t, err) assert.Contains(t, err.Error(), "size mismatch") // Verify temp file was cleaned up _, err = os.Stat(".file.txt.tmp") assert.True(t, os.IsNotExist(err), "temp file should be cleaned up on size mismatch") } //nolint:paralleltest // changes the process-global working directory func TestFetchProgress(t *testing.T) { // Create source filesystem with a larger test file sourceFs := afero.NewMemMapFs() // Create content large enough to trigger multiple progress updates content := bytes.Repeat([]byte("x"), 100*1024) // 100KB require.NoError(t, afero.WriteFile(sourceFs, "/large.txt", content, 0o644)) // Generate and parse manifest manifestData := scanToManifest(t, sourceFs) manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData)) require.NoError(t, err) files := manifest.Files() require.Len(t, files, 1) // Create server that serves the content server := httptest.NewServer( http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/octet-stream") w.Header().Set("Content-Length", "102400") // Write in chunks to allow progress reporting reader := bytes.NewReader(content) _, _ = io.Copy(w, reader) })) defer server.Close() // Work in a fresh temp directory chdirTemp(t) // Set up progress channel and collect updates progress := make(chan DownloadProgress, 100) var progressUpdates []DownloadProgress done := make(chan struct{}) go func() { for p := range progress { progressUpdates = append(progressUpdates, p) } close(done) }() // Download err = downloadFile(context.Background(), testClient(), server.URL+"/large.txt", ".", "large.txt", files[0], progress) close(progress) <-done require.NoError(t, err) // Verify we got progress updates assert.NotEmpty(t, progressUpdates, "should have received progress updates") // Verify final progress shows complete if len(progressUpdates) > 0 { last := progressUpdates[len(progressUpdates)-1] assert.Equal(t, int64(len(content)), last.BytesRead, "final progress should show all bytes read") assert.Equal(t, "large.txt", last.Path) } // Verify file was downloaded correctly downloaded, err := os.ReadFile("large.txt") require.NoError(t, err) assert.Equal(t, content, downloaded) } // TestFetchRefusesSymlinks runs fetch into a destination directory that // holds a symlink pointing outside it, in each of the three places fetch // writes: a parent directory, the temp file, and the file itself, which // the temp file is renamed onto; and once as a directory inside a plain // directory. The fetch must fail and nothing outside may change. // //nolint:paralleltest // changes the process-global working directory func TestFetchRefusesSymlinks(t *testing.T) { tests := []struct { name string entry string // the manifest's only file link string // symlink placed in the destination directory target string // what link points to, relative to the outside directory }{ {"parent directory", "sub/deeper/file.txt", "sub", "."}, {"directory inside a plain directory", "docs/data/passwd", "docs/data", "."}, {"temp file", testFileTxt, ".file.txt.tmp", "new.txt"}, {"file", testFileTxt, testFileTxt, "new.txt"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { content := []byte("fetched") sourceFs := afero.NewMemMapFs() require.NoError(t, sourceFs.MkdirAll(filepath.Dir("/"+tt.entry), 0o755)) require.NoError(t, afero.WriteFile(sourceFs, "/"+tt.entry, content, 0o644)) server := httptest.NewServer(fetchTestHandler( scanToManifest(t, sourceFs), map[string][]byte{tt.entry: content})) defer server.Close() outside := t.TempDir() chdirTemp(t) require.NoError(t, os.MkdirAll(filepath.Dir(tt.link), 0o750)) require.NoError(t, os.Symlink(filepath.Join(outside, tt.target), tt.link)) opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) assert.Equal(t, 1, runCLI(opts)) assert.Contains(t, testStderr(t, opts), "failed to download "+tt.entry+ ": symlink in path not allowed: "+tt.link) written, err := os.ReadDir(outside) require.NoError(t, err) assert.Empty(t, written, "fetch wrote outside the destination") }) } } // TestFetchReplacesHardLinkAtTempName runs fetch into a destination // directory that holds, at the temp file's name, a hard link to a file // outside it. To fetch that is an ordinary leftover from an interrupted // earlier run: it must replace it and succeed, and the outside file must // not change. // //nolint:paralleltest // changes the process-global working directory func TestFetchReplacesHardLinkAtTempName(t *testing.T) { content := []byte("fetched") sourceFs := afero.NewMemMapFs() require.NoError(t, afero.WriteFile(sourceFs, "/"+testFileTxt, content, 0o644)) server := httptest.NewServer(fetchTestHandler( scanToManifest(t, sourceFs), map[string][]byte{testFileTxt: content})) defer server.Close() outsideFile := filepath.Join(t.TempDir(), "secret.txt") require.NoError(t, os.WriteFile(outsideFile, []byte("outside"), 0o600)) chdirTemp(t) require.NoError(t, os.Link(outsideFile, ".file.txt.tmp")) opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) fetched, err := os.ReadFile(testFileTxt) require.NoError(t, err) assert.Equal(t, content, fetched) outside, err := os.ReadFile(outsideFile) //nolint:gosec // test-controlled path 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) } } // TestFetchTree runs fetch on a tree with nested directories. Every file // the manifest lists must land under its own path with its own content, // beside the manifest, and nothing else may be left in the destination. // //nolint:paralleltest // changes the process-global working directory func TestFetchTree(t *testing.T) { files := map[string][]byte{ "top.txt": []byte("at the top"), "sub/one.txt": []byte("one level down"), "sub/deeper/two.txt": []byte("two levels down"), "other/deep/est.txt": []byte("in a second directory"), } manifest := manifestOf(t, files) server := httptest.NewServer(fetchTestHandler(manifest, files)) defer server.Close() dest := chdirTemp(t) opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) want := maps.Clone(files) want[defaultManifestName] = manifest assert.Equal(t, want, filesUnder(t, dest)) } // TestFetchFailsOnHashMismatch runs fetch against a server that serves a // file with the size the manifest lists but different content. fetch must // exit non-zero and leave no file in the destination. // //nolint:paralleltest // changes the process-global working directory func TestFetchFailsOnHashMismatch(t *testing.T) { listed := map[string][]byte{testDirFile: []byte("original")} served := map[string][]byte{testDirFile: []byte("tampered")} server := httptest.NewServer(fetchTestHandler(manifestOf(t, listed), served)) defer server.Close() dest := chdirTemp(t) opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) assert.Equal(t, 1, runCLI(opts)) assert.Empty(t, filesUnder(t, dest)) } // TestFetchIntoPartlyFilledDestination runs fetch where an interrupted // fetch of an older version of the tree left one file current, one file // out of date and one half written to its temp file, beside a file the // manifest does not list. fetch skips the current file and downloads the // other two. The out-of-date file has the same size as the new version, // so only its hash shows that it must be replaced. The file the manifest // does not list is left alone, and the manifest is saved beside the files. // //nolint:paralleltest // changes the process-global working directory func TestFetchIntoPartlyFilledDestination(t *testing.T) { files := map[string][]byte{ "current.txt": []byte("already fetched"), "sub/changed.txt": []byte("new version"), "sub/partial.txt": []byte("cut off partway"), } unlisted := []byte("not in the manifest") manifest := manifestOf(t, files) tree := fetchTestHandler(manifest, 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() dest := chdirTemp(t) require.NoError(t, os.MkdirAll("sub", 0o750)) require.NoError(t, os.WriteFile("current.txt", files["current.txt"], 0o600)) require.NoError(t, os.WriteFile("sub/changed.txt", []byte("old version"), 0o600)) require.NoError(t, os.WriteFile(tempPathFor("sub/partial.txt"), []byte("cut off"), 0o600)) require.NoError(t, os.WriteFile("unlisted.txt", unlisted, 0o600)) opts := testOpts([]string{testApp, cmdFetch, server.URL}, afero.NewOsFs()) require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) assert.Contains(t, testStderr(t, opts), "skipping current.txt: already present") want := maps.Clone(files) want["unlisted.txt"] = unlisted want[defaultManifestName] = manifest assert.Equal(t, want, filesUnder(t, dest)) mu.Lock() defer mu.Unlock() assert.ElementsMatch(t, []string{ "/" + defaultManifestName, "/sub/changed.txt", "/sub/partial.txt", }, requested) } // TestFetchIntoDest fetches a tree with --dest into a directory that does // not exist yet. The files and the manifest must land there and nowhere // else, and check must pass on the result with no extra files. A second // fetch into the same directory must download nothing but the manifest. // //nolint:paralleltest // changes the process-global working directory func TestFetchIntoDest(t *testing.T) { files := map[string][]byte{ "top.txt": []byte("at the top"), "sub/one.txt": []byte("one level down"), } manifest := manifestOf(t, files) tree := fetchTestHandler(manifest, 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() cwd := chdirTemp(t) dest := filepath.Join(t.TempDir(), "mirror") fetch := []string{testApp, cmdFetch, "-q", "--" + flagDest, dest, server.URL} opts := testOpts(fetch, afero.NewOsFs()) require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) want := maps.Clone(files) want[defaultManifestName] = manifest assert.Equal(t, want, filesUnder(t, dest)) assert.Empty(t, filesUnder(t, cwd), "fetch wrote outside --dest") check := testOpts([]string{ testApp, cmdCheck, "-q", testFlagBase, dest, testFlagNoExtra, filepath.Join(dest, defaultManifestName), }, afero.NewOsFs()) require.Equal(t, 0, runCLI(check), testStderr(t, check)) mu.Lock() requested = nil mu.Unlock() opts = testOpts(fetch, afero.NewOsFs()) require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) assert.Equal(t, want, filesUnder(t, dest)) mu.Lock() defer mu.Unlock() assert.Equal(t, []string{"/" + defaultManifestName}, requested) } // TestFetchRequireSignature runs fetch with --require-signature. A // manifest that is unsigned, or signed by another key, must stop fetch // with check's message before it downloads or writes anything; the // required key lets it through. The signed cases need gpg and are skipped // without it, as the other signing tests are. // //nolint:paralleltest // signedManifest calls t.Setenv, which bars t.Parallel func TestFetchRequireSignature(t *testing.T) { files := map[string][]byte{testFileTxt: []byte("signed file")} t.Run("unsigned", func(t *testing.T) { assertFetchRefused(t, manifestOf(t, files), files, msgFpA, "manifest is not signed, but signature from "+msgFpA+" is required") }) t.Run("signed", func(t *testing.T) { manifest := signedManifest(t, files) signer, err := signedChecker(t, manifest). ExtractEmbeddedSigningKeyFP(context.Background()) require.NoError(t, err) assertFetchRefused(t, manifest, files, msgFpB, "embedded signing key fingerprint "+signer+" does not match required "+msgFpB) server := httptest.NewServer(fetchTestHandler(manifest, files)) defer server.Close() dest := t.TempDir() opts := testOpts([]string{ testApp, cmdFetch, "-q", "--" + flagDest, dest, "--" + flagRequireSignature, signer, server.URL, }, afero.NewOsFs()) require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) assert.Equal(t, files[testFileTxt], filesUnder(t, dest)[testFileTxt]) }) } // assertFetchRefused serves manifest, a manifest of files, and fetches it // with --require-signature signer into a directory that does not exist // yet. fetch must fail with message after requesting only the manifest, // and must not create the directory. func assertFetchRefused( t *testing.T, manifest []byte, files map[string][]byte, signer, message string, ) { t.Helper() tree := fetchTestHandler(manifest, 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() dest := filepath.Join(t.TempDir(), "mirror") opts := testOpts([]string{ testApp, cmdFetch, "-q", "--" + flagDest, dest, "--" + flagRequireSignature, signer, server.URL, }, afero.NewOsFs()) assert.Equal(t, 1, runCLI(opts)) assert.Contains(t, testStderr(t, opts), message) assert.NoDirExists(t, dest, "fetch wrote before checking the signer") mu.Lock() defer mu.Unlock() assert.Equal(t, []string{"/" + defaultManifestName}, requested) } // 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") }