1 Commits
Author SHA1 Message Date
sneak 438b73eddf fetch: destination directory, skip files already present, save the manifest, require a signer (closes #101)
check / check (push) Failing after 3s
fetch takes --dest (default .) and writes every file there through the
existing symlink and hard-link guards, which now work relative to that
directory. A file already there with the listed size and hash is
skipped; a leftover temp file is still replaced. Once every file
verifies, the manifest is saved as index.mf through the same temp file
and rename, so check runs on the result. --require-signature is shared
with check and enforced through verifyRequiredSigner, on a Checker over
the manifest held in memory, before anything is downloaded or written.

Model: opus-5-5
2026-10-04 15:20:55 +00:00
6 changed files with 486 additions and 134 deletions
+15 -6
View File
@@ -43,9 +43,11 @@ bin/mfer gen .
# is missing or corrupted. # is missing or corrupted.
bin/mfer check index.mf bin/mfer check index.mf
# Download and cryptographically verify a tree published over HTTP: mfer # Download and cryptographically verify a tree published over HTTP into
# fetches <url>/index.mf, then downloads every file it lists. # ./mirror: mfer fetches <url>/index.mf, downloads every file it lists,
bin/mfer fetch https://example.com/tree/ # skipping any already there with the right hash, then saves the manifest as
# mirror/index.mf.
bin/mfer fetch --dest mirror https://example.com/tree/
``` ```
Run `bin/mfer help` for the full command list, or `bin/mfer <command> --help` Run `bin/mfer help` for the full command list, or `bin/mfer <command> --help`
@@ -252,9 +254,16 @@ are now tracked only in the [issues](https://git.eeqj.de/sneak/mfer/issues).
- verifies checksums of all files in manifest, displaying error and exiting - verifies checksums of all files in manifest, displaying error and exiting
nonzero if any files are missing or corrupted nonzero if any files are missing or corrupted
- `mfer fetch https://example.com/stuff/` - `mfer fetch https://example.com/stuff/`
- fetches `/stuff/index.mf` and downloads all files listed in manifest, - fetches `/stuff/index.mf` and downloads all files listed in manifest into
optionally resuming any that already exist locally, and assures the current directory, or the one given with `--dest`, and assures
cryptographic integrity of downloaded files. cryptographic integrity of downloaded files. A file already there with the
size and hash the manifest lists is skipped. Once every file is in place,
the manifest is saved there as `index.mf`, so `mfer check` can verify the
tree later.
- `mfer fetch --require-signature <fingerprint> https://example.com/stuff/`
- as above, but first refuses a manifest not signed by the key with that
fingerprint, as `mfer check --require-signature` does, before downloading
any file.
# Implementation Plan # Implementation Plan
+1 -1
View File
@@ -317,7 +317,7 @@ func (mfa *CLIApp) checkManifestOperation(ctx *cli.Context) error {
} }
// Check signature requirement // Check signature requirement
requiredSigner := ctx.String("require-signature") requiredSigner := ctx.String(flagRequireSignature)
if requiredSigner != "" { if requiredSigner != "" {
err = verifyRequiredSigner(ctx.Context, chk, requiredSigner) err = verifyRequiredSigner(ctx.Context, chk, requiredSigner)
if err != nil { if err != nil {
+22 -12
View File
@@ -111,9 +111,10 @@ func TestVerifyRequiredSignerMessages(t *testing.T) {
// string; the required signer is a fixed value that cannot match it. Requires // string; the required signer is a fixed value that cannot match it. Requires
// gpg and is skipped where it is absent, as the other signing tests are. // gpg and is skipped where it is absent, as the other signing tests are.
// //
//nolint:paralleltest // signedChecker calls t.Setenv, which bars t.Parallel //nolint:paralleltest // signedManifest calls t.Setenv, which bars t.Parallel
func TestSignerMismatchMessage(t *testing.T) { func TestSignerMismatchMessage(t *testing.T) {
chk := signedChecker(t) chk := signedChecker(t,
signedManifest(t, map[string][]byte{"f.txt": []byte("signed file")}))
embeddedFP, err := chk.ExtractEmbeddedSigningKeyFP(context.Background()) embeddedFP, err := chk.ExtractEmbeddedSigningKeyFP(context.Background())
require.NoError(t, err) require.NoError(t, err)
@@ -125,9 +126,10 @@ func TestSignerMismatchMessage(t *testing.T) {
" does not match required "+msgFpB) " does not match required "+msgFpB)
} }
// signedChecker builds a Checker over a manifest signed by a throwaway GPG // signedManifest returns a manifest of files signed by a throwaway GPG key
// key generated in a temporary GNUPGHOME. // generated in a temporary GNUPGHOME, which it leaves set for the rest of
func signedChecker(t *testing.T) *mfer.Checker { // the test.
func signedManifest(t *testing.T, files map[string][]byte) []byte {
t.Helper() t.Helper()
_, err := exec.LookPath("gpg") _, err := exec.LookPath("gpg")
@@ -159,17 +161,25 @@ func signedChecker(t *testing.T) *mfer.Checker {
b := mfer.NewBuilder() b := mfer.NewBuilder()
b.SetSigningOptions(&mfer.SigningOptions{KeyID: mfer.GPGKeyID("test@mfer.test")}) b.SetSigningOptions(&mfer.SigningOptions{KeyID: mfer.GPGKeyID("test@mfer.test")})
content := []byte("signed file") for path, content := range files {
_, err = b.AddFile("f.txt", mfer.FileSize(len(content)), mfer.ModTime{}, _, err = b.AddFile(mfer.RelFilePath(path), mfer.FileSize(len(content)),
bytes.NewReader(content), nil) mfer.ModTime{}, bytes.NewReader(content), nil)
require.NoError(t, err) require.NoError(t, err)
}
var buf bytes.Buffer var buf bytes.Buffer
require.NoError(t, b.Build(context.Background(), &buf)) require.NoError(t, b.Build(context.Background(), &buf))
return buf.Bytes()
}
// signedChecker builds a Checker over manifest, a signed manifest.
func signedChecker(t *testing.T, manifest []byte) *mfer.Checker {
t.Helper()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
require.NoError(t, afero.WriteFile(fs, "/index.mf", buf.Bytes(), 0o644)) require.NoError(t, afero.WriteFile(fs, "/index.mf", manifest, 0o644))
chk, err := mfer.NewChecker(&mfer.CheckerOptions{ chk, err := mfer.NewChecker(&mfer.CheckerOptions{
ManifestPath: "/index.mf", ManifestPath: "/index.mf",
@@ -300,7 +310,7 @@ func TestFetchFileHTTPStatusMessage(t *testing.T) {
// downloadFile logs each retry of the 500 to the process-global logger. // downloadFile logs each retry of the 500 to the process-global logger.
err := runLocked(func() error { err := runLocked(func() error {
return downloadFile(context.Background(), testClient(), server.URL+"/x", "x", return downloadFile(context.Background(), testClient(), server.URL+"/x", ".", "x",
&mfer.MFFilePath{}, nil) &mfer.MFFilePath{}, nil)
}) })
require.ErrorIs(t, err, errHTTPStatus) require.ErrorIs(t, err, errHTTPStatus)
@@ -355,7 +365,7 @@ func TestSizeMismatchMessage(t *testing.T) {
// finishDownload returns the size-mismatch error before it touches the // finishDownload returns the size-mismatch error before it touches the
// paths, digest, or entry, so those can be zero here. // paths, digest, or entry, so those can be zero here.
err := finishDownload("", "", 9, 10, nil, nil, nil, nil) err := finishDownload("", "", "", 9, 10, nil, nil, nil, nil)
require.ErrorIs(t, err, errSizeMismatch) require.ErrorIs(t, err, errSizeMismatch)
assert.EqualError(t, err, "size mismatch: expected 10 bytes, got 9") assert.EqualError(t, err, "size mismatch: expected 10 bytes, got 9")
} }
+256 -93
View File
@@ -19,6 +19,7 @@ import (
"github.com/dustin/go-humanize" "github.com/dustin/go-humanize"
"github.com/multiformats/go-multihash" "github.com/multiformats/go-multihash"
"github.com/spf13/afero"
"github.com/urfave/cli/v2" "github.com/urfave/cli/v2"
"sneak.berlin/go/mfer/internal/log" "sneak.berlin/go/mfer/internal/log"
"sneak.berlin/go/mfer/mfer" "sneak.berlin/go/mfer/mfer"
@@ -239,20 +240,34 @@ func manifestBaseURL(manifestURL string) (*url.URL, error) {
return parsed.JoinPath(".."), nil return parsed.JoinPath(".."), nil
} }
// downloadManifestFiles downloads every file in the manifest, reporting // downloadManifestFiles downloads every file in the manifest into dest,
// progress on the progress channel. // reporting progress on the progress channel. A file already present in
// dest is skipped. It returns how many files it downloaded and their
// total size.
func downloadManifestFiles( func downloadManifestFiles(
ctx context.Context, ctx context.Context,
client retryingClient, client retryingClient,
baseURL *url.URL, baseURL *url.URL,
dest string,
files []*mfer.MFFilePath, files []*mfer.MFFilePath,
progress chan<- DownloadProgress, progress chan<- DownloadProgress,
) error { ) (int, int64, error) {
var (
downloaded int
downloadedBytes int64
)
for _, f := range files { for _, f := range files {
// Sanitize the path to prevent path traversal attacks // Sanitize the path to prevent path traversal attacks
localPath, err := sanitizePath(f.GetPath()) localPath, err := sanitizePath(f.GetPath())
if err != nil { if err != nil {
return fmt.Errorf("invalid path in manifest: %w", err) return 0, 0, fmt.Errorf("invalid path in manifest: %w", err)
}
if alreadyPresent(dest, localPath, f) {
log.Infof("skipping %s: already present", f.GetPath())
continue
} }
// JoinPath takes escaped path text, so a name such as "100%.txt" // JoinPath takes escaped path text, so a name such as "100%.txt"
@@ -260,13 +275,53 @@ func downloadManifestFiles(
fileURL := baseURL.JoinPath(encodeFilePath(f.GetPath())).String() fileURL := baseURL.JoinPath(encodeFilePath(f.GetPath())).String()
log.Infof("fetching %s", f.GetPath()) log.Infof("fetching %s", f.GetPath())
err = downloadFile(ctx, client, fileURL, localPath, f, progress) err = downloadFile(ctx, client, fileURL, dest, localPath, f, progress)
if err != nil { if err != nil {
return fmt.Errorf("failed to download %s: %w", f.GetPath(), err) return 0, 0, fmt.Errorf("failed to download %s: %w", f.GetPath(), err)
} }
downloaded++
downloadedBytes += f.GetSize()
} }
return nil return downloaded, downloadedBytes, nil
}
// alreadyPresent reports whether localPath under dest is a regular file
// with the size and one of the hashes the manifest lists for entry. It
// hashes the whole file, since a matching size alone would accept a
// corrupted or partly written one. A file it cannot read, or reaches only
// through a symlink, is not present: fetch downloads it, and the download
// reports the problem.
func alreadyPresent(dest, localPath string, entry *mfer.MFFilePath) bool {
if checkNoSymlinks(dest, localPath) != nil {
return false
}
path := filepath.Join(dest, localPath)
info, err := os.Lstat(path)
if err != nil || !info.Mode().IsRegular() || info.Size() != entry.GetSize() {
return false
}
// G304: localPath is a relative path that sanitizePath keeps inside
// dest as text, and checkNoSymlinks just found no symlink in it.
f, err := os.Open(path) //nolint:gosec // G304: see comment above
if err != nil {
return false
}
defer func() { _ = f.Close() }()
h := sha256.New()
_, err = io.Copy(h, f)
if err != nil {
return false
}
return verifyDownloadedHash(h.Sum(nil), entry) == nil
} }
func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error { func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error {
@@ -291,43 +346,22 @@ func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error {
firstDelay: firstRetryDelay, firstDelay: firstRetryDelay,
} }
log.Infof("fetching manifest from %s", manifestURL) manifestData, files, err := fetchManifest(ctx, client, 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 { if err != nil {
return fmt.Errorf("failed to fetch manifest: %w", err) return err
} }
// Parse manifest
manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData))
if err != nil {
return fmt.Errorf("failed to parse manifest: %w", err)
}
files := manifest.Files()
log.Infof("manifest contains %d files", len(files))
// Compute base URL (directory containing manifest) // Compute base URL (directory containing manifest)
baseURL, err := manifestBaseURL(manifestURL) baseURL, err := manifestBaseURL(manifestURL)
if err != nil { if err != nil {
return err return err
} }
// Calculate total bytes to download dest := ctx.String(flagDest)
var totalBytes int64
for _, f := range files { err = os.MkdirAll(dest, dirPerms)
totalBytes += f.GetSize() if err != nil {
return fmt.Errorf("failed to create destination directory %s: %w", dest, err)
} }
// Create progress channel and start progress reporter goroutine // Create progress channel and start progress reporter goroutine
@@ -340,7 +374,8 @@ func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error {
startTime := time.Now() startTime := time.Now()
// Download each file // Download each file
dlErr := downloadManifestFiles(ctx.Context, client, baseURL, files, progress) downloaded, downloadedBytes, dlErr := downloadManifestFiles(
ctx.Context, client, baseURL, dest, files, progress)
close(progress) close(progress)
<-done <-done
@@ -349,15 +384,124 @@ func (mfa *CLIApp) fetchManifestOperation(ctx *cli.Context) error {
return dlErr return dlErr
} }
// Saved only now that every file is in place and verified, so that
// "mfer check" can verify the tree later.
err = saveManifest(dest, manifestData)
if err != nil {
return fmt.Errorf("failed to save manifest: %w", err)
}
// Print summary // Print summary
elapsed := time.Since(startTime) elapsed := time.Since(startTime)
avgBytesPerSec := float64(totalBytes) / elapsed.Seconds() avgBytesPerSec := float64(downloadedBytes) / elapsed.Seconds()
avgRate := formatBitrate(avgBytesPerSec * bitsPerByte) avgRate := formatBitrate(avgBytesPerSec * bitsPerByte)
log.Infof("downloaded %d files (%s) in %.1fs (%s avg)", log.Infof("downloaded %d files (%s) in %.1fs (%s avg), skipped %d already present",
len(files), downloaded,
humanize.IBytes(safeUint64(totalBytes)), humanize.IBytes(safeUint64(downloadedBytes)),
elapsed.Seconds(), elapsed.Seconds(),
avgRate) avgRate,
len(files)-downloaded)
log.Infof("saved manifest to %s", filepath.Join(dest, defaultManifestName))
return nil
}
// fetchManifest downloads the manifest at manifestURL and parses it,
// enforcing --require-signature if it is given. It returns the manifest as
// downloaded, to be saved once the files are in place, and the files it
// lists.
func fetchManifest(
ctx *cli.Context, client retryingClient, manifestURL string,
) ([]byte, []*mfer.MFFilePath, error) {
log.Infof("fetching manifest from %s", 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 nil, nil, fmt.Errorf("failed to fetch manifest: %w", err)
}
// Parse manifest
manifest, err := mfer.NewManifestFromReader(bytes.NewReader(manifestData))
if err != nil {
return nil, nil, fmt.Errorf("failed to parse manifest: %w", err)
}
requiredSigner := ctx.String(flagRequireSignature)
if requiredSigner != "" {
err = verifyFetchedSigner(ctx, manifestData, requiredSigner)
if err != nil {
return nil, nil, err
}
}
files := manifest.Files()
log.Infof("manifest contains %d files", len(files))
return manifestData, files, nil
}
// verifyFetchedSigner enforces --require-signature on the fetched manifest
// exactly as check does. verifyRequiredSigner takes a Checker, which loads
// its manifest from a file, so the manifest is handed to it as a file in
// memory.
func verifyFetchedSigner(
ctx *cli.Context, manifestData []byte, requiredSigner string,
) error {
memFs := afero.NewMemMapFs()
manifestPath := "/" + defaultManifestName
err := afero.WriteFile(memFs, manifestPath, manifestData, filePerms)
if err != nil {
return err
}
chk, err := mfer.NewChecker(&mfer.CheckerOptions{
ManifestPath: manifestPath,
BasePath: "/",
Fs: memFs,
})
if err != nil {
return fmt.Errorf("failed to load manifest: %w", err)
}
return verifyRequiredSigner(ctx.Context, chk, requiredSigner)
}
// saveManifest writes the fetched manifest into dest under the default
// manifest name, the way fetch writes every file: to a new temp file that
// is then renamed into place.
func saveManifest(dest string, manifestData []byte) error {
tmpPath := tempPathFor(defaultManifestName)
out, err := createTempFile(dest, tmpPath)
if err != nil {
return err
}
_, writeErr := out.Write(manifestData)
closeErr := out.Close()
err = errors.Join(writeErr, closeErr)
if err == nil {
err = moveIntoPlace(dest, tmpPath, defaultManifestName)
}
if err != nil {
_ = os.Remove(filepath.Join(dest, tmpPath))
return err
}
return nil return nil
} }
@@ -401,14 +545,15 @@ func sanitizePath(p string) (string, error) {
return cleaned, nil return cleaned, nil
} }
// checkNoSymlinks returns an error if any part of the relative path p // checkNoSymlinks returns an error if any part of p, a path relative to
// already exists as a symlink. sanitizePath checks p only as text, so // dest, already exists under dest as a symlink. dest itself is the user's
// without this a symlink inside the target directory could send a write // choice and may be one. sanitizePath checks p only as text, so without
// to p outside of it. Parts that do not exist yet are fine: fetch creates // this a symlink inside dest could send a write to p outside of it. Parts
// them as plain directories and files. Call it immediately before each // that do not exist yet are fine: fetch creates them as plain directories
// write: a symlink created after it returns is not caught. // and files. Call it immediately before each write: a symlink created
func checkNoSymlinks(p string) error { // after it returns is not caught.
current := "" func checkNoSymlinks(dest, p string) error {
current := dest
for _, part := range strings.Split(p, string(filepath.Separator)) { for _, part := range strings.Split(p, string(filepath.Separator)) {
current = filepath.Join(current, part) current = filepath.Join(current, part)
@@ -548,13 +693,14 @@ func verifyDownloadedHash(digest []byte, entry *mfer.MFFilePath) error {
return errHashMismatch return errHashMismatch
} }
// downloadFile downloads a URL to a local file path with hash verification. // downloadFile downloads a URL to localPath, a path relative to dest, with
// It downloads to a temporary file, verifies the hash, then renames to the final path. // hash verification. It downloads to a temporary file, verifies the hash,
// Progress is reported via the progress channel. // then renames to the final path. Progress is reported via the progress
// channel.
func downloadFile( func downloadFile(
ctx context.Context, ctx context.Context,
client retryingClient, client retryingClient,
fileURL, localPath string, fileURL, dest, localPath string,
entry *mfer.MFFilePath, entry *mfer.MFFilePath,
progress chan<- DownloadProgress, progress chan<- DownloadProgress,
) error { ) error {
@@ -568,11 +714,13 @@ func downloadFile(
// Create parent directories if needed // Create parent directories if needed
dir := filepath.Dir(localPath) dir := filepath.Dir(localPath)
if dir != "" && dir != "." { if dir != "" && dir != "." {
err = checkNoSymlinks(dir) err = checkNoSymlinks(dest, dir)
if err != nil { if err != nil {
return err return err
} }
dir = filepath.Join(dest, dir)
err = os.MkdirAll(dir, dirPerms) err = os.MkdirAll(dir, dirPerms)
if err != nil { if err != nil {
return fmt.Errorf("failed to create directory %s: %w", dir, err) return fmt.Errorf("failed to create directory %s: %w", dir, err)
@@ -582,17 +730,61 @@ func downloadFile(
tmpPath := tempPathFor(localPath) tmpPath := tempPathFor(localPath)
return client.get(ctx, fileURL, func(resp *http.Response) error { return client.get(ctx, fileURL, func(resp *http.Response) error {
return saveResponse(resp, tmpPath, localPath, entry, progress) return saveResponse(resp, dest, tmpPath, localPath, entry, progress)
}) })
} }
// createTempFile creates tmpPath, a path relative to dest, as a new empty
// file.
func createTempFile(dest, tmpPath string) (*os.File, error) {
err := checkNoSymlinks(dest, tmpPath)
if err != nil {
return nil, err
}
path := filepath.Join(dest, tmpPath)
// Remove whatever is at tmpPath, such as a leftover from an
// interrupted run, rather than write into it: it may be a hard link
// to a file outside dest, and removing a hard link removes only this
// name. If the removal fails, O_EXCL below makes the create fail.
_ = os.Remove(path)
// Create the temp file only if nothing is at tmpPath (O_EXCL).
//
// G304: tmpPath is a relative path that sanitizePath keeps inside dest
// as text, and checkNoSymlinks just found no symlink in it.
out, err := os.OpenFile( //nolint:gosec // G304: see comment above
path, os.O_RDWR|os.O_CREATE|os.O_EXCL, filePerms)
if err != nil {
return nil, fmt.Errorf("failed to create temp file: %w", err)
}
return out, nil
}
// moveIntoPlace renames tmpPath to localPath, both relative to dest.
func moveIntoPlace(dest, tmpPath, localPath string) error {
err := checkNoSymlinks(dest, localPath)
if err != nil {
return err
}
err = os.Rename(filepath.Join(dest, tmpPath), filepath.Join(dest, localPath))
if err != nil {
return fmt.Errorf("failed to rename temp file: %w", err)
}
return nil
}
// saveResponse writes resp's body to tmpPath, verifies it against entry, // 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 // and renames it to localPath, both paths relative to dest. It starts a
// removes it on failure, so a retry after a failed try never appends to // new temp file each time and removes it on failure, so a retry after a
// or keeps a partial file. // failed try never appends to or keeps a partial file.
func saveResponse( func saveResponse(
resp *http.Response, resp *http.Response,
tmpPath, localPath string, dest, tmpPath, localPath string,
entry *mfer.MFFilePath, entry *mfer.MFFilePath,
progress chan<- DownloadProgress, progress chan<- DownloadProgress,
) error { ) error {
@@ -604,29 +796,11 @@ func saveResponse(
totalBytes = expectedSize totalBytes = expectedSize
} }
err := checkNoSymlinks(tmpPath) out, err := createTempFile(dest, tmpPath)
if err != nil { if err != nil {
return err return err
} }
// Remove whatever is at tmpPath, such as a leftover from an
// interrupted run, rather than write into it: it may be a hard link
// to a file outside the target directory, and removing a hard link
// removes only this name. If the removal fails, O_EXCL below makes
// the create fail.
_ = os.Remove(tmpPath)
// Create the temp file only if nothing is at tmpPath (O_EXCL).
//
// G304: tmpPath is a relative path that sanitizePath keeps inside the
// target directory as text, and checkNoSymlinks just found no symlink
// in it.
out, err := os.OpenFile( //nolint:gosec // G304: see comment above
tmpPath, os.O_RDWR|os.O_CREATE|os.O_EXCL, filePerms)
if err != nil {
return fmt.Errorf("failed to create temp file: %w", err)
}
// Set up hash computation // Set up hash computation
h := sha256.New() h := sha256.New()
@@ -646,10 +820,10 @@ func saveResponse(
closeErr := out.Close() closeErr := out.Close()
err = finishDownload( err = finishDownload(
tmpPath, localPath, written, expectedSize, h.Sum(nil), entry, dest, tmpPath, localPath, written, expectedSize, h.Sum(nil), entry,
copyErr, closeErr) copyErr, closeErr)
if err != nil { if err != nil {
_ = os.Remove(tmpPath) _ = os.Remove(filepath.Join(dest, tmpPath))
return err return err
} }
@@ -660,7 +834,7 @@ func saveResponse(
// finishDownload validates the copy result, verifies size and hash, and // finishDownload validates the copy result, verifies size and hash, and
// moves the temp file into place. On error the caller removes tmpPath. // moves the temp file into place. On error the caller removes tmpPath.
func finishDownload( func finishDownload(
tmpPath, localPath string, dest, tmpPath, localPath string,
written, expectedSize int64, written, expectedSize int64,
digest []byte, digest []byte,
entry *mfer.MFFilePath, entry *mfer.MFFilePath,
@@ -686,16 +860,5 @@ func finishDownload(
return err return err
} }
err = checkNoSymlinks(localPath) return moveIntoPlace(dest, tmpPath, localPath)
if err != nil {
return err
}
// Rename temp file to final path
err = os.Rename(tmpPath, localPath)
if err != nil {
return fmt.Errorf("failed to rename temp file: %w", err)
}
return nil
} }
+169 -14
View File
@@ -338,7 +338,7 @@ func TestFetchFromHTTP(t *testing.T) {
fileURL := baseURL + f.GetPath() fileURL := baseURL + f.GetPath()
err = downloadFile(context.Background(), testClient(), err = downloadFile(context.Background(), testClient(),
fileURL, localPath, f, progress) fileURL, ".", localPath, f, progress)
require.NoError(t, err, "failed to download %s", f.GetPath()) require.NoError(t, err, "failed to download %s", f.GetPath())
} }
@@ -386,7 +386,7 @@ func TestFetchHashMismatch(t *testing.T) {
// Try to download - should fail with hash mismatch // Try to download - should fail with hash mismatch
err = downloadFile(context.Background(), testClient(), err = downloadFile(context.Background(), testClient(),
server.URL+"/file.txt", testFileTxt, files[0], nil) server.URL+"/file.txt", ".", testFileTxt, files[0], nil)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "mismatch") assert.Contains(t, err.Error(), "mismatch")
@@ -432,7 +432,7 @@ func TestFetchSizeMismatch(t *testing.T) {
// Try to download - should fail with size mismatch // Try to download - should fail with size mismatch
err = downloadFile(context.Background(), testClient(), err = downloadFile(context.Background(), testClient(),
server.URL+"/file.txt", testFileTxt, files[0], nil) server.URL+"/file.txt", ".", testFileTxt, files[0], nil)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "size mismatch") assert.Contains(t, err.Error(), "size mismatch")
@@ -490,7 +490,7 @@ func TestFetchProgress(t *testing.T) {
// Download // Download
err = downloadFile(context.Background(), testClient(), err = downloadFile(context.Background(), testClient(),
server.URL+"/large.txt", "large.txt", files[0], progress) server.URL+"/large.txt", ".", "large.txt", files[0], progress)
close(progress) close(progress)
<-done <-done
@@ -787,7 +787,7 @@ func TestDownloadFileRetriesToSuccess(t *testing.T) {
chdirTemp(t) chdirTemp(t)
err = downloadFile(context.Background(), testClient(), err = downloadFile(context.Background(), testClient(),
server.URL+"/"+testFileTxt, testFileTxt, manifest.Files()[0], nil) server.URL+"/"+testFileTxt, ".", testFileTxt, manifest.Files()[0], nil)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, int32(3), requests.Load()) assert.Equal(t, int32(3), requests.Load())
@@ -878,7 +878,7 @@ func TestFetchEscapedPaths(t *testing.T) {
// TestFetchTree runs fetch on a tree with nested directories. Every file // TestFetchTree runs fetch on a tree with nested directories. Every file
// the manifest lists must land under its own path with its own content, // the manifest lists must land under its own path with its own content,
// and nothing else may be left in the destination. // beside the manifest, and nothing else may be left in the destination.
// //
//nolint:paralleltest // changes the process-global working directory //nolint:paralleltest // changes the process-global working directory
func TestFetchTree(t *testing.T) { func TestFetchTree(t *testing.T) {
@@ -889,7 +889,9 @@ func TestFetchTree(t *testing.T) {
"other/deep/est.txt": []byte("in a second directory"), "other/deep/est.txt": []byte("in a second directory"),
} }
server := httptest.NewServer(fetchTestHandler(manifestOf(t, files), files)) manifest := manifestOf(t, files)
server := httptest.NewServer(fetchTestHandler(manifest, files))
defer server.Close() defer server.Close()
dest := chdirTemp(t) dest := chdirTemp(t)
@@ -897,7 +899,9 @@ func TestFetchTree(t *testing.T) {
opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs())
require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) require.Equal(t, 0, runCLI(opts), testStderr(t, opts))
assert.Equal(t, files, filesUnder(t, dest)) want := maps.Clone(files)
want[defaultManifestName] = manifest
assert.Equal(t, want, filesUnder(t, dest))
} }
// TestFetchFailsOnHashMismatch runs fetch against a server that serves a // TestFetchFailsOnHashMismatch runs fetch against a server that serves a
@@ -922,9 +926,10 @@ func TestFetchFailsOnHashMismatch(t *testing.T) {
// TestFetchIntoPartlyFilledDestination runs fetch where an interrupted // TestFetchIntoPartlyFilledDestination runs fetch where an interrupted
// fetch of an older version of the tree left one file current, one file // 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 // out of date and one half written to its temp file, beside a file the
// manifest does not list. fetch downloads every file the manifest lists, // manifest does not list. fetch skips the current file and downloads the
// those already present included, and replaces what is there; the file // other two. The out-of-date file has the same size as the new version,
// the manifest does not list is left alone. // 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 //nolint:paralleltest // changes the process-global working directory
func TestFetchIntoPartlyFilledDestination(t *testing.T) { func TestFetchIntoPartlyFilledDestination(t *testing.T) {
@@ -935,7 +940,8 @@ func TestFetchIntoPartlyFilledDestination(t *testing.T) {
} }
unlisted := []byte("not in the manifest") unlisted := []byte("not in the manifest")
tree := fetchTestHandler(manifestOf(t, files), files) manifest := manifestOf(t, files)
tree := fetchTestHandler(manifest, files)
var ( var (
mu sync.Mutex mu sync.Mutex
@@ -963,21 +969,170 @@ func TestFetchIntoPartlyFilledDestination(t *testing.T) {
os.WriteFile(tempPathFor("sub/partial.txt"), []byte("cut off"), 0o600)) os.WriteFile(tempPathFor("sub/partial.txt"), []byte("cut off"), 0o600))
require.NoError(t, os.WriteFile("unlisted.txt", unlisted, 0o600)) require.NoError(t, os.WriteFile("unlisted.txt", unlisted, 0o600))
opts := testOpts([]string{testApp, cmdFetch, "-q", server.URL}, afero.NewOsFs()) opts := testOpts([]string{testApp, cmdFetch, server.URL}, afero.NewOsFs())
require.Equal(t, 0, runCLI(opts), testStderr(t, opts)) 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 := maps.Clone(files)
want["unlisted.txt"] = unlisted want["unlisted.txt"] = unlisted
want[defaultManifestName] = manifest
assert.Equal(t, want, filesUnder(t, dest)) assert.Equal(t, want, filesUnder(t, dest))
mu.Lock() mu.Lock()
defer mu.Unlock() defer mu.Unlock()
assert.ElementsMatch(t, []string{ assert.ElementsMatch(t, []string{
"/" + defaultManifestName, "/current.txt", "/sub/changed.txt", "/sub/partial.txt", "/" + defaultManifestName, "/sub/changed.txt", "/sub/partial.txt",
}, requested) }, 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 // TestFetchTimeoutFlag runs fetch with --timeout against a server that
// never answers. Without the flag's limit the request would wait forever; // 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 // once fetch gives up on it, the server cancels fetch's context so that
+23 -8
View File
@@ -22,8 +22,10 @@ const (
cmdFetch = "fetch" cmdFetch = "fetch"
cmdVersion = "version" cmdVersion = "version"
flagProgress = "progress" flagProgress = "progress"
flagTimeout = "timeout" flagTimeout = "timeout"
flagDest = "dest"
flagRequireSignature = "require-signature"
manifestArgsUsage = "[manifest file]" manifestArgsUsage = "[manifest file]"
@@ -142,6 +144,17 @@ func commonFlags() []cli.Flag {
} }
} }
// requireSignatureFlag returns the --require-signature flag taken by the
// check and fetch subcommands.
func requireSignatureFlag() *cli.StringFlag {
return &cli.StringFlag{
Name: flagRequireSignature,
Aliases: []string{"S"},
Usage: "Require manifest to be signed by the specified GPG key ID",
EnvVars: []string{"MFER_REQUIRE_SIGNATURE"},
}
}
func (mfa *CLIApp) generateCommand() *cli.Command { func (mfa *CLIApp) generateCommand() *cli.Command {
return &cli.Command{ return &cli.Command{
Name: cmdGenerate, Name: cmdGenerate,
@@ -229,12 +242,7 @@ func (mfa *CLIApp) checkCommand() *cli.Command {
Name: "no-extra-files", Name: "no-extra-files",
Usage: "Fail if files exist in base directory that are not in manifest", Usage: "Fail if files exist in base directory that are not in manifest",
}, },
&cli.StringFlag{ requireSignatureFlag(),
Name: "require-signature",
Aliases: []string{"S"},
Usage: "Require manifest to be signed by the specified GPG key ID",
EnvVars: []string{"MFER_REQUIRE_SIGNATURE"},
},
), ),
} }
} }
@@ -354,6 +362,13 @@ func (mfa *CLIApp) fetchCommand() *cli.Command {
Usage: "Time limit for each HTTP request, including the download " + Usage: "Time limit for each HTTP request, including the download " +
"of its body", "of its body",
}, },
&cli.StringFlag{
Name: flagDest,
Aliases: []string{"d"},
Value: ".",
Usage: "Directory to download the files and the manifest into",
},
requireSignatureFlag(),
), ),
} }
} }