check / check (push) Failing after 2s
fetch downloads each file to a temp name beside it and first removes whatever is there. A manifest listing both a.txt and .a.txt.tmp had fetch delete the second while fetching the first, then exit 0 with a tree check rejects. The refusal of a manifest that lists the saved manifest's own name or temp name now covers this too: a listed file, or a directory a listed file is in, may not sit at any name fetch writes besides the listed files themselves. Temp names come from tempPathFor, names are compared ignoring case as before, and the refusal still happens before the destination is created or any file requested. Model: opus-5-5
913 lines
25 KiB
Go
913 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")
|
|
// errNameClash indicates a manifest that lists a file where fetch
|
|
// writes another file.
|
|
errNameClash = errors.New(
|
|
"manifest lists a file where fetch writes another file")
|
|
)
|
|
|
|
// 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 fetch writes another. 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 = checkNoNameClash(files)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
|
|
log.Infof("manifest contains %d files", len(files))
|
|
|
|
return manifestData, files, nil
|
|
}
|
|
|
|
// checkNoNameClash returns an error if files lists a file, or a directory
|
|
// a file is in, under a name where fetch writes another file: the temp
|
|
// file it downloads a listed file to, or, at the top of the tree, the
|
|
// saved manifest or its temp file. fetch would remove or replace what is
|
|
// listed there, or fail partway, 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 checkNoNameClash(files []*mfer.MFFilePath) error {
|
|
sep := string(filepath.Separator)
|
|
|
|
// written maps each name fetch writes, other than the listed files
|
|
// themselves, in lower case, to the file it writes there.
|
|
written := map[string]string{
|
|
defaultManifestName: "the saved manifest",
|
|
tempPathFor(defaultManifestName): "the saved manifest's temp file",
|
|
}
|
|
|
|
for _, f := range files {
|
|
tmpPath := tempPathFor(filepath.Clean(f.GetPath()))
|
|
written[strings.ToLower(tmpPath)] = "the temp file for " + f.GetPath()
|
|
}
|
|
|
|
for _, f := range files {
|
|
// Look up each directory on the file's path, then the file itself.
|
|
parts := strings.Split(strings.ToLower(filepath.Clean(f.GetPath())), sep)
|
|
|
|
for i := range parts {
|
|
what, ok := written[strings.Join(parts[:i+1], sep)]
|
|
if ok {
|
|
return fmt.Errorf("%w: %s (%s)", errNameClash, f.GetPath(), what)
|
|
}
|
|
}
|
|
}
|
|
|
|
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)
|
|
}
|