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; a manifest that lists index.mf or its temp name at the top of the tree is refused, since saving would replace that file. --require-signature is shared with check and enforced through verifyRequiredSigner. Both refusals come before any file is downloaded or anything is written. Model: opus-5-5
895 lines
25 KiB
Go
895 lines
25 KiB
Go
package cli
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/sha256"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"math/rand/v2"
|
|
"net"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"path/filepath"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/dustin/go-humanize"
|
|
"github.com/multiformats/go-multihash"
|
|
"github.com/spf13/afero"
|
|
"github.com/urfave/cli/v2"
|
|
"sneak.berlin/go/mfer/internal/log"
|
|
"sneak.berlin/go/mfer/mfer"
|
|
)
|
|
|
|
const (
|
|
// progressChanBuffer is the buffer size of the download progress
|
|
// channel.
|
|
progressChanBuffer = 10
|
|
|
|
// bitsPerByte converts a bytes-per-second rate to bits per second.
|
|
bitsPerByte = 8
|
|
|
|
// dirPerms is the permission mode for directories created for
|
|
// downloaded files. Fetched trees are content that is normally
|
|
// published (served by a web server, read by another uid), so the
|
|
// traversal bit for group and other must stay set.
|
|
dirPerms os.FileMode = 0o755
|
|
|
|
// filePerms is the permission mode, before the umask, for downloaded
|
|
// files. It is the mode os.Create uses; like dirPerms, it keeps group
|
|
// and other read access.
|
|
filePerms os.FileMode = 0o666
|
|
|
|
// Bitrate unit thresholds in bits per second.
|
|
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.
|
|
errAbsolutePath = errors.New("absolute path not allowed")
|
|
// errPathTraversal indicates a manifest path escaping the target
|
|
// directory.
|
|
errPathTraversal = errors.New("path traversal not allowed")
|
|
// errSymlinkInPath indicates a manifest path running through a
|
|
// symlink that already exists in the target directory.
|
|
errSymlinkInPath = errors.New("symlink in path not allowed")
|
|
// errSizeMismatch indicates a downloaded file with an unexpected
|
|
// size.
|
|
errSizeMismatch = errors.New("size mismatch")
|
|
// errHashMismatch indicates a downloaded file whose hash matches no
|
|
// manifest hash.
|
|
errHashMismatch = errors.New("hash mismatch")
|
|
// errManifestNameListed indicates a manifest that lists a file where
|
|
// fetch saves the manifest.
|
|
errManifestNameListed = errors.New(
|
|
"manifest lists a file where fetch saves the manifest")
|
|
)
|
|
|
|
// DownloadProgress reports the progress of a single file download.
|
|
type DownloadProgress struct {
|
|
Path string // File path being downloaded
|
|
BytesRead int64 // Bytes downloaded so far
|
|
TotalBytes int64 // Total expected bytes (-1 if unknown)
|
|
BytesPerSec float64 // Current download rate
|
|
ETA time.Duration // Estimated time to completion
|
|
}
|
|
|
|
// 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.
|
|
//
|
|
// 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 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
|
|
}
|
|
|
|
when, err := http.ParseTime(value)
|
|
if err == nil {
|
|
return time.Until(when), true
|
|
}
|
|
|
|
return 0, false
|
|
}
|
|
|
|
// reportDownloadProgress renders download progress until the channel
|
|
// closes, then closes done.
|
|
func reportDownloadProgress(progress <-chan DownloadProgress, done chan<- struct{}) {
|
|
defer close(done)
|
|
|
|
for p := range progress {
|
|
rate := formatBitrate(p.BytesPerSec * bitsPerByte)
|
|
if p.ETA > 0 {
|
|
log.Infof("%s: %s/%s, %s, ETA %s",
|
|
p.Path, humanize.IBytes(safeUint64(p.BytesRead)),
|
|
humanize.IBytes(safeUint64(p.TotalBytes)),
|
|
rate, p.ETA.Round(time.Second))
|
|
} else {
|
|
log.Infof("%s: %s/%s, %s",
|
|
p.Path, humanize.IBytes(safeUint64(p.BytesRead)),
|
|
humanize.IBytes(safeUint64(p.TotalBytes)), rate)
|
|
}
|
|
}
|
|
}
|
|
|
|
// manifestBaseURL returns the URL of the directory containing the
|
|
// manifest.
|
|
func manifestBaseURL(manifestURL string) (*url.URL, error) {
|
|
parsed, err := url.Parse(manifestURL)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("fetch: invalid manifest URL: %w", err)
|
|
}
|
|
|
|
// JoinPath cleans the path it builds, so ".." drops the manifest's
|
|
// file name.
|
|
return parsed.JoinPath(".."), nil
|
|
}
|
|
|
|
// downloadManifestFiles downloads every file in the manifest into dest,
|
|
// 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(
|
|
ctx context.Context,
|
|
client retryingClient,
|
|
baseURL *url.URL,
|
|
dest string,
|
|
files []*mfer.MFFilePath,
|
|
progress chan<- DownloadProgress,
|
|
) (int, int64, error) {
|
|
var (
|
|
downloaded int
|
|
downloadedBytes int64
|
|
)
|
|
|
|
for _, f := range files {
|
|
// Sanitize the path to prevent path traversal attacks
|
|
localPath, err := sanitizePath(f.GetPath())
|
|
if err != nil {
|
|
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"
|
|
// must be escaped first.
|
|
fileURL := baseURL.JoinPath(encodeFilePath(f.GetPath())).String()
|
|
log.Infof("fetching %s", f.GetPath())
|
|
|
|
err = downloadFile(ctx, client, fileURL, dest, localPath, f, progress)
|
|
if err != nil {
|
|
return 0, 0, fmt.Errorf("failed to download %s: %w", f.GetPath(), err)
|
|
}
|
|
|
|
downloaded++
|
|
downloadedBytes += f.GetSize()
|
|
}
|
|
|
|
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 {
|
|
log.Debug("fetchManifestOperation()")
|
|
|
|
if ctx.Args().Len() == 0 {
|
|
return errURLRequired
|
|
}
|
|
|
|
timeout := ctx.Duration(flagTimeout)
|
|
if timeout <= 0 {
|
|
return errInvalidTimeout
|
|
}
|
|
|
|
manifestURL, err := resolveManifestURL(ctx.Args().Get(0))
|
|
if err != nil {
|
|
return fmt.Errorf("invalid URL: %w", err)
|
|
}
|
|
|
|
client := retryingClient{
|
|
client: &http.Client{Timeout: timeout},
|
|
firstDelay: firstRetryDelay,
|
|
}
|
|
|
|
manifestData, files, err := fetchManifest(ctx, client, manifestURL)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// Compute base URL (directory containing manifest)
|
|
baseURL, err := manifestBaseURL(manifestURL)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
dest := ctx.String(flagDest)
|
|
|
|
err = os.MkdirAll(dest, dirPerms)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to create destination directory %s: %w", dest, err)
|
|
}
|
|
|
|
// Create progress channel and start progress reporter goroutine
|
|
progress := make(chan DownloadProgress, progressChanBuffer)
|
|
done := make(chan struct{})
|
|
|
|
go reportDownloadProgress(progress, done)
|
|
|
|
// Track download start time
|
|
startTime := time.Now()
|
|
|
|
// Download each file
|
|
downloaded, downloadedBytes, dlErr := downloadManifestFiles(
|
|
ctx.Context, client, baseURL, dest, files, progress)
|
|
|
|
close(progress)
|
|
<-done
|
|
|
|
if dlErr != nil {
|
|
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
|
|
elapsed := time.Since(startTime)
|
|
avgBytesPerSec := float64(downloadedBytes) / elapsed.Seconds()
|
|
avgRate := formatBitrate(avgBytesPerSec * bitsPerByte)
|
|
log.Infof("downloaded %d files (%s) in %.1fs (%s avg), skipped %d already present",
|
|
downloaded,
|
|
humanize.IBytes(safeUint64(downloadedBytes)),
|
|
elapsed.Seconds(),
|
|
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 and refusing a manifest
|
|
// that lists a file where it will be saved. 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()
|
|
|
|
err = checkManifestNameUnlisted(files)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
|
|
log.Infof("manifest contains %d files", len(files))
|
|
|
|
return manifestData, files, nil
|
|
}
|
|
|
|
// checkManifestNameUnlisted returns an error if files lists a file or
|
|
// directory at the top of the tree under the name fetch saves the
|
|
// manifest as, or under that name's temp file. Saving the manifest would
|
|
// replace or remove it, or fail once every file was downloaded, leaving a
|
|
// tree check rejects. Names are compared ignoring case, since on a
|
|
// case-insensitive filesystem INDEX.MF and index.mf are one file.
|
|
func checkManifestNameUnlisted(files []*mfer.MFFilePath) error {
|
|
for _, f := range files {
|
|
top, _, _ := strings.Cut(filepath.Clean(f.GetPath()), string(filepath.Separator))
|
|
|
|
if strings.EqualFold(top, defaultManifestName) ||
|
|
strings.EqualFold(top, tempPathFor(defaultManifestName)) {
|
|
return fmt.Errorf("%w: %s", errManifestNameListed, f.GetPath())
|
|
}
|
|
}
|
|
|
|
return 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
|
|
}
|
|
|
|
// encodeFilePath URL-encodes each segment of a file path while preserving slashes.
|
|
func encodeFilePath(p string) string {
|
|
segments := strings.Split(p, "/")
|
|
for i, seg := range segments {
|
|
segments[i] = url.PathEscape(seg)
|
|
}
|
|
|
|
return strings.Join(segments, "/")
|
|
}
|
|
|
|
// sanitizePath validates and sanitizes a file path from the manifest.
|
|
// It prevents path traversal attacks and rejects unsafe paths.
|
|
func sanitizePath(p string) (string, error) {
|
|
// Reject empty paths
|
|
if p == "" {
|
|
return "", errEmptyPath
|
|
}
|
|
|
|
// Reject absolute paths
|
|
if filepath.IsAbs(p) {
|
|
return "", fmt.Errorf("%w: %s", errAbsolutePath, p)
|
|
}
|
|
|
|
// Clean the path to resolve . and ..
|
|
cleaned := filepath.Clean(p)
|
|
|
|
// Reject paths that escape the current directory
|
|
if strings.HasPrefix(cleaned, ".."+string(filepath.Separator)) || cleaned == ".." {
|
|
return "", fmt.Errorf("%w: %s", errPathTraversal, p)
|
|
}
|
|
|
|
// Also check for absolute paths after cleaning (handles edge cases)
|
|
if filepath.IsAbs(cleaned) {
|
|
return "", fmt.Errorf("%w: %s", errAbsolutePath, p)
|
|
}
|
|
|
|
return cleaned, nil
|
|
}
|
|
|
|
// checkNoSymlinks returns an error if any part of p, a path relative to
|
|
// dest, already exists under dest as a symlink. dest itself is the user's
|
|
// choice and may be one. sanitizePath checks p only as text, so without
|
|
// this a symlink inside dest could send a write to p outside of it. Parts
|
|
// that do not exist yet are fine: fetch creates them as plain directories
|
|
// and files. Call it immediately before each write: a symlink created
|
|
// after it returns is not caught.
|
|
func checkNoSymlinks(dest, p string) error {
|
|
current := dest
|
|
|
|
for _, part := range strings.Split(p, string(filepath.Separator)) {
|
|
current = filepath.Join(current, part)
|
|
|
|
info, err := os.Lstat(current)
|
|
if errors.Is(err, os.ErrNotExist) {
|
|
return nil
|
|
}
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("failed to check %s for a symlink: %w", current, err)
|
|
}
|
|
|
|
if info.Mode()&os.ModeSymlink != 0 {
|
|
return fmt.Errorf("%w: %s", errSymlinkInPath, current)
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// resolveManifestURL takes a URL and returns the manifest URL.
|
|
// If the URL already ends with .mf, it's returned as-is.
|
|
// Otherwise, the default manifest name is appended.
|
|
func resolveManifestURL(inputURL string) (string, error) {
|
|
parsed, err := url.Parse(inputURL)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
// Check if URL already ends with .mf
|
|
if strings.HasSuffix(parsed.Path, ".mf") {
|
|
return inputURL, nil
|
|
}
|
|
|
|
return parsed.JoinPath(defaultManifestName).String(), nil
|
|
}
|
|
|
|
// progressWriter wraps an io.Writer and reports progress to a channel.
|
|
type progressWriter struct {
|
|
w io.Writer
|
|
path string
|
|
total int64
|
|
written int64
|
|
startTime time.Time
|
|
progress chan<- DownloadProgress
|
|
}
|
|
|
|
func (pw *progressWriter) Write(p []byte) (int, error) {
|
|
n, err := pw.w.Write(p)
|
|
|
|
pw.written += int64(n)
|
|
if pw.progress != nil {
|
|
var (
|
|
bytesPerSec float64
|
|
eta time.Duration
|
|
)
|
|
|
|
elapsed := time.Since(pw.startTime)
|
|
if elapsed > 0 && pw.written > 0 {
|
|
bytesPerSec = float64(pw.written) / elapsed.Seconds()
|
|
if bytesPerSec > 0 && pw.total > 0 {
|
|
remainingBytes := pw.total - pw.written
|
|
eta = time.Duration(float64(remainingBytes)/bytesPerSec) * time.Second
|
|
}
|
|
}
|
|
|
|
sendProgress(pw.progress, DownloadProgress{
|
|
Path: pw.path,
|
|
BytesRead: pw.written,
|
|
TotalBytes: pw.total,
|
|
BytesPerSec: bytesPerSec,
|
|
ETA: eta,
|
|
})
|
|
}
|
|
|
|
return n, err
|
|
}
|
|
|
|
// formatBitrate formats a bits-per-second value with appropriate unit prefix.
|
|
func formatBitrate(bps float64) string {
|
|
switch {
|
|
case bps >= bpsPerGbps:
|
|
return fmt.Sprintf("%.1f Gbps", bps/bpsPerGbps)
|
|
case bps >= bpsPerMbps:
|
|
return fmt.Sprintf("%.1f Mbps", bps/bpsPerMbps)
|
|
case bps >= bpsPerKbps:
|
|
return fmt.Sprintf("%.1f Kbps", bps/bpsPerKbps)
|
|
default:
|
|
return fmt.Sprintf("%.0f bps", bps)
|
|
}
|
|
}
|
|
|
|
// sendProgress sends a progress update without blocking.
|
|
func sendProgress(ch chan<- DownloadProgress, p DownloadProgress) {
|
|
select {
|
|
case ch <- p:
|
|
default:
|
|
}
|
|
}
|
|
|
|
// tempPathFor computes the temporary download path for a local file.
|
|
// For dotfiles, just append .tmp (they're already hidden); for regular
|
|
// files, prefix with . and append .tmp.
|
|
func tempPathFor(localPath string) string {
|
|
dir := filepath.Dir(localPath)
|
|
base := filepath.Base(localPath)
|
|
|
|
var tmpName string
|
|
if strings.HasPrefix(base, ".") {
|
|
tmpName = base + ".tmp"
|
|
} else {
|
|
tmpName = "." + base + ".tmp"
|
|
}
|
|
|
|
if dir == "" || dir == "." {
|
|
return tmpName
|
|
}
|
|
|
|
return filepath.Join(dir, tmpName)
|
|
}
|
|
|
|
// verifyDownloadedHash checks the computed sha256 digest against the
|
|
// manifest entry's hashes; at least one must match.
|
|
func verifyDownloadedHash(digest []byte, entry *mfer.MFFilePath) error {
|
|
computed, err := multihash.Encode(digest, multihash.SHA2_256)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to encode hash: %w", err)
|
|
}
|
|
|
|
for _, hash := range entry.GetHashes() {
|
|
if bytes.Equal(computed, hash.GetMultiHash()) {
|
|
return nil
|
|
}
|
|
}
|
|
|
|
return errHashMismatch
|
|
}
|
|
|
|
// downloadFile downloads a URL to localPath, a path relative to dest, with
|
|
// hash verification. It downloads to a temporary file, verifies the hash,
|
|
// then renames to the final path. Progress is reported via the progress
|
|
// channel.
|
|
func downloadFile(
|
|
ctx context.Context,
|
|
client retryingClient,
|
|
fileURL, dest, localPath string,
|
|
entry *mfer.MFFilePath,
|
|
progress chan<- DownloadProgress,
|
|
) error {
|
|
// Enforce the path invariant here rather than relying on the caller,
|
|
// so every entry point to downloadFile gets the same treatment.
|
|
localPath, err := sanitizePath(localPath)
|
|
if err != nil {
|
|
return fmt.Errorf("invalid path: %w", err)
|
|
}
|
|
|
|
// Create parent directories if needed
|
|
dir := filepath.Dir(localPath)
|
|
if dir != "" && dir != "." {
|
|
err = checkNoSymlinks(dest, dir)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
dir = filepath.Join(dest, dir)
|
|
|
|
err = os.MkdirAll(dir, dirPerms)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to create directory %s: %w", dir, err)
|
|
}
|
|
}
|
|
|
|
tmpPath := tempPathFor(localPath)
|
|
|
|
return client.get(ctx, fileURL, func(resp *http.Response) error {
|
|
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,
|
|
// and renames it to localPath, both paths relative to dest. 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,
|
|
dest, tmpPath, localPath string,
|
|
entry *mfer.MFFilePath,
|
|
progress chan<- DownloadProgress,
|
|
) error {
|
|
// Determine expected size
|
|
expectedSize := entry.GetSize()
|
|
|
|
totalBytes := resp.ContentLength
|
|
if totalBytes < 0 {
|
|
totalBytes = expectedSize
|
|
}
|
|
|
|
out, err := createTempFile(dest, tmpPath)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// Set up hash computation
|
|
h := sha256.New()
|
|
|
|
// Create progress-reporting writer that also computes hash
|
|
pw := &progressWriter{
|
|
w: io.MultiWriter(out, h),
|
|
path: localPath,
|
|
total: totalBytes,
|
|
startTime: time.Now(),
|
|
progress: progress,
|
|
}
|
|
|
|
// Copy content while hashing and reporting progress
|
|
written, copyErr := io.Copy(pw, resp.Body)
|
|
|
|
// Close file before checking errors (to flush writes)
|
|
closeErr := out.Close()
|
|
|
|
err = finishDownload(
|
|
dest, tmpPath, localPath, written, expectedSize, h.Sum(nil), entry,
|
|
copyErr, closeErr)
|
|
if err != nil {
|
|
_ = os.Remove(filepath.Join(dest, tmpPath))
|
|
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// finishDownload validates the copy result, verifies size and hash, and
|
|
// moves the temp file into place. On error the caller removes tmpPath.
|
|
func finishDownload(
|
|
dest, tmpPath, localPath string,
|
|
written, expectedSize int64,
|
|
digest []byte,
|
|
entry *mfer.MFFilePath,
|
|
copyErr, closeErr error,
|
|
) error {
|
|
if copyErr != nil {
|
|
return copyErr
|
|
}
|
|
|
|
if closeErr != nil {
|
|
return closeErr
|
|
}
|
|
|
|
// Verify size
|
|
if written != expectedSize {
|
|
return fmt.Errorf("%w: expected %d bytes, got %d",
|
|
errSizeMismatch, expectedSize, written)
|
|
}
|
|
|
|
// Verify hash against manifest (at least one must match)
|
|
err := verifyDownloadedHash(digest, entry)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
return moveIntoPlace(dest, tmpPath, localPath)
|
|
}
|