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/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") ) // 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, reporting // progress on the progress channel. func downloadManifestFiles( ctx context.Context, client retryingClient, baseURL *url.URL, files []*mfer.MFFilePath, progress chan<- DownloadProgress, ) error { for _, f := range files { // Sanitize the path to prevent path traversal attacks localPath, err := sanitizePath(f.GetPath()) if err != nil { return fmt.Errorf("invalid path in manifest: %w", err) } // 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, localPath, f, progress) if err != nil { return fmt.Errorf("failed to download %s: %w", f.GetPath(), err) } } return 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 } inputURL := ctx.Args().Get(0) manifestURL, err := resolveManifestURL(inputURL) if err != nil { return fmt.Errorf("invalid URL: %w", err) } client := retryingClient{ client: &http.Client{Timeout: timeout}, firstDelay: firstRetryDelay, } 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 fmt.Errorf("failed to fetch manifest: %w", 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) baseURL, err := manifestBaseURL(manifestURL) if err != nil { return err } // Calculate total bytes to download var totalBytes int64 for _, f := range files { totalBytes += f.GetSize() } // 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 dlErr := downloadManifestFiles(ctx.Context, client, baseURL, files, progress) close(progress) <-done if dlErr != nil { return dlErr } // Print summary elapsed := time.Since(startTime) avgBytesPerSec := float64(totalBytes) / elapsed.Seconds() avgRate := formatBitrate(avgBytesPerSec * bitsPerByte) log.Infof("downloaded %d files (%s) in %.1fs (%s avg)", len(files), humanize.IBytes(safeUint64(totalBytes)), elapsed.Seconds(), avgRate) 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 the relative path p // already exists as a symlink. sanitizePath checks p only as text, so // without this a symlink inside the target directory 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(p string) error { current := "" 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 a local file path 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, 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(dir) if err != nil { return err } 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, tmpPath, localPath, entry, progress) }) } // saveResponse writes resp's body to tmpPath, verifies it against entry, // and renames it to localPath. It starts a new temp file each time and // removes it on failure, so a retry after a failed try never appends to // or keeps a partial file. func saveResponse( resp *http.Response, tmpPath, localPath string, entry *mfer.MFFilePath, progress chan<- DownloadProgress, ) error { // Determine expected size expectedSize := entry.GetSize() totalBytes := resp.ContentLength if totalBytes < 0 { totalBytes = expectedSize } err := checkNoSymlinks(tmpPath) if err != nil { 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 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( tmpPath, localPath, written, expectedSize, h.Sum(nil), entry, copyErr, closeErr) if err != nil { _ = os.Remove(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( 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 } err = checkNoSymlinks(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 }