diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 97fe26d..b890554 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -23,7 +23,7 @@ jobs: with: go-version: "1.25.13" cache-dependency-path: go.sum - - run: go install golang.org/x/vuln/cmd/govulncheck@latest + - run: go install golang.org/x/vuln/cmd/govulncheck@v1.7.0 - run: govulncheck ./... - run: go vet ./... - run: go test ./... @@ -74,12 +74,32 @@ jobs: with: go-version: "1.25.13" cache-dependency-path: go.sum - - run: go install golang.org/x/vuln/cmd/govulncheck@latest + - run: go install golang.org/x/vuln/cmd/govulncheck@v1.7.0 - run: govulncheck ./... - run: ./build.sh env: VERSION: ci - run: ./packaging/package-release.sh ci + - name: Upload Linux binaries + uses: actions/upload-artifact@v7 + with: + name: dbxcli-linux-binaries + path: | + dist/dbxcli-linux-amd64 + dist/dbxcli-linux-arm64 + dist/dbxcli-linux-arm + if-no-files-found: error + retention-days: 30 + - name: Upload release archives + uses: actions/upload-artifact@v7 + with: + name: dbxcli-release-archives + path: | + dist/dbxcli_ci_*.tar.gz + dist/dbxcli_ci_*.zip + dist/SHA256SUMS + if-no-files-found: error + retention-days: 30 docs: name: Generated docs diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 8db4d13..c5e5172 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -29,7 +29,7 @@ jobs: asset_version="${GITHUB_REF_NAME#v}" echo "asset_version=${asset_version}" >> "$GITHUB_OUTPUT" echo "windows_archive=dbxcli_${asset_version}_windows_amd64.zip" >> "$GITHUB_OUTPUT" - - run: go install golang.org/x/vuln/cmd/govulncheck@latest + - run: go install golang.org/x/vuln/cmd/govulncheck@v1.7.0 - run: govulncheck ./... - uses: golangci/golangci-lint-action@v9 with: diff --git a/CHANGELOG.md b/CHANGELOG.md index c2e861f..97b8a19 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,19 @@ [Full Changelog](https://github.com/dropbox/dbxcli/compare/v3.7.3...HEAD) +**Added:** + +- `get --recursive` now downloads multiple files in parallel. The number of concurrent downloads is tuned automatically from the measured throughput, and can be pinned with the new `--workers`/`-w` flag. + +**Changed:** + +- Download progress for `get` and `share-link download` now shows a progress bar, current throughput (MiB/s), elapsed time, and estimated time remaining, and parallel downloads end with a summary of the average throughput. + +**Infrastructure:** + +- The CI release build now uploads the Linux binaries and packaged release archives as workflow artifacts. +- Pinned `govulncheck` to v1.7.0 in CI and release workflows, because newer releases require Go 1.26. + ## [v3.7.3](https://github.com/dropbox/dbxcli/tree/v3.7.3) (2026-08-18) [Full Changelog](https://github.com/dropbox/dbxcli/compare/v3.7.2...v3.7.3) diff --git a/README.md b/README.md index a8de6b5..51d6b0c 100644 --- a/README.md +++ b/README.md @@ -87,6 +87,15 @@ Download to stdout: dbxcli get /Backups/project.tgz - | tar tz ``` +Download a folder. Files are fetched in parallel, and the number of +concurrent downloads is tuned automatically from the measured throughput. +Use `--workers` (`-w`) to pin it instead: + +```sh +dbxcli get -r /Photos/2026 ./photos +dbxcli get -r -w 8 /Photos/2026 ./photos +``` + Create a shared link: ```sh diff --git a/cmd/get.go b/cmd/get.go index 4d3323b..40608fd 100644 --- a/cmd/get.go +++ b/cmd/get.go @@ -18,6 +18,7 @@ import ( "errors" "fmt" "io" + "math" "os" "path" "path/filepath" @@ -27,8 +28,6 @@ import ( "github.com/dropbox/dbxcli/v3/internal/output" "github.com/dropbox/dropbox-sdk-go-unofficial/v6/dropbox/files" "github.com/dropbox/dropbox-sdk-go-unofficial/v6/dropbox/filetransfer" - "github.com/dustin/go-humanize" - "github.com/mitchellh/ioprogress" "github.com/spf13/cobra" ) @@ -43,6 +42,9 @@ const ( type getOptions struct { errOut io.Writer + // workers is the number of files downloaded concurrently by recursive + // downloads; downloadWorkersAuto tunes it from measured throughput. + workers int } type getCommandInput struct { @@ -83,7 +85,10 @@ func get(cmd *cobra.Command, args []string) (err error) { } recursive, _ := cmd.Flags().GetBool("recursive") - opts := parseGetOptions(cmd) + opts, err := parseGetOptions(cmd) + if err != nil { + return err + } if dst == "-" { if commandOutputFormat(cmd) == output.FormatJSON { @@ -138,7 +143,7 @@ func get(cmd *cobra.Command, args []string) (err error) { dst = filepath.Join(dst, sourceName) } if commandOutputFormat(cmd) == output.FormatText { - return withJSONErrorDetails(getRecursiveWithRootMetadata(dbx, src, dst, meta), operationErrorDetails("download"), pathErrorDetails(src), relocationErrorDetails(src, dst)) + return withJSONErrorDetails(getRecursiveWithRootMetadata(dbx, src, dst, meta, opts), operationErrorDetails("download"), pathErrorDetails(src), relocationErrorDetails(src, dst)) } results, err := getRecursiveWithResults(dbx, src, dst, meta, opts) if err != nil { @@ -174,10 +179,22 @@ func get(cmd *cobra.Command, args []string) (err error) { }, []getResult{result}) } -func parseGetOptions(cmd *cobra.Command) getOptions { - return getOptions{ - errOut: cmd.ErrOrStderr(), +func parseGetOptions(cmd *cobra.Command) (getOptions, error) { + opts := getOptions{ + errOut: cmd.ErrOrStderr(), + workers: downloadWorkersAuto, + } + if cmd.Flags().Lookup("workers") != nil { + workers, err := cmd.Flags().GetInt("workers") + if err != nil { + return getOptions{}, err + } + if workers < 0 { + return getOptions{}, invalidArgumentsErrorWithDetails("`--workers` must be greater than or equal to 0 (0 selects automatic tuning)", flagErrorDetails("workers")) + } + opts.workers = workers } + return opts, nil } func getErrorOutput(opts getOptions) io.Writer { @@ -245,8 +262,8 @@ func getRecursive(dbx filesClient, src, dst string) error { return err } -func getRecursiveWithRootMetadata(dbx filesClient, src, dst string, rootMeta files.IsMetadata) error { - _, err := getRecursiveInternal(dbx, src, dst, rootMeta, getOptions{}, false) +func getRecursiveWithRootMetadata(dbx filesClient, src, dst string, rootMeta files.IsMetadata, opts getOptions) error { + _, err := getRecursiveInternal(dbx, src, dst, rootMeta, opts, false) return err } @@ -292,14 +309,22 @@ func getRecursiveInternal(dbx filesClient, src, dst string, rootMeta files.IsMet } } - var downloadErrors []error + // Folders are created up front and in listing order; files are queued as + // jobs and downloaded concurrently. Results and errors are kept in listing + // order so output is stable regardless of completion order. + entryResults := make([]*getResult, len(entries)) + entryErrors := make([]error, len(entries)) + var jobIndexes []int + var jobFiles []*files.FileMetadata + var jobTargets []string + var totalBytes int64 - for _, entry := range entries { + for i, entry := range entries { switch f := entry.(type) { case *files.FolderMetadata: relPath, err := relativeTo(rootPath, f.PathDisplay) if err != nil { - downloadErrors = append(downloadErrors, err) + entryErrors[i] = err continue } if relPath == "" { @@ -309,51 +334,98 @@ func getRecursiveInternal(dbx filesClient, src, dst string, rootMeta files.IsMet if collectResults { result, err := ensureLocalDirectoryResult(f.PathDisplay, localDir, f) if err != nil { - downloadErrors = append(downloadErrors, fmt.Errorf("mkdir %s: %w", localDir, err)) + entryErrors[i] = fmt.Errorf("mkdir %s: %w", localDir, err) continue } - results = append(results, result) + entryResults[i] = &result } else { if err := os.MkdirAll(localDir, 0755); err != nil { - downloadErrors = append(downloadErrors, fmt.Errorf("mkdir %s: %w", localDir, err)) + entryErrors[i] = fmt.Errorf("mkdir %s: %w", localDir, err) } } case *files.FileMetadata: relPath, err := relativeTo(rootPath, f.PathDisplay) if err != nil { - downloadErrors = append(downloadErrors, err) + entryErrors[i] = err continue } - localPath := filepath.Join(dst, filepath.FromSlash(relPath)) + jobIndexes = append(jobIndexes, i) + jobFiles = append(jobFiles, f) + jobTargets = append(jobTargets, filepath.Join(dst, filepath.FromSlash(relPath))) + if f.Size <= math.MaxInt64 { + totalBytes += int64(f.Size) + } + } + } + + errOut := getErrorOutput(opts) + status := newDownloadStatusWriter(errOut) + pool := newDownloadPool(opts.workers, len(jobFiles), totalBytes, status) + concurrent := pool.concurrent() + + jobs := make([]func(), len(jobFiles)) + for j := range jobFiles { + i, f, localPath := jobIndexes[j], jobFiles[j], jobTargets[j] + jobs[j] = func() { if err := os.MkdirAll(filepath.Dir(localPath), 0755); err != nil { - downloadErrors = append(downloadErrors, fmt.Errorf("mkdir %s: %w", filepath.Dir(localPath), err)) - continue + entryErrors[i] = fmt.Errorf("mkdir %s: %w", filepath.Dir(localPath), err) + return + } + + // With several downloads in flight the per-file progress bars + // would overwrite each other, so concurrent mode reports only + // one line per file plus the pool's aggregate status line. + fileErrOut := errOut + var progress downloadProgressFunc + if concurrent { + status.message("Downloading %s -> %s\n", f.PathDisplay, localPath) + fileErrOut = io.Discard + if !isExportOnlyFile(f) { + progress = pool.fileProgress() + } + } else { + fmt.Fprintf(errOut, "Downloading %s -> %s\n", f.PathDisplay, localPath) + } + + metadata, actualDst, err := downloadFileWithProgress(dbx, f.PathDisplay, localPath, f, false, fileErrOut, progress) + if err != nil { + entryErrors[i] = fmt.Errorf("%s: %w", f.PathDisplay, err) + return + } + // Exported files report no streaming progress, so count them once + // they have been written in full. + if concurrent && progress == nil && metadata != nil && metadata.Size <= math.MaxInt64 { + pool.addBytes(int64(metadata.Size)) } - fmt.Fprintf(getErrorOutput(opts), "Downloading %s -> %s\n", f.PathDisplay, localPath) if collectResults { - result, err := downloadFileWithResult(dbx, f.PathDisplay, localPath, f, false, opts) + result, err := newGetResult(getStatusDownloaded, getKindFile, f.PathDisplay, actualDst, metadata) if err != nil { - downloadErrors = append( - downloadErrors, - fmt.Errorf("%s: %w", f.PathDisplay, err), - ) - continue + entryErrors[i] = fmt.Errorf("%s: %w", f.PathDisplay, err) + return } - results = append(results, result) - continue - } - if _, _, err := downloadFileWithMetadata(dbx, f.PathDisplay, localPath, f, false, getErrorOutput(opts)); err != nil { - downloadErrors = append( - downloadErrors, - fmt.Errorf("%s: %w", f.PathDisplay, err), - ) + entryResults[i] = &result } } } + pool.run(currentContext(), jobs, func(j int, err error) { + entryErrors[jobIndexes[j]] = fmt.Errorf("%s: %w", jobFiles[j].PathDisplay, err) + }) + + var downloadErrors []error + for i := range entries { + if entryErrors[i] != nil { + downloadErrors = append(downloadErrors, entryErrors[i]) + continue + } + if entryResults[i] != nil { + results = append(results, *entryResults[i]) + } + } + if len(downloadErrors) > 0 { for _, e := range downloadErrors { - fmt.Fprintf(getErrorOutput(opts), "Error: %v\n", e) + fmt.Fprintf(errOut, "Error: %v\n", e) } return nil, commandFailedErrorfWithDetails("get: %d error(s)", mergeJSONErrorDetails(operationErrorDetails("download"), pathErrorDetails(src), relocationErrorDetails(src, dst)), len(downloadErrors)) } @@ -416,9 +488,25 @@ func downloadFileWithMetadata( metadata *files.FileMetadata, dstExplicit bool, errOut io.Writer, +) (*files.FileMetadata, string, error) { + return downloadFileWithProgress(dbx, src, dst, metadata, dstExplicit, errOut, nil) +} + +// downloadProgressFunc receives the monotonic number of bytes committed so +// far for one file and the file's total size. +type downloadProgressFunc func(committed, total int64) + +func downloadFileWithProgress( + dbx filesClient, + src string, + dst string, + metadata *files.FileMetadata, + dstExplicit bool, + errOut io.Writer, + progress downloadProgressFunc, ) (*files.FileMetadata, string, error) { if !isExportOnlyFile(metadata) { - result, err := downloadFileOnce(dbx, src, dst, errOut) + result, err := downloadFileOnce(dbx, src, dst, errOut, progress) return result, dst, err } @@ -472,7 +560,7 @@ func downloadDestinationPath(dst string) (string, error) { return "", fmt.Errorf("too many symlinks resolving %s", dst) } -func downloadFileOnce(dbx filesClient, src string, dst string, errOut io.Writer) (*files.FileMetadata, error) { +func downloadFileOnce(dbx filesClient, src string, dst string, errOut io.Writer, onProgress downloadProgressFunc) (*files.FileMetadata, error) { finalDst, err := downloadDestinationPath(dst) if err != nil { return nil, err @@ -481,11 +569,8 @@ func downloadFileOnce(dbx filesClient, src string, dst string, errOut io.Writer) errOut = io.Discard } - draw := ioprogress.DrawTerminalf(errOut, func(progress, total int64) string { - return fmt.Sprintf("Downloading %s/%s", - humanize.IBytes(uint64(progress)), humanize.IBytes(uint64(total))) - }) - defer func() { _ = draw(-1, -1) }() + drawer := newTransferProgressDrawer(errOut, "Downloading ") + defer drawer.finish() result, err := filetransfer.NewDownloader(dbx).Download( currentContext(), @@ -494,7 +579,10 @@ func downloadFileOnce(dbx filesClient, src string, dst string, errOut io.Writer) filetransfer.DownloadOptions{ MaxAttempts: maxRetries + 1, Progress: func(progress filetransfer.DownloadProgress) { - _ = draw(progress.BytesCommitted, progress.TotalBytes) + drawer.update(progress.BytesCommitted, progress.TotalBytes) + if onProgress != nil { + onProgress(progress.BytesCommitted, progress.TotalBytes) + } }, }, ) @@ -574,12 +662,16 @@ var getCmd = &cobra.Command{ - Source may be a Dropbox path, file ID (id:), revision (rev:), or namespace-relative path (ns:). - Use --recursive (-r) to download entire directories. + - Recursive downloads fetch several files in parallel. By default the + number of concurrent downloads is tuned automatically from the measured + throughput; use --workers (-w) to set a fixed number instead. - Use - as target to write file bytes to stdout. Stdout is byte-clean: all progress and errors go to stderr. `, Example: ` dbxcli get /remote/file.txt ./local-file.txt dbxcli get rev:a1c10ce0dd78 ./historical-file.txt dbxcli get -r /remote/folder ./local-folder + dbxcli get -r -w 8 /remote/folder ./local-folder dbxcli get /backups/src.tgz - | tar tz dbxcli get /file.txt - > local-copy.txt`, RunE: get, @@ -588,5 +680,6 @@ var getCmd = &cobra.Command{ func init() { RootCmd.AddCommand(getCmd) getCmd.Flags().BoolP("recursive", "r", false, "Recursively download a folder") + getCmd.Flags().IntP("workers", "w", downloadWorkersAuto, "Number of files to download concurrently with --recursive (0 = auto-tune from measured bandwidth)") enableStructuredOutput(getCmd) } diff --git a/cmd/get_parallel.go b/cmd/get_parallel.go new file mode 100644 index 0000000..6f8f691 --- /dev/null +++ b/cmd/get_parallel.go @@ -0,0 +1,430 @@ +package cmd + +import ( + "context" + "fmt" + "io" + "os" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/dustin/go-humanize" + "golang.org/x/term" +) + +const ( + // downloadWorkersAuto selects the number of concurrent file downloads + // automatically from the measured download throughput. + downloadWorkersAuto = 0 + + // autoDownloadWorkersInitial is the concurrency the automatic tuner starts + // from before it has measured any throughput. + autoDownloadWorkersInitial = 4 + + // autoDownloadWorkersMax caps the concurrency the automatic tuner may reach. + // Dropbox rate limits make very high connection counts counterproductive. + autoDownloadWorkersMax = 16 + + // autoTuneGainThreshold is the relative throughput improvement a higher + // concurrency level must deliver to be kept. Smaller gains are treated as + // noise, which means the link is saturated and the tuner settles. + autoTuneGainThreshold = 0.10 +) + +var ( + // downloadMonitorInterval controls how often the status line is redrawn + // and throughput samples are collected. + downloadMonitorInterval = 500 * time.Millisecond + + // downloadAutoTuneWindow is the sampling window used to measure throughput + // for one concurrency level before deciding whether to change it. + downloadAutoTuneWindow = 2 * time.Second +) + +// concurrencyLimiter bounds the number of concurrently running download jobs. +// Unlike a channel semaphore, the limit can be raised or lowered while jobs +// are running; lowering it only delays new jobs and never interrupts running +// ones. +type concurrencyLimiter struct { + mu sync.Mutex + cond *sync.Cond + limit int + active int + waiting int + waited bool +} + +func newConcurrencyLimiter(limit int) *concurrencyLimiter { + if limit < 1 { + limit = 1 + } + l := &concurrencyLimiter{limit: limit} + l.cond = sync.NewCond(&l.mu) + return l +} + +// acquire blocks until a slot is available or ctx is done. +func (l *concurrencyLimiter) acquire(ctx context.Context) error { + stop := context.AfterFunc(ctx, func() { + l.mu.Lock() + l.cond.Broadcast() + l.mu.Unlock() + }) + defer stop() + + l.mu.Lock() + defer l.mu.Unlock() + for l.active >= l.limit { + if err := ctx.Err(); err != nil { + return err + } + l.waited = true + l.waiting++ + l.cond.Wait() + l.waiting-- + } + if err := ctx.Err(); err != nil { + return err + } + l.active++ + return nil +} + +func (l *concurrencyLimiter) release() { + l.mu.Lock() + l.active-- + l.cond.Broadcast() + l.mu.Unlock() +} + +func (l *concurrencyLimiter) setLimit(limit int) { + if limit < 1 { + limit = 1 + } + l.mu.Lock() + l.limit = limit + l.cond.Broadcast() + l.mu.Unlock() +} + +func (l *concurrencyLimiter) currentLimit() int { + l.mu.Lock() + defer l.mu.Unlock() + return l.limit +} + +// takeSaturated reports whether the limit was binding at any point since the +// previous call, meaning a job had to wait for a slot. The flag is reset. +func (l *concurrencyLimiter) takeSaturated() bool { + l.mu.Lock() + defer l.mu.Unlock() + saturated := l.waited || (l.waiting > 0 && l.active >= l.limit) + l.waited = false + return saturated +} + +// adaptiveConcurrency is a hill-climbing controller that discovers how many +// concurrent downloads the available bandwidth can use. It raises the +// concurrency step by step while each step improves aggregate throughput by +// at least autoTuneGainThreshold, and settles on the last level that did as +// soon as a step stops paying off, which indicates the link is saturated. +type adaptiveConcurrency struct { + limit int + max int + best float64 + bestLimit int + warmup bool + settled bool +} + +func newAdaptiveConcurrency(initial, max int) *adaptiveConcurrency { + if max < 1 { + max = 1 + } + if initial < 1 { + initial = 1 + } + if initial > max { + initial = max + } + return &adaptiveConcurrency{ + limit: initial, + max: max, + bestLimit: initial, + settled: initial >= max, + } +} + +// observe feeds one throughput sample (bytes per second) measured at the +// current limit and returns the limit to use for the next window. saturated +// reports whether the limit was actually binding during the window; samples +// taken while it was not are uninformative and leave the limit unchanged. +func (c *adaptiveConcurrency) observe(bytesPerSecond float64, saturated bool) int { + if c.settled || !saturated || bytesPerSecond <= 0 { + return c.limit + } + if c.warmup { + // The first window after a change still contains transfers that + // started under the old limit, so skip it. + c.warmup = false + return c.limit + } + + if c.best == 0 || bytesPerSecond >= c.best*(1+autoTuneGainThreshold) { + c.best = bytesPerSecond + c.bestLimit = c.limit + if c.limit >= c.max { + c.settled = true + return c.limit + } + c.limit = nextConcurrencyStep(c.limit, c.max) + c.warmup = true + return c.limit + } + + // The last increase did not buy a meaningful gain: the link is saturated. + c.limit = c.bestLimit + c.settled = true + return c.limit +} + +func nextConcurrencyStep(limit, max int) int { + next := limit + limit/2 + if next <= limit { + next = limit + 1 + } + if next > max { + next = max + } + return next +} + +// downloadStatusWriter serialises stderr output from concurrent downloads and +// keeps a single live status line at the bottom when stderr is a terminal. +type downloadStatusWriter struct { + mu sync.Mutex + w io.Writer + live bool + lastLen int +} + +func newDownloadStatusWriter(w io.Writer) *downloadStatusWriter { + if w == nil { + w = io.Discard + } + live := false + if f, ok := w.(*os.File); ok { + live = term.IsTerminal(int(f.Fd())) + } + return &downloadStatusWriter{w: w, live: live} +} + +func (s *downloadStatusWriter) clearLineLocked() { + if s.lastLen == 0 { + return + } + _, _ = fmt.Fprint(s.w, "\r"+strings.Repeat(" ", s.lastLen)+"\r") + s.lastLen = 0 +} + +// message prints a complete line, temporarily clearing the live status line. +func (s *downloadStatusWriter) message(format string, args ...any) { + s.mu.Lock() + defer s.mu.Unlock() + s.clearLineLocked() + _, _ = fmt.Fprintf(s.w, format, args...) +} + +// status redraws the live status line. It is a no-op unless stderr is a +// terminal, so logs and pipes only receive complete lines. +func (s *downloadStatusWriter) status(line string) { + if !s.live { + return + } + s.mu.Lock() + defer s.mu.Unlock() + // Pad to the widest line drawn so far so shorter updates leave no + // stale characters behind, and remember that width for clearing. + padded := line + if len(padded) < s.lastLen { + padded += strings.Repeat(" ", s.lastLen-len(padded)) + } + s.lastLen = len(padded) + _, _ = fmt.Fprint(s.w, "\r"+padded) +} + +// finish clears the live status line, if any, and prints a final summary +// line in its place. +func (s *downloadStatusWriter) finish(summary string) { + s.mu.Lock() + defer s.mu.Unlock() + s.clearLineLocked() + _, _ = fmt.Fprintln(s.w, summary) +} + +// downloadPool runs file download jobs with bounded, optionally self-tuning, +// concurrency and tracks aggregate progress across all of them. +type downloadPool struct { + limiter *concurrencyLimiter + tuner *adaptiveConcurrency + status *downloadStatusWriter + totalFiles int + totalBytes int64 + bytes atomic.Int64 + done atomic.Int64 + peakLimit int + started time.Time + meter *transferRateMeter +} + +// newDownloadPool creates a pool for jobs files. workers selects a fixed +// concurrency, or downloadWorkersAuto to tune it from measured throughput. +func newDownloadPool(workers int, totalFiles int, totalBytes int64, status *downloadStatusWriter) *downloadPool { + if status == nil { + status = newDownloadStatusWriter(io.Discard) + } + p := &downloadPool{ + status: status, + totalFiles: totalFiles, + totalBytes: totalBytes, + } + if workers == downloadWorkersAuto { + p.tuner = newAdaptiveConcurrency( + min(autoDownloadWorkersInitial, max(totalFiles, 1)), + min(autoDownloadWorkersMax, max(totalFiles, 1)), + ) + workers = p.tuner.limit + } + p.limiter = newConcurrencyLimiter(workers) + p.peakLimit = workers + return p +} + +// concurrent reports whether more than one download may run at a time. +func (p *downloadPool) concurrent() bool { + if p.tuner != nil { + return p.tuner.max > 1 + } + return p.limiter.currentLimit() > 1 +} + +// addBytes records newly committed download bytes for throughput tracking. +func (p *downloadPool) addBytes(n int64) { + if n > 0 { + p.bytes.Add(n) + } +} + +// fileProgress returns a progress callback for a single file that feeds only +// newly committed bytes into the pool's aggregate counter. +func (p *downloadPool) fileProgress() downloadProgressFunc { + var last int64 + var mu sync.Mutex + return func(committed, _ int64) { + mu.Lock() + defer mu.Unlock() + if committed > last { + p.addBytes(committed - last) + last = committed + } + } +} + +// run executes jobs in order, starting each as soon as a slot is free, and +// returns when all of them have finished. Jobs are never skipped: when ctx is +// cancelled the remaining jobs receive ctx's error through onError. +func (p *downloadPool) run(ctx context.Context, jobs []func(), onError func(index int, err error)) { + p.started = time.Now() + p.meter = newTransferRateMeter(p.started) + + // Sequential downloads keep their own per-file progress bar, so the + // aggregate status line and tuner only run in concurrent mode. + concurrent := p.concurrent() + stop := make(chan struct{}) + var monitor sync.WaitGroup + if concurrent { + monitor.Add(1) + go func() { + defer monitor.Done() + p.monitor(stop) + }() + } + + var wg sync.WaitGroup + for i, job := range jobs { + if err := p.limiter.acquire(ctx); err != nil { + onError(i, err) + p.done.Add(1) + continue + } + wg.Add(1) + go func() { + defer wg.Done() + defer p.limiter.release() + defer p.done.Add(1) + job() + }() + } + wg.Wait() + + close(stop) + monitor.Wait() + if concurrent { + p.status.finish(p.summary()) + } +} + +func (p *downloadPool) monitor(stop <-chan struct{}) { + ticker := time.NewTicker(downloadMonitorInterval) + defer ticker.Stop() + + windowStart := p.started + windowBytes := p.bytes.Load() + for { + select { + case <-stop: + return + case now := <-ticker.C: + total := p.bytes.Load() + p.meter.update(now, total) + if p.tuner != nil && now.Sub(windowStart) >= downloadAutoTuneWindow { + elapsed := now.Sub(windowStart).Seconds() + rate := float64(total-windowBytes) / elapsed + limit := p.tuner.observe(rate, p.limiter.takeSaturated()) + p.limiter.setLimit(limit) + if limit > p.peakLimit { + p.peakLimit = limit + } + windowStart, windowBytes = now, total + } + p.status.status(p.statusLine(now, total)) + } + } +} + +// statusLine renders the live aggregate line, for example: +// +// Downloading 12/40 files 45%|========= | 1.2 GiB/2.7 GiB [00:12<00:15, 98 MiB/s, 6 workers] +func (p *downloadPool) statusLine(now time.Time, bytes int64) string { + return fmt.Sprintf("Downloading %d/%d files %s", + p.done.Load(), p.totalFiles, + formatTransferProgress(bytes, p.totalBytes, p.meter.elapsed(now), p.meter.rate(), + fmt.Sprintf(", %d workers", p.limiter.currentLimit()))) +} + +// summary renders the final line printed once every job has finished, with +// the average throughput over the whole run. +func (p *downloadPool) summary() string { + elapsed := time.Since(p.started) + bytes := max(p.bytes.Load(), 0) + average := 0.0 + if seconds := elapsed.Seconds(); seconds > 0 { + average = float64(bytes) / seconds + } + return fmt.Sprintf("Downloaded %d/%d files, %s in %s (%s), up to %d workers", + p.done.Load(), p.totalFiles, + humanize.IBytes(uint64(bytes)), formatTransferClock(elapsed), formatTransferRate(average), + p.peakLimit) +} diff --git a/cmd/get_parallel_test.go b/cmd/get_parallel_test.go new file mode 100644 index 0000000..5dbac5f --- /dev/null +++ b/cmd/get_parallel_test.go @@ -0,0 +1,666 @@ +package cmd + +import ( + "bytes" + "context" + "io" + "os" + "path/filepath" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/dropbox/dropbox-sdk-go-unofficial/v6/dropbox/files" + "github.com/spf13/cobra" +) + +func TestAdaptiveConcurrencyClimbsWhileThroughputImproves(t *testing.T) { + c := newAdaptiveConcurrency(4, 16) + + if got := c.observe(100, true); got != 6 { + t.Fatalf("first sample: limit = %d, want 6", got) + } + // The window right after a change is a warm-up and must not change anything. + if got := c.observe(150, true); got != 6 { + t.Fatalf("warm-up sample: limit = %d, want 6", got) + } + if got := c.observe(150, true); got != 9 { + t.Fatalf("improved sample: limit = %d, want 9", got) + } + if got := c.observe(160, true); got != 9 { + t.Fatalf("warm-up sample: limit = %d, want 9", got) + } + // Less than a 10% gain means the link is saturated: fall back and settle. + if got := c.observe(155, true); got != 6 { + t.Fatalf("plateau sample: limit = %d, want 6", got) + } + if !c.settled { + t.Fatal("expected controller to settle after a plateau") + } + if got := c.observe(1000, true); got != 6 { + t.Fatalf("settled controller changed limit to %d", got) + } +} + +func TestAdaptiveConcurrencyIgnoresUninformativeSamples(t *testing.T) { + c := newAdaptiveConcurrency(4, 16) + + if got := c.observe(100, false); got != 4 { + t.Fatalf("unsaturated sample: limit = %d, want 4", got) + } + if got := c.observe(0, true); got != 4 { + t.Fatalf("zero sample: limit = %d, want 4", got) + } + if c.best != 0 { + t.Fatalf("best = %v, want no measurement recorded", c.best) + } +} + +func TestAdaptiveConcurrencySettlesAtMax(t *testing.T) { + c := newAdaptiveConcurrency(4, 6) + + if got := c.observe(100, true); got != 6 { + t.Fatalf("limit = %d, want 6", got) + } + c.observe(100, true) // warm-up + if got := c.observe(200, true); got != 6 { + t.Fatalf("limit = %d, want 6", got) + } + if !c.settled { + t.Fatal("expected controller to settle at max") + } +} + +func TestNewAdaptiveConcurrencyClampsInitialToMax(t *testing.T) { + c := newAdaptiveConcurrency(4, 2) + if c.limit != 2 || !c.settled { + t.Fatalf("limit = %d, settled = %v, want 2 and settled", c.limit, c.settled) + } + + c = newAdaptiveConcurrency(0, 0) + if c.limit != 1 || c.max != 1 { + t.Fatalf("limit = %d, max = %d, want 1 and 1", c.limit, c.max) + } +} + +func TestNextConcurrencyStep(t *testing.T) { + steps := []int{4} + for steps[len(steps)-1] < 16 { + steps = append(steps, nextConcurrencyStep(steps[len(steps)-1], 16)) + } + want := []int{4, 6, 9, 13, 16} + if len(steps) != len(want) { + t.Fatalf("steps = %v, want %v", steps, want) + } + for i := range want { + if steps[i] != want[i] { + t.Fatalf("steps = %v, want %v", steps, want) + } + } + if got := nextConcurrencyStep(1, 16); got != 2 { + t.Fatalf("nextConcurrencyStep(1) = %d, want 2", got) + } +} + +// concurrencyProbe records how many jobs run at once and holds the first +// `barrier` jobs until that many are in flight, which proves they overlap. +type concurrencyProbe struct { + barrier int32 + active atomic.Int32 + peak atomic.Int32 + calls atomic.Int32 + ready chan struct{} + once sync.Once + timeout atomic.Bool +} + +func newConcurrencyProbe(barrier int) *concurrencyProbe { + return &concurrencyProbe{barrier: int32(barrier), ready: make(chan struct{})} +} + +func (p *concurrencyProbe) enter() { + n := p.active.Add(1) + for { + peak := p.peak.Load() + if n <= peak || p.peak.CompareAndSwap(peak, n) { + break + } + } + if p.calls.Add(1) >= p.barrier { + p.once.Do(func() { close(p.ready) }) + } + select { + case <-p.ready: + case <-time.After(10 * time.Second): + p.timeout.Store(true) + } +} + +func (p *concurrencyProbe) leave() { + p.active.Add(-1) +} + +func (p *concurrencyProbe) check(t *testing.T, wantPeak int) { + t.Helper() + if p.timeout.Load() { + t.Fatalf("jobs never reached %d in flight", p.barrier) + } + if got := int(p.peak.Load()); got != wantPeak { + t.Fatalf("peak concurrency = %d, want %d", got, wantPeak) + } +} + +func TestDownloadPoolRunsJobsUpToFixedLimit(t *testing.T) { + probe := newConcurrencyProbe(3) + pool := newDownloadPool(3, 8, 0, nil) + if !pool.concurrent() { + t.Fatal("expected a fixed limit above 1 to be concurrent") + } + + var ran atomic.Int32 + jobs := make([]func(), 8) + for i := range jobs { + jobs[i] = func() { + probe.enter() + defer probe.leave() + ran.Add(1) + } + } + pool.run(context.Background(), jobs, func(int, error) { + t.Error("unexpected dispatch error") + }) + + probe.check(t, 3) + if ran.Load() != 8 { + t.Fatalf("ran %d jobs, want 8", ran.Load()) + } + if pool.done.Load() != 8 { + t.Fatalf("done = %d, want 8", pool.done.Load()) + } +} + +func TestDownloadPoolSingleWorkerRunsSequentially(t *testing.T) { + probe := newConcurrencyProbe(1) + pool := newDownloadPool(1, 4, 0, nil) + if pool.concurrent() { + t.Fatal("expected a single worker not to be concurrent") + } + + jobs := make([]func(), 4) + for i := range jobs { + jobs[i] = func() { + probe.enter() + defer probe.leave() + } + } + pool.run(context.Background(), jobs, func(int, error) { + t.Error("unexpected dispatch error") + }) + probe.check(t, 1) +} + +func TestDownloadPoolAutoStartsWithInitialConcurrency(t *testing.T) { + pool := newDownloadPool(downloadWorkersAuto, 40, 0, nil) + if pool.tuner == nil { + t.Fatal("expected automatic tuner") + } + if got := pool.limiter.currentLimit(); got != autoDownloadWorkersInitial { + t.Fatalf("initial limit = %d, want %d", got, autoDownloadWorkersInitial) + } + if pool.tuner.max != autoDownloadWorkersMax { + t.Fatalf("max = %d, want %d", pool.tuner.max, autoDownloadWorkersMax) + } + + // Small folders never need more workers than files. + pool = newDownloadPool(downloadWorkersAuto, 2, 0, nil) + if got := pool.limiter.currentLimit(); got != 2 { + t.Fatalf("initial limit for 2 files = %d, want 2", got) + } + if pool.tuner.max != 2 { + t.Fatalf("max for 2 files = %d, want 2", pool.tuner.max) + } + + pool = newDownloadPool(downloadWorkersAuto, 1, 0, nil) + if pool.concurrent() { + t.Fatal("a single file must download sequentially with per-file progress") + } +} + +func TestDownloadPoolAutoTunesWhileRunning(t *testing.T) { + restoreMonitorIntervals(t, time.Millisecond, 2*time.Millisecond) + + pool := newDownloadPool(downloadWorkersAuto, 64, 64<<20, nil) + var maxSeen atomic.Int32 + jobs := make([]func(), 64) + for i := range jobs { + jobs[i] = func() { + progress := pool.fileProgress() + for step := int64(1); step <= 4; step++ { + progress(step<<18, 1<<20) + time.Sleep(time.Millisecond) + } + limit := int32(pool.limiter.currentLimit()) + for { + seen := maxSeen.Load() + if limit <= seen || maxSeen.CompareAndSwap(seen, limit) { + break + } + } + } + } + pool.run(context.Background(), jobs, func(int, error) { + t.Error("unexpected dispatch error") + }) + + if pool.done.Load() != 64 { + t.Fatalf("done = %d, want 64", pool.done.Load()) + } + if got := pool.bytes.Load(); got != 64<<20 { + t.Fatalf("bytes = %d, want %d", got, 64<<20) + } + if limit := int(maxSeen.Load()); limit < 1 || limit > autoDownloadWorkersMax { + t.Fatalf("observed limit %d outside [1, %d]", limit, autoDownloadWorkersMax) + } + if pool.peakLimit < autoDownloadWorkersInitial || pool.peakLimit > autoDownloadWorkersMax { + t.Fatalf("peak limit %d outside [%d, %d]", pool.peakLimit, autoDownloadWorkersInitial, autoDownloadWorkersMax) + } +} + +func TestDownloadPoolFileProgressCountsOnlyNewBytes(t *testing.T) { + pool := newDownloadPool(2, 1, 0, nil) + progress := pool.fileProgress() + progress(10, 100) + progress(10, 100) + progress(35, 100) + progress(20, 100) // never goes backwards + if got := pool.bytes.Load(); got != 35 { + t.Fatalf("bytes = %d, want 35", got) + } +} + +func TestDownloadPoolCancelledContextFailsRemainingJobs(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + pool := newDownloadPool(2, 3, 0, nil) + var ran atomic.Int32 + jobs := []func(){ + func() { ran.Add(1) }, + func() { ran.Add(1) }, + func() { ran.Add(1) }, + } + var mu sync.Mutex + var failed []int + pool.run(ctx, jobs, func(i int, err error) { + if err != context.Canceled { + t.Errorf("job %d error = %v, want context.Canceled", i, err) + } + mu.Lock() + failed = append(failed, i) + mu.Unlock() + }) + + if ran.Load() != 0 { + t.Fatalf("ran %d jobs after cancellation, want 0", ran.Load()) + } + if len(failed) != 3 { + t.Fatalf("failed jobs = %v, want all three", failed) + } +} + +func TestConcurrencyLimiterReportsSaturation(t *testing.T) { + l := newConcurrencyLimiter(1) + if l.takeSaturated() { + t.Fatal("fresh limiter reported saturation") + } + if err := l.acquire(context.Background()); err != nil { + t.Fatal(err) + } + + acquired := make(chan struct{}) + go func() { + _ = l.acquire(context.Background()) + close(acquired) + }() + + deadline := time.Now().Add(5 * time.Second) + for !l.takeSaturated() { + if time.Now().After(deadline) { + t.Fatal("limiter never reported saturation while a job waited") + } + time.Sleep(time.Millisecond) + } + + l.setLimit(2) + select { + case <-acquired: + case <-time.After(5 * time.Second): + t.Fatal("raising the limit did not wake the waiting job") + } + l.release() + l.release() +} + +func TestDownloadStatusWriterKeepsPipesLineOriented(t *testing.T) { + var buf bytes.Buffer + w := newDownloadStatusWriter(&buf) + if w.live { + t.Fatal("a buffer must not be treated as a terminal") + } + w.status("Downloading 1/2 files") + w.message("Downloading %s -> %s\n", "/a", "a") + w.finish("done") + if got := buf.String(); got != "Downloading /a -> a\ndone\n" { + t.Fatalf("output = %q, want only complete lines", got) + } +} + +func TestDownloadStatusWriterClearsStatusLineBeforeMessages(t *testing.T) { + var buf bytes.Buffer + w := newDownloadStatusWriter(&buf) + w.live = true + + w.status("abcdef") + w.status("xy") + w.message("hi\n") + w.finish("done") + + want := "\rabcdef" + "\rxy " + "\r \r" + "hi\n" + "done\n" + if got := buf.String(); got != want { + t.Fatalf("output = %q, want %q", got, want) + } +} + +func TestParseGetOptionsDefaultsToAutomaticWorkers(t *testing.T) { + opts, err := parseGetOptions(testGetCmd()) + if err != nil { + t.Fatal(err) + } + if opts.workers != downloadWorkersAuto { + t.Fatalf("workers = %d, want auto (%d)", opts.workers, downloadWorkersAuto) + } + + opts, err = parseGetOptions(getCmd) + if err != nil { + t.Fatal(err) + } + if opts.workers != downloadWorkersAuto { + t.Fatalf("registered flag default workers = %d, want auto (%d)", opts.workers, downloadWorkersAuto) + } +} + +func TestParseGetOptionsRejectsNegativeWorkers(t *testing.T) { + cmd := testGetWorkersCmd(-1) + _, err := parseGetOptions(cmd) + if err == nil { + t.Fatal("expected error for negative workers") + } + if code := jsonErrorCode(err); code != jsonErrorCodeInvalidArguments { + t.Fatalf("error code = %q, want %q", code, jsonErrorCodeInvalidArguments) + } + if details := jsonErrorDetails(err); details["flag"] != "workers" { + t.Fatalf("details = %#v, want flag=workers", details) + } +} + +func testGetWorkersCmd(workers int) *cobra.Command { + cmd := testGetCmd() + cmd.Flags().IntP("workers", "w", downloadWorkersAuto, "") + if err := cmd.Flags().Set("recursive", "true"); err != nil { + panic(err) + } + if err := cmd.Flags().Set("workers", strconv.Itoa(workers)); err != nil { + panic(err) + } + return cmd +} + +func restoreMonitorIntervals(t *testing.T, monitor, window time.Duration) { + t.Helper() + previousMonitor, previousWindow := downloadMonitorInterval, downloadAutoTuneWindow + downloadMonitorInterval, downloadAutoTuneWindow = monitor, window + t.Cleanup(func() { + downloadMonitorInterval, downloadAutoTuneWindow = previousMonitor, previousWindow + }) +} + +func parallelGetMock(t *testing.T, paths []string, probe *concurrencyProbe) *mockFilesClient { + t.Helper() + entries := []files.IsMetadata{getTestFolderMetadata("/remote")} + for _, p := range paths { + entries = append(entries, getTestFileMetadata(p, 4)) + } + return &mockFilesClient{ + getMetadataFn: func(arg *files.GetMetadataArg) (files.IsMetadata, error) { + return getTestFolderMetadata(arg.Path), nil + }, + listFolderFn: func(arg *files.ListFolderArg) (*files.ListFolderResult, error) { + return &files.ListFolderResult{Entries: entries}, nil + }, + downloadFn: func(arg *files.DownloadArg) (*files.FileMetadata, io.ReadCloser, error) { + if probe != nil { + probe.enter() + defer probe.leave() + } + return getTestFileMetadata(arg.Path, 4), io.NopCloser(strings.NewReader("data")), nil + }, + } +} + +func TestGetRecursiveDownloadsFilesInParallelByDefault(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out") + paths := []string{"/remote/a.txt", "/remote/b.txt", "/remote/c.txt", "/remote/d.txt", "/remote/e.txt", "/remote/f.txt"} + probe := newConcurrencyProbe(autoDownloadWorkersInitial) + stubFilesClient(t, parallelGetMock(t, paths, probe)) + + var stderr bytes.Buffer + cmd := testGetCmd() + cmd.SetErr(&stderr) + if err := cmd.Flags().Set("recursive", "true"); err != nil { + t.Fatal(err) + } + if err := get(cmd, []string{"/remote", dst}); err != nil { + t.Fatalf("get error: %v", err) + } + + probe.check(t, autoDownloadWorkersInitial) + for _, p := range paths { + got, err := os.ReadFile(filepath.Join(dst, filepath.Base(p))) + if err != nil { + t.Fatalf("read %s: %v", p, err) + } + if string(got) != "data" { + t.Fatalf("%s content = %q, want data", p, got) + } + if !strings.Contains(stderr.String(), "Downloading "+p+" -> ") { + t.Fatalf("stderr = %q, want a line for %s", stderr.String(), p) + } + } + if strings.Contains(stderr.String(), "\r") { + t.Fatalf("stderr = %q, want no progress-bar control characters on a pipe", stderr.String()) + } + summary := stderr.String()[strings.LastIndex(strings.TrimSpace(stderr.String()), "\n")+1:] + if !strings.HasPrefix(summary, "Downloaded 6/6 files, 24 B in 00:0") || !strings.Contains(summary, "/s), up to 4 workers") { + t.Fatalf("summary = %q, want file count, bytes, elapsed time, and average throughput", summary) + } +} + +func TestGetRecursiveWorkersFlagFixesConcurrency(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out") + paths := []string{"/remote/a.txt", "/remote/b.txt", "/remote/c.txt", "/remote/d.txt", "/remote/e.txt"} + probe := newConcurrencyProbe(2) + stubFilesClient(t, parallelGetMock(t, paths, probe)) + + cmd := testGetWorkersCmd(2) + cmd.SetErr(io.Discard) + if err := get(cmd, []string{"/remote", dst}); err != nil { + t.Fatalf("get error: %v", err) + } + probe.check(t, 2) + for _, p := range paths { + if _, err := os.Stat(filepath.Join(dst, filepath.Base(p))); err != nil { + t.Fatalf("%s not downloaded: %v", p, err) + } + } +} + +func TestGetRecursiveSingleWorkerKeepsSequentialProgress(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out") + paths := []string{"/remote/a.txt", "/remote/b.txt", "/remote/c.txt"} + probe := newConcurrencyProbe(1) + stubFilesClient(t, parallelGetMock(t, paths, probe)) + + var stderr bytes.Buffer + cmd := testGetWorkersCmd(1) + cmd.SetErr(&stderr) + if err := get(cmd, []string{"/remote", dst}); err != nil { + t.Fatalf("get error: %v", err) + } + probe.check(t, 1) + // The sequential path keeps the per-file progress bar on stderr, with + // throughput and ETA, and never prints the aggregate summary. + if !strings.Contains(stderr.String(), "Downloading 100%|====================| 4 B/4 B [00:00<00:00, ") { + t.Fatalf("stderr = %q, want per-file progress", stderr.String()) + } + if strings.Contains(stderr.String(), "Downloaded 3/3 files") { + t.Fatalf("stderr = %q, want no aggregate summary in sequential mode", stderr.String()) + } +} + +func TestGetJSONRecursiveResultsKeepListingOrderUnderConcurrency(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out") + paths := []string{"/remote/a.txt", "/remote/b.txt", "/remote/c.txt"} + mock := parallelGetMock(t, paths, nil) + + // Make the first-listed file finish last so completion order differs + // from listing order. + lastStarted := make(chan struct{}) + var startOnce sync.Once + download := mock.downloadFn + mock.downloadFn = func(arg *files.DownloadArg) (*files.FileMetadata, io.ReadCloser, error) { + switch arg.Path { + case "/remote/c.txt": + startOnce.Do(func() { close(lastStarted) }) + case "/remote/a.txt": + select { + case <-lastStarted: + case <-time.After(10 * time.Second): + t.Error("c.txt never started while a.txt waited") + } + } + return download(arg) + } + stubFilesClient(t, mock) + + var stdout bytes.Buffer + cmd := testGetJSONCmd(&stdout, nil) + cmd.Flags().IntP("workers", "w", downloadWorkersAuto, "") + if err := cmd.Flags().Set("recursive", "true"); err != nil { + t.Fatal(err) + } + if err := cmd.Flags().Set("workers", "3"); err != nil { + t.Fatal(err) + } + if err := get(cmd, []string{"/remote", dst}); err != nil { + t.Fatalf("get error: %v", err) + } + + got := decodeGetOutput(t, &stdout) + wantTargets := []string{dst, filepath.Join(dst, "a.txt"), filepath.Join(dst, "b.txt"), filepath.Join(dst, "c.txt")} + if len(got.Results) != len(wantTargets) { + t.Fatalf("results = %+v, want %d entries", got.Results, len(wantTargets)) + } + for i, want := range wantTargets { + if got.Results[i].Input.Target != want { + t.Fatalf("results[%d].target = %q, want %q (results: %+v)", i, got.Results[i].Input.Target, want, got.Results) + } + } + if got.Results[0].Kind != getKindFolder || got.Results[1].Kind != getKindFile { + t.Fatalf("kinds = %s, %s; want folder then file", got.Results[0].Kind, got.Results[1].Kind) + } +} + +func TestGetRecursiveConcurrentReportsErrorsInListingOrder(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out") + paths := []string{"/remote/bad1.txt", "/remote/good.txt", "/remote/bad2.txt", "/remote/also-good.txt"} + mock := parallelGetMock(t, paths, nil) + download := mock.downloadFn + mock.downloadFn = func(arg *files.DownloadArg) (*files.FileMetadata, io.ReadCloser, error) { + if strings.Contains(arg.Path, "bad") { + return nil, nil, &files.DownloadAPIError{} + } + return download(arg) + } + stubFilesClient(t, mock) + + var stderr bytes.Buffer + cmd := testGetWorkersCmd(4) + cmd.SetErr(&stderr) + err := get(cmd, []string{"/remote", dst}) + if err == nil { + t.Fatal("expected error for failed downloads") + } + if !strings.Contains(err.Error(), "2 error(s)") { + t.Fatalf("error = %q, want 2 error(s)", err.Error()) + } + first := strings.Index(stderr.String(), "Error: /remote/bad1.txt") + second := strings.Index(stderr.String(), "Error: /remote/bad2.txt") + if first < 0 || second < 0 || first > second { + t.Fatalf("stderr = %q, want errors reported in listing order", stderr.String()) + } + for _, name := range []string{"good.txt", "also-good.txt"} { + if _, statErr := os.Stat(filepath.Join(dst, name)); statErr != nil { + t.Fatalf("%s not downloaded: %v", name, statErr) + } + } +} + +func TestGetRecursiveConcurrentExportOnlyFilesCountBytes(t *testing.T) { + dst := filepath.Join(t.TempDir(), "out") + exported := "exported paper" + paper := getTestFileMetadata("/remote/doc.paper", uint64(len(exported))) + paper.ExportInfo = &files.ExportInfo{ExportAs: "markdown"} + entries := []files.IsMetadata{ + getTestFolderMetadata("/remote"), + paper, + getTestFileMetadata("/remote/a.txt", 4), + } + mock := &mockFilesClient{ + getMetadataFn: func(arg *files.GetMetadataArg) (files.IsMetadata, error) { + return getTestFolderMetadata(arg.Path), nil + }, + listFolderFn: func(arg *files.ListFolderArg) (*files.ListFolderResult, error) { + return &files.ListFolderResult{Entries: entries}, nil + }, + downloadFn: func(arg *files.DownloadArg) (*files.FileMetadata, io.ReadCloser, error) { + return getTestFileMetadata(arg.Path, 4), io.NopCloser(strings.NewReader("data")), nil + }, + exportFn: func(arg *files.ExportArg) (*files.ExportResult, io.ReadCloser, error) { + meta := getTestFileMetadata(arg.Path, uint64(len(exported))) + return &files.ExportResult{ + ExportMetadata: &files.ExportMetadata{Name: "doc.md", Size: uint64(len(exported))}, + FileMetadata: meta, + }, io.NopCloser(strings.NewReader(exported)), nil + }, + } + stubFilesClient(t, mock) + + cmd := testGetWorkersCmd(2) + cmd.SetErr(io.Discard) + if err := get(cmd, []string{"/remote", dst}); err != nil { + t.Fatalf("get error: %v", err) + } + got, err := os.ReadFile(filepath.Join(dst, "doc.md")) + if err != nil { + t.Fatalf("read exported file: %v", err) + } + if string(got) != exported { + t.Fatalf("exported content = %q, want %q", got, exported) + } + if _, err := os.Stat(filepath.Join(dst, "a.txt")); err != nil { + t.Fatalf("a.txt not downloaded: %v", err) + } +} diff --git a/cmd/help_manifest.go b/cmd/help_manifest.go index 3b179a2..2b708e4 100644 --- a/cmd/help_manifest.go +++ b/cmd/help_manifest.go @@ -131,7 +131,7 @@ var commandManifestRegistry = map[string]jsonCommandManifestMetadata{ {Description: "Download a file", Command: "dbxcli get /remote.txt ./remote.txt"}, {Description: "Download a file revision", Command: "dbxcli get rev:a1c10ce0dd78 ./historical.txt"}, }, - Flags: map[string]jsonCommandFlagMetadata{"recursive": {ValueKind: "boolean"}}, + Flags: map[string]jsonCommandFlagMetadata{"recursive": {ValueKind: "boolean"}, "workers": {ValueKind: "integer"}}, DropboxScopes: []string{"files.content.read", "files.metadata.read"}, StdinStdout: jsonCommandStdinStdout{WritesBinaryStdout: true}, Known: true, diff --git a/cmd/share_link_download.go b/cmd/share_link_download.go index 69f0eff..f57209a 100644 --- a/cmd/share_link_download.go +++ b/cmd/share_link_download.go @@ -26,7 +26,6 @@ import ( "github.com/dropbox/dbxcli/v3/internal/output" "github.com/dropbox/dropbox-sdk-go-unofficial/v6/dropbox/files" "github.com/dropbox/dropbox-sdk-go-unofficial/v6/dropbox/sharing" - "github.com/dustin/go-humanize" "github.com/mitchellh/ioprogress" "github.com/spf13/cobra" ) @@ -451,12 +450,9 @@ func copySharedLinkContentToFile(contents io.Reader, size uint64, dst string, er }() progressbar := &ioprogress.Reader{ - Reader: contents, - DrawFunc: ioprogress.DrawTerminalf(errOut, func(progress, total int64) string { - return fmt.Sprintf("Downloading %s/%s", - humanize.IBytes(uint64(progress)), humanize.IBytes(uint64(total))) - }), - Size: int64(size), + Reader: contents, + DrawFunc: newTransferProgressDrawer(errOut, "Downloading ").drawFunc(), + Size: int64(size), } _, copyErr := io.Copy(f, progressbar) diff --git a/cmd/transfer_progress.go b/cmd/transfer_progress.go new file mode 100644 index 0000000..2af83b3 --- /dev/null +++ b/cmd/transfer_progress.go @@ -0,0 +1,219 @@ +package cmd + +import ( + "fmt" + "io" + "strings" + "time" + + "github.com/dustin/go-humanize" + "github.com/mitchellh/ioprogress" +) + +const ( + // transferRateWindow is how far back the rate meter looks when computing + // the current throughput, so the figure tracks recent speed rather than + // the average since the start. + transferRateWindow = 5 * time.Second + + // transferRateSampleSpacing coalesces updates that arrive closer together + // than this so the sample window stays small. + transferRateSampleSpacing = 100 * time.Millisecond + + // transferProgressBarWidth is the width of the ASCII progress bar. + transferProgressBarWidth = 20 +) + +// transferProgressDrawInterval throttles progress redraws so a fast transfer +// does not flood stderr. +var transferProgressDrawInterval = 100 * time.Millisecond + +type transferRateSample struct { + at time.Time + bytes int64 +} + +// transferRateMeter estimates throughput from cumulative byte counts using a +// sliding window, the way tqdm smooths its rate display. +type transferRateMeter struct { + start time.Time + window time.Duration + samples []transferRateSample +} + +func newTransferRateMeter(start time.Time) *transferRateMeter { + return &transferRateMeter{ + start: start, + window: transferRateWindow, + samples: []transferRateSample{{at: start}}, + } +} + +// update records the cumulative number of bytes transferred as of now. +func (m *transferRateMeter) update(now time.Time, bytes int64) { + last := m.samples[len(m.samples)-1] + if len(m.samples) > 1 && now.Sub(last.at) < transferRateSampleSpacing { + m.samples[len(m.samples)-1] = transferRateSample{at: now, bytes: bytes} + } else { + m.samples = append(m.samples, transferRateSample{at: now, bytes: bytes}) + } + + // Drop samples that fell out of the window, keeping one older sample as + // the anchor so the window always spans at least its full width. + cutoff := now.Add(-m.window) + for len(m.samples) > 2 && !m.samples[1].at.After(cutoff) { + m.samples = m.samples[1:] + } +} + +// rate returns the recent throughput in bytes per second. +func (m *transferRateMeter) rate() float64 { + first := m.samples[0] + last := m.samples[len(m.samples)-1] + seconds := last.at.Sub(first.at).Seconds() + if seconds <= 0 || last.bytes <= first.bytes { + return 0 + } + return float64(last.bytes-first.bytes) / seconds +} + +// elapsed returns the time since the transfer started. +func (m *transferRateMeter) elapsed(now time.Time) time.Duration { + return now.Sub(m.start) +} + +// formatTransferProgress renders a tqdm-style progress line: +// +// 45%|========= | 45 MiB/100 MiB [00:12<00:15, 3.7 MiB/s] +// +// extra is appended inside the brackets, for example ", 6 workers". When the +// total is unknown the percentage, bar, and ETA are omitted. +func formatTransferProgress(done, total int64, elapsed time.Duration, rate float64, extra string) string { + if done < 0 { + done = 0 + } + if total <= 0 { + return fmt.Sprintf("%s [%s, %s%s]", + humanize.IBytes(uint64(done)), formatTransferClock(elapsed), formatTransferRate(rate), extra) + } + if done > total { + done = total + } + + percent := int(done * 100 / total) + filled := int(done * transferProgressBarWidth / total) + bar := strings.Repeat("=", filled) + strings.Repeat(" ", transferProgressBarWidth-filled) + + return fmt.Sprintf("%3d%%|%s| %s/%s [%s<%s, %s%s]", + percent, bar, + humanize.IBytes(uint64(done)), humanize.IBytes(uint64(total)), + formatTransferClock(elapsed), formatTransferETA(total-done, rate), + formatTransferRate(rate), extra) +} + +func formatTransferRate(rate float64) string { + if rate < 0 { + rate = 0 + } + return humanize.IBytes(uint64(rate)) + "/s" +} + +// formatTransferETA estimates the remaining time from the bytes left and the +// current rate, or "?" when no rate is known yet. +func formatTransferETA(remaining int64, rate float64) string { + if remaining <= 0 { + return formatTransferClock(0) + } + if rate <= 0 { + return "?" + } + seconds := float64(remaining) / rate + if seconds > 359999 { // cap at 99:59:59 to keep the line tidy + return ">99:59:59" + } + return formatTransferClock(time.Duration(seconds * float64(time.Second))) +} + +// formatTransferClock renders a duration as MM:SS, or H:MM:SS past an hour. +func formatTransferClock(d time.Duration) string { + if d < 0 { + d = 0 + } + total := int64(d.Round(time.Second) / time.Second) + hours, minutes, seconds := total/3600, (total%3600)/60, total%60 + if hours > 0 { + return fmt.Sprintf("%d:%02d:%02d", hours, minutes, seconds) + } + return fmt.Sprintf("%02d:%02d", minutes, seconds) +} + +// transferProgressDrawer draws a single-transfer progress line with throughput +// and ETA, throttled to transferProgressDrawInterval. +type transferProgressDrawer struct { + prefix string + draw ioprogress.DrawFunc + meter *transferRateMeter + now func() time.Time + lastDraw time.Time + finished bool +} + +func newTransferProgressDrawer(w io.Writer, prefix string) *transferProgressDrawer { + if w == nil { + w = io.Discard + } + d := &transferProgressDrawer{prefix: prefix, now: time.Now} + d.meter = newTransferRateMeter(d.now()) + d.draw = ioprogress.DrawTerminalf(w, func(progress, total int64) string { + return d.line(progress, total) + }) + return d +} + +func (d *transferProgressDrawer) line(progress, total int64) string { + now := d.now() + return d.prefix + formatTransferProgress(progress, total, d.meter.elapsed(now), d.meter.rate(), "") +} + +// update records progress and redraws the line if enough time has passed or +// the transfer is complete. +func (d *transferProgressDrawer) update(progress, total int64) { + if d.finished { + return + } + now := d.now() + d.meter.update(now, progress) + complete := total > 0 && progress >= total + if !complete && !d.lastDraw.IsZero() && now.Sub(d.lastDraw) < transferProgressDrawInterval { + return + } + d.lastDraw = now + _ = d.draw(progress, total) +} + +// finish ends the progress line. It is safe to call more than once. +func (d *transferProgressDrawer) finish() { + if d.finished { + return + } + d.finished = true + _ = d.draw(-1, -1) +} + +// drawFunc adapts the drawer to ioprogress.Reader, which performs its own +// throttling and signals completion with (-1, -1). +func (d *transferProgressDrawer) drawFunc() ioprogress.DrawFunc { + return func(progress, total int64) error { + if progress == -1 && total == -1 { + d.finish() + return nil + } + if d.finished { + return nil + } + now := d.now() + d.meter.update(now, progress) + d.lastDraw = now + return d.draw(progress, total) + } +} diff --git a/cmd/transfer_progress_test.go b/cmd/transfer_progress_test.go new file mode 100644 index 0000000..3c795c7 --- /dev/null +++ b/cmd/transfer_progress_test.go @@ -0,0 +1,176 @@ +package cmd + +import ( + "bytes" + "strings" + "testing" + "time" +) + +func TestFormatTransferProgress(t *testing.T) { + got := formatTransferProgress(45<<20, 100<<20, 12*time.Second, 3.7*float64(1<<20), "") + want := " 45%|========= | 45 MiB/100 MiB [00:12<00:15, 3.7 MiB/s]" + if got != want { + t.Fatalf("progress = %q, want %q", got, want) + } + + got = formatTransferProgress(2700<<20, 2700<<20, 90*time.Second, 30<<20, ", 6 workers") + want = "100%|====================| 2.6 GiB/2.6 GiB [01:30<00:00, 30 MiB/s, 6 workers]" + if got != want { + t.Fatalf("complete progress = %q, want %q", got, want) + } + + got = formatTransferProgress(0, 1<<30, 0, 0, "") + want = " 0%| | 0 B/1.0 GiB [00:0099:59:59" { + t.Fatalf("huge ETA = %q, want cap", got) + } +} + +func TestTransferRateMeterUsesSlidingWindow(t *testing.T) { + start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + m := newTransferRateMeter(start) + if m.rate() != 0 { + t.Fatalf("initial rate = %v, want 0", m.rate()) + } + + // 1 MiB/s for the first ten seconds. + for i := 1; i <= 10; i++ { + m.update(start.Add(time.Duration(i)*time.Second), int64(i)<<20) + } + if got := m.rate(); got != float64(1<<20) { + t.Fatalf("steady rate = %v, want %v", got, 1<<20) + } + + // Then 10 MiB/s: after the window has rolled over, only the recent + // speed should be reported, not the average since the start. + bytes := int64(10 << 20) + for i := 11; i <= 20; i++ { + bytes += 10 << 20 + m.update(start.Add(time.Duration(i)*time.Second), bytes) + } + if got := m.rate(); got != float64(10<<20) { + t.Fatalf("recent rate = %v, want %v", got, 10<<20) + } + if got := m.elapsed(start.Add(20 * time.Second)); got != 20*time.Second { + t.Fatalf("elapsed = %v, want 20s", got) + } + if len(m.samples) > 8 { + t.Fatalf("meter kept %d samples, want the window only", len(m.samples)) + } +} + +func TestTransferRateMeterCoalescesRapidUpdates(t *testing.T) { + start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + m := newTransferRateMeter(start) + for i := 1; i <= 1000; i++ { + m.update(start.Add(time.Duration(i)*time.Millisecond), int64(i)*1024) + } + if len(m.samples) > 3 { + t.Fatalf("meter kept %d samples for sub-interval updates, want coalesced", len(m.samples)) + } + if got := m.rate(); got != float64(1024*1000) { + t.Fatalf("rate = %v, want %v", got, 1024*1000) + } +} + +func TestTransferProgressDrawerThrottlesAndReportsRate(t *testing.T) { + var buf bytes.Buffer + d := newTransferProgressDrawer(&buf, "Downloading ") + now := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + d.now = func() time.Time { return now } + d.meter = newTransferRateMeter(now) + + total := int64(100 << 20) + d.update(1<<20, total) // first update always draws + now = now.Add(10 * time.Millisecond) + d.update(2<<20, total) // too soon: skipped + now = now.Add(time.Second) + d.update(10<<20, total) // ~1s elapsed: drawn with rate and ETA + now = now.Add(time.Second) + d.update(total, total) // completion always draws + d.finish() + d.finish() // idempotent + + out := buf.String() + frames := strings.Split(strings.TrimSuffix(out, "\n"), "\r") + if len(frames) != 4 || frames[3] != "" { + t.Fatalf("frames = %q, want three drawn frames and a trailing newline", frames) + } + if !strings.HasPrefix(frames[0], "Downloading 1%|") { + t.Fatalf("first frame = %q", frames[0]) + } + if want := "Downloading 10%|== | 10 MiB/100 MiB [00:01<00:09, 9.9 MiB/s]"; !strings.HasPrefix(frames[1], want) { + t.Fatalf("second frame = %q, want prefix %q", frames[1], want) + } + if want := "Downloading 100%|====================| 100 MiB/100 MiB [00:02<00:00, 50 MiB/s]"; !strings.HasPrefix(frames[2], want) { + t.Fatalf("final frame = %q, want prefix %q", frames[2], want) + } + if strings.Count(out, "\n") != 1 { + t.Fatalf("output = %q, want exactly one newline from finish", out) + } +} + +func TestTransferProgressDrawerDrawFuncHandlesReaderProtocol(t *testing.T) { + var buf bytes.Buffer + d := newTransferProgressDrawer(&buf, "Downloading ") + draw := d.drawFunc() + if err := draw(512, 1024); err != nil { + t.Fatal(err) + } + if err := draw(-1, -1); err != nil { + t.Fatal(err) + } + if err := draw(1024, 1024); err != nil { + t.Fatal(err) + } + out := buf.String() + if !strings.Contains(out, "Downloading 50%|========== | 512 B/1.0 KiB [") { + t.Fatalf("output = %q, want a half-way frame", out) + } + if !strings.HasSuffix(out, "\n") || strings.Contains(out, "1.0 KiB/1.0 KiB") { + t.Fatalf("output = %q, want nothing drawn after finish", out) + } +} diff --git a/docs/commands/dbxcli_get.md b/docs/commands/dbxcli_get.md index 1afdf7f..c69b32e 100644 --- a/docs/commands/dbxcli_get.md +++ b/docs/commands/dbxcli_get.md @@ -10,6 +10,9 @@ Download a file or folder from Dropbox. - Source may be a Dropbox path, file ID (id:), revision (rev:), or namespace-relative path (ns:). - Use --recursive (-r) to download entire directories. + - Recursive downloads fetch several files in parallel. By default the + number of concurrent downloads is tuned automatically from the measured + throughput; use --workers (-w) to set a fixed number instead. - Use - as target to write file bytes to stdout. Stdout is byte-clean: all progress and errors go to stderr. @@ -24,6 +27,7 @@ dbxcli get [flags] [] dbxcli get /remote/file.txt ./local-file.txt dbxcli get rev:a1c10ce0dd78 ./historical-file.txt dbxcli get -r /remote/folder ./local-folder + dbxcli get -r -w 8 /remote/folder ./local-folder dbxcli get /backups/src.tgz - | tar tz dbxcli get /file.txt - > local-copy.txt ``` @@ -31,8 +35,9 @@ dbxcli get [flags] [] ### Options ``` - -h, --help help for get - -r, --recursive Recursively download a folder + -h, --help help for get + -r, --recursive Recursively download a folder + -w, --workers int Number of files to download concurrently with --recursive (0 = auto-tune from measured bandwidth) ``` ### Options inherited from parent commands