diff --git a/download/download.go b/download/download.go index a71c500..3b422f1 100644 --- a/download/download.go +++ b/download/download.go @@ -34,6 +34,7 @@ import ( "github.com/google/googet/v2/client" "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/oswrap" + "github.com/google/googet/v2/progress" "github.com/google/logger" ) @@ -84,11 +85,9 @@ func packageHTTP(ctx context.Context, url, dst, chksum string, downloader *clien if err != nil { return err } - if ok && size < length { - logger.Infof("resuming download of %s (%d bytes remaining)", url, length-size) - req.Header.Add("Range", fmt.Sprintf("bytes=%d-", size)) - } else { - // Get rid of the old file and download from start, resetting hash. + // restart discards any partial download so the response body is written + // from the beginning of the file. + restart := func() error { if err := f.Truncate(0); err != nil { return err } @@ -96,6 +95,14 @@ func packageHTTP(ctx context.Context, url, dst, chksum string, downloader *clien return err } hash.Reset() + size = 0 + return nil + } + if ok && size < length { + logger.Infof("resuming download of %s (%d bytes remaining)", url, length-size) + req.Header.Add("Range", fmt.Sprintf("bytes=%d-", size)) + } else if err := restart(); err != nil { + return err } resp, err := downloader.HTTPClient.Do(req) if err != nil { @@ -103,13 +110,30 @@ func packageHTTP(ctx context.Context, url, dst, chksum string, downloader *clien } defer resp.Body.Close() if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent { - return fmt.Errorf("downloading %s: %v", url, err) + return fmt.Errorf("downloading %s: unexpected status %s", url, resp.Status) + } + if resp.StatusCode == http.StatusOK && size > 0 { + // The server ignored the Range header and is sending the whole file; + // appending it to the partial download would corrupt the file and + // overstate the progress total. + logger.Infof("server ignored range request for %s, restarting download", url) + if err := restart(); err != nil { + return err + } + } + // The total is what is already on disk plus what the server will send. + total := int64(-1) + if resp.ContentLength >= 0 { + total = size + resp.ContentLength } + bar := progress.NewBar(fmt.Sprintf("Downloading %s", filepath.Base(dst)), total, size) // Continue hashing the file as we download it. - n, err := io.Copy(io.MultiWriter(hash, f), resp.Body) + n, err := io.Copy(io.MultiWriter(hash, f, bar), resp.Body) if err != nil { + bar.Abort() return fmt.Errorf("downloading %s: %v", url, err) } + bar.Finish() // Verify the checksum of the fully downloaded file. if sum := hex.EncodeToString(hash.Sum(nil)); sum != chksum { os.RemoveAll(dst) // delete the bad file @@ -134,7 +158,7 @@ func packageGCS(ctx context.Context, bucket, object string, dst, chksum string) defer r.Close() logger.Infof("Downloading gs://%s/%s", bucket, object) - return download(r, dst, chksum) + return download(r, r.Attrs.Size, dst, chksum) } // FromRepo downloads a package from a repo. It returns the path to the @@ -162,7 +186,9 @@ func Latest(ctx context.Context, name, dir string, rm client.RepoMap, archs []st return FromRepo(ctx, rs, repo, dir, downloader) } -func download(r io.Reader, dst, chksum string) (err error) { +// download copies r to dst, verifying the SHA256 checksum, and renders a +// progress bar when enabled. +func download(r io.Reader, size int64, dst, chksum string) (err error) { f, err := oswrap.Create(dst) if err != nil { return err @@ -173,13 +199,16 @@ func download(r io.Reader, dst, chksum string) (err error) { } }() + bar := progress.NewBar(fmt.Sprintf("Downloading %s", filepath.Base(dst)), size, 0) hash := sha256.New() - tw := io.MultiWriter(f, hash) + tw := io.MultiWriter(f, hash, bar) b, err := io.Copy(tw, r) if err != nil { + bar.Abort() return err } + bar.Finish() if hex.EncodeToString(hash.Sum(nil)) != chksum { fmt.Println(hex.EncodeToString(hash.Sum(nil)), chksum) diff --git a/download/download_test.go b/download/download_test.go index dfc5664..381f13b 100644 --- a/download/download_test.go +++ b/download/download_test.go @@ -17,23 +17,31 @@ import ( "archive/tar" "bytes" "compress/gzip" - "io/ioutil" + "context" + "io" + "net/http" + "net/http/httptest" + "os" "path" "path/filepath" + "slices" + "strconv" + "strings" "testing" + "github.com/google/googet/v2/client" "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/oswrap" "github.com/google/logger" ) func init() { - logger.Init("test", true, false, ioutil.Discard) + logger.Init("test", true, false, io.Discard) } func TestDownload(t *testing.T) { r := bytes.NewReader([]byte("some content")) - tempDir, err := ioutil.TempDir("", "") + tempDir, err := os.MkdirTemp("", "") if err != nil { t.Fatalf("error creating temp directory: %v", err) } @@ -44,16 +52,132 @@ func TestDownload(t *testing.T) { t.Errorf("error seeking to front of reader: %v", err) } tempFile := path.Join(tempDir, "test") - if err := download(r, tempFile, chksum); err != nil { + if err := download(r, int64(r.Len()), tempFile, chksum); err != nil { t.Errorf("error downloading and checking checksum: %v", err) } - if err := download(r, tempFile, "notachecksum"); err == nil { + if err := download(r, int64(r.Len()), tempFile, "notachecksum"); err == nil { t.Error("wanted but did not recieve checksum error") } } +func TestDownloadUnknownSize(t *testing.T) { + content := []byte("some content") + chksum := goolib.Checksum(bytes.NewReader(content)) + tempFile := path.Join(t.TempDir(), "test") + // A size of 0 means the total is unknown; the checksum must still verify. + if err := download(bytes.NewReader(content), 0, tempFile, chksum); err != nil { + t.Errorf("error downloading with unknown size: %v", err) + } + got, err := os.ReadFile(tempFile) + if err != nil { + t.Fatalf("error reading downloaded file: %v", err) + } + if !bytes.Equal(got, content) { + t.Errorf("downloaded contents = %q, want %q", got, content) + } +} + +// rangeServer serves payload, advertising range support on HEAD. When +// honorRange is false it answers ranged GETs with the whole file and 200 OK, +// as some proxies do. +func rangeServer(t *testing.T, payload []byte, honorRange bool, requests *[]string) *httptest.Server { + t.Helper() + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + *requests = append(*requests, r.Method+" "+r.Header.Get("Range")) + w.Header().Set("Accept-Ranges", "bytes") + if r.Method == http.MethodHead { + w.Header().Set("Content-Length", strconv.Itoa(len(payload))) + return + } + if rng := r.Header.Get("Range"); rng != "" && honorRange { + start, err := strconv.Atoi(strings.TrimSuffix(strings.TrimPrefix(rng, "bytes="), "-")) + if err != nil { + t.Errorf("bad Range header %q: %v", rng, err) + w.WriteHeader(http.StatusBadRequest) + return + } + w.Header().Set("Content-Length", strconv.Itoa(len(payload)-start)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[start:]) + return + } + w.Header().Set("Content-Length", strconv.Itoa(len(payload))) + w.Write(payload) + })) +} + +func TestPackageHTTP(t *testing.T) { + payload := bytes.Repeat([]byte("0123456789"), 100) + chksum := goolib.Checksum(bytes.NewReader(payload)) + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader: %v", err) + } + for _, tc := range []struct { + desc string + existing []byte // Contents written to dst before the download. + honorRange bool + wantGETs []string + }{ + { + // An empty destination still asks for bytes=0-; the server may + // answer 200 or 206 and either way nothing is on disk to keep. + desc: "fresh download", + wantGETs: []string{"GET bytes=0-"}, + }, + { + desc: "resumed download", + existing: payload[:400], + honorRange: true, + wantGETs: []string{"GET bytes=400-"}, + }, + { + desc: "server ignores range", + existing: payload[:400], + wantGETs: []string{"GET bytes=400-"}, + }, + { + desc: "already downloaded", + existing: payload, + wantGETs: nil, + }, + } { + t.Run(tc.desc, func(t *testing.T) { + var requests []string + srv := rangeServer(t, payload, tc.honorRange, &requests) + defer srv.Close() + + dst := filepath.Join(t.TempDir(), "pkg.goo") + if tc.existing != nil { + if err := os.WriteFile(dst, tc.existing, 0644); err != nil { + t.Fatalf("writing existing file: %v", err) + } + } + if err := packageHTTP(context.Background(), srv.URL+"/pkg.goo", dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP() = %v, want nil", err) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("reading downloaded file: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("downloaded %d bytes, want %d bytes matching the payload", len(got), len(payload)) + } + var gets []string + for _, r := range requests { + if strings.HasPrefix(r, "GET") { + gets = append(gets, r) + } + } + if !slices.Equal(gets, tc.wantGETs) { + t.Errorf("GET requests = %q, want %q", gets, tc.wantGETs) + } + }) + } +} + func TestExtractPkg(t *testing.T) { - tempDir, err := ioutil.TempDir("", "") + tempDir, err := os.MkdirTemp("", "") if err != nil { t.Fatalf("error creating temp directory: %v", err) } @@ -94,7 +218,7 @@ func TestExtractPkg(t *testing.T) { t.Fatalf("error running ExtractPkg: %v", err) } - cts, err := ioutil.ReadFile(filepath.Join(dst, filepath.Clean(name))) + cts, err := os.ReadFile(filepath.Join(dst, filepath.Clean(name))) if err != nil { t.Fatalf("error opening test file: %v", err) } @@ -104,7 +228,7 @@ func TestExtractPkg(t *testing.T) { } func TestExtractPkgPathTraversal(t *testing.T) { - tempDir, err := ioutil.TempDir("", "") + tempDir, err := os.MkdirTemp("", "") if err != nil { t.Fatalf("error creating temp directory: %v", err) } diff --git a/go.mod b/go.mod index 412293a..558c1cd 100644 --- a/go.mod +++ b/go.mod @@ -13,6 +13,7 @@ require ( github.com/yusufpapurcu/wmi v1.2.4 golang.org/x/oauth2 v0.37.0 golang.org/x/sys v0.48.0 + golang.org/x/term v0.46.0 google.golang.org/api v0.299.0 gopkg.in/yaml.v3 v3.0.1 modernc.org/sqlite v1.59.0 diff --git a/go.sum b/go.sum index f5bfa1c..742406a 100644 --- a/go.sum +++ b/go.sum @@ -153,6 +153,8 @@ golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= +golang.org/x/term v0.46.0 h1:3+OXuTbaKDgwk8jTi3aSLHRlmWqHEUDUtxnbFigO4YE= +golang.org/x/term v0.46.0/go.mod h1:+K02xbkittuwc0Am4abfA3Fc+XRGXkvBXNO88NCXPoc= golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI= golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E= golang.org/x/time v0.16.0 h1:vMb6ptszcQMkcwiRTAuNNU50gom6++Q/6gY2hDM6VDE= diff --git a/googet.go b/googet.go index 49a7793..f467e24 100644 --- a/googet.go +++ b/googet.go @@ -22,6 +22,7 @@ import ( "os" "github.com/google/googet/v2/googetdb" + "github.com/google/googet/v2/progress" "github.com/google/googet/v2/settings" "github.com/google/googet/v2/system" "github.com/google/logger" @@ -83,6 +84,7 @@ func run(ctx context.Context) int { noConfirm := flag.Bool("noconfirm", false, "skip confirmation") verbose := flag.Bool("verbose", false, "print info level logs to stdout") systemLog := flag.Bool("system_log", true, "log to Linux Syslog or Windows Event Log") + progressFlag := flag.Bool("progress", true, "show a download progress bar and install spinner when stderr is a terminal (never with -verbose); overrides progress in googet.conf") showVer := flag.Bool("version", false, "display GooGet version and exit") if flagParse != nil { @@ -102,6 +104,7 @@ func run(ctx context.Context) int { cmdr.Register(cmdr.HelpCommand(), "") cmdr.ImportantFlag("verbose") cmdr.ImportantFlag("noconfirm") + cmdr.ImportantFlag("progress") // These commands may execute without a lock and before any initialization. cmdName := flag.Arg(0) // empty string if no args @@ -161,6 +164,14 @@ func run(ctx context.Context) int { logger.Init("GooGet", *verbose, *systemLog, lf) defer logger.Close() + progressSet := false + flag.Visit(func(f *flag.Flag) { + if f.Name == "progress" { + progressSet = true + } + }) + progress.Init(wantProgress(settings.Progress, progressSet, *progressFlag, *verbose)) + if err := googetdb.CreateIfMissing(dbFile); err != nil { logger.Errorf("Unable to create initial db file; if db is not created, run again as admin: %v", err) return 1 @@ -175,3 +186,19 @@ func run(ctx context.Context) int { } return int(cmdr.Execute(ctx)) } + +// wantProgress reports whether progress output should be rendered, given the +// progress setting from googet.conf, whether -progress was set explicitly and +// its value, and -verbose. Progress is still subject to the terminal checks in +// progress.Init. An explicit -progress, true or false, overrides the config; +// -verbose always disables progress because it interleaves INFO logs on +// stdout, which would corrupt a redrawn line. +func wantProgress(conf, flagSet, flagVal, verbose bool) bool { + if verbose { + return false + } + if flagSet { + return flagVal + } + return conf +} diff --git a/googet_test.go b/googet_test.go index c9bad4e..7f6b75d 100644 --- a/googet_test.go +++ b/googet_test.go @@ -71,3 +71,26 @@ func TestRotateLog(t *testing.T) { } } } + +func TestWantProgress(t *testing.T) { + for _, tc := range []struct { + desc string + conf, flagSet, flagVal, verbose bool + want bool + }{ + {desc: "defaults", conf: true, want: true}, + {desc: "config off", conf: false, want: false}, + {desc: "flag off overrides config on", conf: true, flagSet: true, flagVal: false, want: false}, + {desc: "flag on overrides config off", conf: false, flagSet: true, flagVal: true, want: true}, + {desc: "unset flag value is ignored", conf: false, flagSet: false, flagVal: true, want: false}, + {desc: "verbose disables", conf: true, verbose: true, want: false}, + {desc: "verbose beats explicit flag on", conf: true, flagSet: true, flagVal: true, verbose: true, want: false}, + } { + t.Run(tc.desc, func(t *testing.T) { + if got := wantProgress(tc.conf, tc.flagSet, tc.flagVal, tc.verbose); got != tc.want { + t.Errorf("wantProgress(conf=%v, flagSet=%v, flagVal=%v, verbose=%v) = %v, want %v", + tc.conf, tc.flagSet, tc.flagVal, tc.verbose, got, tc.want) + } + }) + } +} diff --git a/goolib/goolib.go b/goolib/goolib.go index a5d6218..7de5a1e 100644 --- a/goolib/goolib.go +++ b/goolib/goolib.go @@ -20,7 +20,6 @@ import ( "encoding/hex" "fmt" "io" - "os" "os/exec" "path/filepath" "regexp" @@ -28,6 +27,8 @@ import ( "slices" "strings" "syscall" + + "github.com/google/googet/v2/progress" ) var interpreter = map[string]string{ @@ -79,10 +80,13 @@ func Exec(s string, args []string, ec []int, w io.Writer) error { // Run runs a command. // The process is successful if the exit code matches any of those provided or '0'. -// stdout and stderr are sent to the writer and to this process's stdout and stderr. +// stdout and stderr are sent to the writer and to this process's stdout and +// stderr. While a progress spinner is active they are still shown as they are +// produced, a line at a time after clearing the spinner line; nothing is +// withheld or discarded. func Run(c *exec.Cmd, ec []int, w io.Writer) error { - c.Stdout = io.MultiWriter(os.Stdout, w) - c.Stderr = io.MultiWriter(os.Stderr, w) + c.Stdout = io.MultiWriter(progress.Stdout(), w) + c.Stderr = io.MultiWriter(progress.Stderr(), w) if err := c.Run(); err != nil { e, ok := err.(*exec.ExitError) if !ok { diff --git a/goolib/goolib_test.go b/goolib/goolib_test.go index 404dbda..51a89b4 100644 --- a/goolib/goolib_test.go +++ b/goolib/goolib_test.go @@ -14,11 +14,16 @@ limitations under the License. package goolib import ( + "bytes" "fmt" + "io" "math/rand" + "os" + "os/exec" + "runtime" "strings" + "sync" "testing" - "time" ) func TestScriptInterpreter(t *testing.T) { @@ -59,7 +64,6 @@ func randString(runes []rune, min, max int) string { } func TestSplitGCSUrl(t *testing.T) { - rand.Seed(time.Now().UnixNano()) const alphanum = "abcdefghijklmnopqrstuvwxyz0123456789" objChars := alphanum + "ABCDEFGHIJKLMNOPQRSTUVWXYZ-_.~@%^=+" bucket := randString([]rune(alphanum), 1, 1) + randString([]rune(alphanum+"-_."), 0, 61) + randString([]rune(alphanum), 1, 1) @@ -131,3 +135,90 @@ func TestSplitGCSUrl(t *testing.T) { } } } + +// syncBuffer is a bytes.Buffer that is safe for concurrent writers. +type syncBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *syncBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *syncBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +// captureStdio runs fn with os.Stdout and os.Stderr redirected to pipes and +// returns what was written to each. +func captureStdio(t *testing.T, fn func()) (stdout, stderr string) { + t.Helper() + capture := func(f **os.File) (restore func() string) { + r, w, err := os.Pipe() + if err != nil { + t.Fatalf("os.Pipe: %v", err) + } + orig := *f + *f = w + done := make(chan string) + go func() { + b, _ := io.ReadAll(r) + done <- string(b) + }() + return func() string { + w.Close() + *f = orig + return <-done + } + } + restoreOut := capture(&os.Stdout) + restoreErr := capture(&os.Stderr) + fn() + return restoreOut(), restoreErr() +} + +func TestRun(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses /bin/sh") + } + for _, tt := range []struct { + name string + script string + ec []int + wantErr bool + }{ + {"success", "echo out; echo err >&2", nil, false}, + {"accepted exit code", "echo out; echo err >&2; exit 3", []int{3}, false}, + {"rejected exit code", "echo out; echo err >&2; exit 3", nil, true}, + } { + t.Run(tt.name, func(t *testing.T) { + // Run writes stdout and stderr to w from separate goroutines, so + // the shared writer must be safe for concurrent use. + var w syncBuffer + var err error + stdout, stderr := captureStdio(t, func() { + err = Run(exec.Command("/bin/sh", "-c", tt.script), tt.ec, &w) + }) + if (err != nil) != tt.wantErr { + t.Errorf("Run() error = %v, wantErr %v", err, tt.wantErr) + } + // Without an active spinner the child's output must reach the + // process's own stdout and stderr unchanged, so unattended callers + // see exactly what they saw before progress reporting existed. + if stdout != "out\n" { + t.Errorf("stdout = %q, want %q", stdout, "out\n") + } + if stderr != "err\n" { + t.Errorf("stderr = %q, want %q", stderr, "err\n") + } + if got := w.String(); !strings.Contains(got, "out\n") || !strings.Contains(got, "err\n") { + t.Errorf("writer got %q, want both stdout and stderr lines", got) + } + }) + } +} diff --git a/install/install.go b/install/install.go index cf8e3c6..62b5431 100644 --- a/install/install.go +++ b/install/install.go @@ -18,6 +18,7 @@ import ( "context" "crypto/sha256" "encoding/hex" + "errors" "fmt" "io" "os" @@ -30,6 +31,7 @@ import ( "github.com/google/googet/v2/googetdb" "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/oswrap" + "github.com/google/googet/v2/progress" "github.com/google/googet/v2/remove" "github.com/google/googet/v2/settings" "github.com/google/googet/v2/system" @@ -419,7 +421,7 @@ func makeInstallFunction(src, dst string, insFiles map[string]string, dbOnly, fo } else { logger.Infof("Warning: file conflict: %s is already owned by package %s, overwriting because `StrictConflicts` is not set", outPath, owner) } - fmt.Printf("Warning: file conflict: %s is already owned by package %s, overwriting...\n", outPath, owner) + progress.Printf("Warning: file conflict: %s is already owned by package %s, overwriting...\n", outPath, owner) } if dbOnly { @@ -549,7 +551,24 @@ func buildConflictMap(db *googetdb.GooDB, currentPkg string) (map[string]string, return conflictMap, nil } -func installPkg(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (map[string]string, error) { +// errInstallInterrupted is the spinner status error used when an install +// unwinds without returning, such as on a panic. +var errInstallInterrupted = errors.New("install interrupted") + +// installPkg extracts and installs a package, rendering a spinner on +// interactive terminals for the duration of the install. +func installPkg(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (insFiles map[string]string, err error) { + sp := progress.NewSpinner(fmt.Sprintf("Installing %s.%s.%s", ps.Name, ps.Arch, ps.Version)) + // The spinner is stopped by a deferred call so that no exit path leaves it + // redrawing. A normal return overwrites err before the deferred call runs; + // a panic leaves errInstallInterrupted in place, so it renders "failed". + err = errInstallInterrupted + defer func() { sp.Stop(err) }() + return installPkgInner(pkg, ps, dbOnly, force, db) +} + +// installPkgInner extracts the package, copies its files and runs its install script. +func installPkgInner(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (map[string]string, error) { dir, err := download.ExtractPkg(pkg) if err != nil { return nil, err diff --git a/progress/progress.go b/progress/progress.go new file mode 100644 index 0000000..dca7dac --- /dev/null +++ b/progress/progress.go @@ -0,0 +1,478 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +// Package progress renders a download progress bar and an install spinner on +// interactive terminals. +// +// Rendering is opt-in: nothing is written unless Init has enabled it and +// stderr is a terminal. Every constructor returns nil when rendering is +// disabled and every method is nil-safe, so callers never need to branch on +// whether progress is active. +// +// Child process output is never hidden. While a spinner is active, installer +// output is passed through to the real stdout and stderr a complete line at a +// time, after first clearing the spinner line, so it scrolls above the +// spinner instead of being interleaved with it. An unterminated line is held +// back for at most about a second before it is written out anyway. +package progress + +import ( + "bytes" + "fmt" + "io" + "os" + "strings" + "sync" + "time" + + "github.com/dustin/go-humanize" + "golang.org/x/term" +) + +const ( + // barWidth is the number of cells in the rendered progress bar. + barWidth = 35 + // redrawInterval throttles how often the download bar is redrawn. Serial + // console loggers record every frame, so this is deliberately coarser + // than what an interactive terminal alone would need. + redrawInterval = 200 * time.Millisecond + // spinInterval is how often the spinner advances a frame. + spinInterval = 250 * time.Millisecond + // maxPartialLine bounds how much of an unterminated child output line is + // held back before it is written out anyway. + maxPartialLine = 4096 + // partialLineDelay bounds how long an unterminated child output line is + // held back, so a partial line followed by a long silent step still shows. + partialLineDelay = 500 * time.Millisecond +) + +// frames are the ASCII spinner frames; ASCII keeps rendering identical on +// conhost, Windows Terminal, SSH sessions and serial consoles. +var frames = []byte{'-', '\\', '|', '/'} + +var ( + // mu guards all package state and serializes writes to out. + mu sync.Mutex + // enabled reports whether progress output is rendered. + enabled bool + // out is where progress lines are rendered; stderr, like curl and Cabbie. + out io.Writer = os.Stderr + // stdout is where non-progress console text is written. + stdout io.Writer = os.Stdout + // active is the spinner currently owning the console line, if any. + active *Spinner + // lastLine is the line currently rendered with a carriage return, used to + // blank stale characters when a shorter line replaces it and to skip + // redraws that would not change anything (serial console loggers record + // every frame, so identical frames are pure noise). + lastLine string + // now is overridable by tests. + now = time.Now +) + +// Init enables progress rendering when allow is true, stderr is a terminal +// and the terminal is not declared dumb. It should be called once from main +// after flag parsing; until then rendering is disabled. +func Init(allow bool) { + mu.Lock() + defer mu.Unlock() + // TERM=dumb (Emacs shells, some CI and serial consoles) means carriage + // returns are not interpreted, so a redrawn line would render as garbage. + enabled = allow && os.Getenv("TERM") != "dumb" && isTerminal(os.Stderr) +} + +// isTerminal reports whether f is an interactive terminal or Windows console. +// It is a variable so tests can stub it. +var isTerminal = func(f *os.File) bool { + return term.IsTerminal(int(f.Fd())) +} + +// termWidth returns the width of the stderr terminal in columns, or 0 if it +// is unknown. It is a variable so tests can stub it. +var termWidth = func() int { + w, _, err := term.GetSize(int(os.Stderr.Fd())) + if err != nil { + return 0 + } + return w +} + +// spinnerSuffixLen is the widest text a spinner renders after its title. +const spinnerSuffixLen = len("... failed [00:00]") + +// fitTitle cuts title so a spinner line fits within the terminal width, +// keeping the last column free. A line that wraps cannot be redrawn with a +// carriage return, so each frame would otherwise land on a new line. +func fitTitle(title string) string { + room := termWidth() - 1 - spinnerSuffixLen + if r := []rune(title); room > 0 && len(r) > room { + return strings.TrimRight(string(r[:room]), " ") + } + return title +} + +// Enabled reports whether progress output is being rendered. +func Enabled() bool { + mu.Lock() + defer mu.Unlock() + return enabled +} + +// Stdout returns the writer that child processes should use for console +// stdout. With no active spinner it is os.Stdout itself. While a spinner is +// active it is a pass-through to os.Stdout that clears the spinner line before +// writing each batch of complete lines; nothing is withheld or discarded. +func Stdout() io.Writer { + mu.Lock() + defer mu.Unlock() + if active != nil { + return active.stdout + } + return os.Stdout +} + +// Stderr returns the writer that child processes should use for console +// stderr. It behaves like Stdout, targeting os.Stderr. +func Stderr() io.Writer { + mu.Lock() + defer mu.Unlock() + if active != nil { + return active.stderr + } + return os.Stderr +} + +// Printf writes to stdout like fmt.Printf, first clearing any active bar or +// spinner line so the text starts at column zero. The bar or spinner redraws +// itself on its next update. +func Printf(format string, a ...any) { + mu.Lock() + defer mu.Unlock() + clearLocked() + fmt.Fprintf(stdout, format, a...) +} + +// redrawLocked overwrites the current console line with s, blanking any +// trailing characters left over from a longer previous line. It writes +// nothing if s is already what is on the line. +func redrawLocked(s string) { + if s == lastLine { + return + } + pad := "" + if n := len(lastLine) - len(s); n > 0 { + pad = strings.Repeat(" ", n) + } + fmt.Fprintf(out, "\r%s%s", s, pad) + lastLine = s +} + +// clearLocked blanks the current console line, if one was rendered. +func clearLocked() { + if lastLine == "" { + return + } + fmt.Fprintf(out, "\r%s\r", strings.Repeat(" ", len(lastLine))) + lastLine = "" +} + +// endLineLocked terminates the current console line. +func endLineLocked() { + fmt.Fprintln(out) + lastLine = "" +} + +// elapsed formats the time since start as mm:ss. +func elapsed(start, t time.Time) string { + d := t.Sub(start).Round(time.Second) + if d < 0 { + d = 0 + } + return fmt.Sprintf("%02d:%02d", int(d/time.Minute), int((d%time.Minute)/time.Second)) +} + +// Bar is a byte-counting progress bar that implements io.Writer so it can be +// added to an io.MultiWriter alongside the real destination. +type Bar struct { + title string + total int64 // A total <= 0 means the size is unknown. + cur int64 + start time.Time + lastDraw time.Time +} + +// NewBar prints a title line and returns a bar expecting total bytes, of +// which initial bytes are already complete (for resumed downloads). It returns +// nil when rendering is disabled. +func NewBar(title string, total, initial int64) *Bar { + mu.Lock() + defer mu.Unlock() + if !enabled { + return nil + } + clearLocked() + if total > 0 { + fmt.Fprintf(out, "%s (%s)\n", title, humanize.IBytes(uint64(total))) + } else { + fmt.Fprintf(out, "%s\n", title) + } + b := &Bar{title: title, total: total, cur: initial, start: now()} + b.drawLocked(b.start) + return b +} + +// Write records len(p) bytes of progress and redraws the bar at most once per +// redrawInterval. It never fails, so it cannot abort the surrounding copy. +func (b *Bar) Write(p []byte) (int, error) { + if b == nil { + return len(p), nil + } + mu.Lock() + defer mu.Unlock() + b.cur += int64(len(p)) + if t := now(); t.Sub(b.lastDraw) >= redrawInterval { + b.drawLocked(t) + } + return len(p), nil +} + +// Finish renders the bar as complete and terminates the line. +func (b *Bar) Finish() { + if b == nil { + return + } + mu.Lock() + defer mu.Unlock() + if b.total > 0 { + b.cur = b.total + } + b.drawLocked(now()) + endLineLocked() +} + +// Abort terminates the line without marking the bar complete, leaving the +// partial state visible above whatever error follows. +func (b *Bar) Abort() { + if b == nil { + return + } + mu.Lock() + defer mu.Unlock() + endLineLocked() +} + +// drawLocked renders the bar line for time t. +func (b *Bar) drawLocked(t time.Time) { + b.lastDraw = t + if b.total <= 0 { + redrawLocked(fmt.Sprintf(" %s [%s]", humanize.IBytes(uint64(b.cur)), elapsed(b.start, t))) + return + } + cur := b.cur + if cur > b.total { + cur = b.total + } + pct := int(cur * 100 / b.total) + bar := " " + if w := fitBarWidth(); w > 0 { + n := pct * w / 100 + bar = "|" + strings.Repeat("=", n) + strings.Repeat("-", w-n) + "| " + } + redrawLocked(fmt.Sprintf("%s%3d%% %s / %s [%s]", bar, pct, + humanize.IBytes(uint64(cur)), humanize.IBytes(uint64(b.total)), elapsed(b.start, t))) +} + +// barTextLen is the widest text a bar line renders besides its cells. +const barTextLen = len("|| 100% 1023 KiB / 1023 KiB [00:00]") + +// minBarWidth is the narrowest bar worth drawing; below it only the +// percentage and byte counts are shown. +const minBarWidth = 10 + +// fitBarWidth returns how many bar cells fit within the terminal width, +// keeping the last column free, or 0 if the bar should be omitted. An unknown +// width gets the full barWidth. +func fitBarWidth() int { + w := termWidth() + if w <= 0 { + return barWidth + } + n := min(barWidth, w-1-barTextLen) + if n < minBarWidth { + return 0 + } + return n +} + +// Spinner renders an indeterminate "title... /" line until Stop is called. +// At most one spinner is active at a time. +type Spinner struct { + title string + start time.Time + done chan struct{} + wg sync.WaitGroup + stopped bool + // stdout and stderr pass child process output through to the real + // streams while the spinner is active. + stdout *lineWriter + stderr *lineWriter +} + +// NewSpinner starts rendering a spinner for title. It returns nil when +// rendering is disabled or another spinner is already active. +func NewSpinner(title string) *Spinner { + mu.Lock() + defer mu.Unlock() + if !enabled || active != nil { + return nil + } + clearLocked() + s := &Spinner{title: title, start: now(), done: make(chan struct{})} + s.stdout = &lineWriter{owner: s, dst: os.Stdout} + s.stderr = &lineWriter{owner: s, dst: os.Stderr} + active = s + s.drawLocked(0, s.start) + s.wg.Add(1) + go s.run() + return s +} + +// run redraws the spinner until done is closed. +func (s *Spinner) run() { + defer s.wg.Done() + ticker := time.NewTicker(spinInterval) + defer ticker.Stop() + for i := 1; ; i++ { + select { + case <-s.done: + return + case t := <-ticker.C: + mu.Lock() + // A console write error here has nowhere better to be reported. + _ = s.stdout.flushStaleLocked(now()) + _ = s.stderr.flushStaleLocked(now()) + s.drawLocked(i, t) + mu.Unlock() + } + } +} + +// drawLocked renders spinner frame i for time t. +func (s *Spinner) drawLocked(i int, t time.Time) { + redrawLocked(fmt.Sprintf("%s... %c [%s]", fitTitle(s.title), frames[i%len(frames)], elapsed(s.start, t))) +} + +// Stop ends the spinner, writing out any unterminated child output line and +// then rendering "done" or, when err is non-nil, "failed". Child output has +// already been shown as it was produced, so nothing else is printed. Stop is +// idempotent. +func (s *Spinner) Stop(err error) { + if s == nil { + return + } + mu.Lock() + if s.stopped { + mu.Unlock() + return + } + s.stopped = true + mu.Unlock() + + close(s.done) + s.wg.Wait() + + mu.Lock() + defer mu.Unlock() + // A console write error here has nowhere better to be reported. + _ = s.stdout.flushLocked() + _ = s.stderr.flushLocked() + status := "done" + if err != nil { + status = "failed" + } + redrawLocked(fmt.Sprintf("%s... %s [%s]", fitTitle(s.title), status, elapsed(s.start, now()))) + endLineLocked() + active = nil +} + +// lineWriter passes child process output through to dst while its owning +// spinner is active. Complete lines are written immediately after clearing +// the spinner line; an unterminated tail is held until its newline arrives, +// it grows past maxPartialLine, it has waited partialLineDelay, or the +// spinner stops, so a spinner redraw never lands in the middle of an +// installer's line that is still being written. Once the owner is no longer +// active, writes go straight to dst. +type lineWriter struct { + owner *Spinner + dst io.Writer + pending []byte + // since is when the oldest byte in pending was written. + since time.Time +} + +// Write passes p through to dst. It reports an error only if writing to dst +// fails, matching the behavior of writing to dst directly. +func (w *lineWriter) Write(p []byte) (int, error) { + mu.Lock() + defer mu.Unlock() + if active != w.owner { + if err := w.flushLocked(); err != nil { + return 0, err + } + return w.dst.Write(p) + } + fresh := len(w.pending) == 0 + w.pending = append(w.pending, p...) + if i := bytes.LastIndexByte(w.pending, '\n'); i >= 0 { + clearLocked() + if _, err := w.dst.Write(w.pending[:i+1]); err != nil { + w.pending = nil + return 0, err + } + w.pending = append(w.pending[:0], w.pending[i+1:]...) + fresh = true + } + if fresh && len(w.pending) > 0 { + w.since = now() + } + if len(w.pending) >= maxPartialLine { + if err := w.flushLocked(); err != nil { + return 0, err + } + } + return len(p), nil +} + +// flushStaleLocked writes out a held unterminated line once it has waited +// partialLineDelay at time t, so an installer that prints a partial line and +// then works silently is still shown promptly. +func (w *lineWriter) flushStaleLocked(t time.Time) error { + if len(w.pending) == 0 || t.Sub(w.since) < partialLineDelay { + return nil + } + return w.flushLocked() +} + +// flushLocked writes out any held unterminated line. It then ends the +// console line on the progress stream so the next spinner frame does not +// overwrite the partial text; the child's own stream is left byte-exact. +func (w *lineWriter) flushLocked() error { + if len(w.pending) == 0 { + return nil + } + clearLocked() + _, err := w.dst.Write(w.pending) + w.pending = nil + endLineLocked() + return err +} diff --git a/progress/progress_test.go b/progress/progress_test.go new file mode 100644 index 0000000..e835c67 --- /dev/null +++ b/progress/progress_test.go @@ -0,0 +1,780 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package progress + +import ( + "bytes" + "errors" + "fmt" + "io" + "os" + "strings" + "testing" + "time" +) + +// fakeClock is a deterministic replacement for the package-level now. +type fakeClock struct { + t time.Time +} + +// now returns the current fake time. +func (c *fakeClock) now() time.Time { return c.t } + +// advance moves the fake time forward by d. +func (c *fakeClock) advance(d time.Duration) { c.t = c.t.Add(d) } + +// setup redirects the package output to fresh buffers, sets enabled and +// installs a fake clock, restoring the package defaults when the test ends. +// The fake clock starts at the real current time so that spinner frames drawn +// from real ticker timestamps still render as 00:00. +func setup(t *testing.T, on bool) (outBuf, stdoutBuf *bytes.Buffer, clock *fakeClock) { + t.Helper() + outBuf, stdoutBuf = &bytes.Buffer{}, &bytes.Buffer{} + clock = &fakeClock{t: time.Now()} + origWidth := termWidth + mu.Lock() + enabled = on + out = outBuf + stdout = stdoutBuf + active = nil + lastLine = "" + now = clock.now + termWidth = func() int { return 0 } + mu.Unlock() + t.Cleanup(func() { + mu.Lock() + defer mu.Unlock() + enabled = false + out = os.Stderr + stdout = os.Stdout + active = nil + lastLine = "" + now = time.Now + termWidth = origWidth + }) + return outBuf, stdoutBuf, clock +} + +// currentLastLen returns the length of the currently rendered line under the package lock. +func currentLastLen() int { + mu.Lock() + defer mu.Unlock() + return len(lastLine) +} + +// currentActive returns active under the package lock. +func currentActive() *Spinner { + mu.Lock() + defer mu.Unlock() + return active +} + +// snapshot returns the contents of buf under the package lock, which makes it +// safe to call while a spinner goroutine may be redrawing. +func snapshot(buf *bytes.Buffer) string { + mu.Lock() + defer mu.Unlock() + return buf.String() +} + +// reset empties buf under the package lock. +func reset(buf *bytes.Buffer) { + mu.Lock() + defer mu.Unlock() + buf.Reset() +} + +// barLine renders the expected bar line for the given progress. +func barLine(eq int, pct int, cur, total string, elapsed string) string { + return fmt.Sprintf("\r|%s%s| %3d%% %s / %s [%s]", + strings.Repeat("=", eq), strings.Repeat("-", barWidth-eq), pct, cur, total, elapsed) +} + +func TestInit(t *testing.T) { + origIsTerminal := isTerminal + t.Cleanup(func() { + isTerminal = origIsTerminal + mu.Lock() + enabled = false + mu.Unlock() + }) + for _, tc := range []struct { + desc string + allow bool + terminal bool + term string + want bool + }{ + {desc: "terminal", allow: true, terminal: true, term: "xterm-256color", want: true}, + {desc: "not allowed", allow: false, terminal: true, term: "xterm-256color", want: false}, + {desc: "not a terminal", allow: true, terminal: false, term: "xterm-256color", want: false}, + {desc: "dumb terminal", allow: true, terminal: true, term: "dumb", want: false}, + {desc: "no TERM", allow: true, terminal: true, term: "", want: true}, + } { + t.Run(tc.desc, func(t *testing.T) { + isTerminal = func(*os.File) bool { return tc.terminal } + t.Setenv("TERM", tc.term) + Init(tc.allow) + if got := Enabled(); got != tc.want { + t.Errorf("Init(%v) with terminal=%v TERM=%q: Enabled() = %v, want %v", tc.allow, tc.terminal, tc.term, got, tc.want) + } + }) + } +} + +func TestDisabled(t *testing.T) { + outBuf, stdoutBuf, _ := setup(t, false) + + if Enabled() { + t.Error("Enabled() = true, want false") + } + b := NewBar("Title", 100, 0) + if b != nil { + t.Errorf("NewBar() = %v, want nil", b) + } + if n, err := b.Write([]byte("12345")); n != 5 || err != nil { + t.Errorf("nil Bar.Write() = %d, %v, want 5, nil", n, err) + } + b.Finish() + b.Abort() + + s := NewSpinner("Title") + if s != nil { + t.Errorf("NewSpinner() = %v, want nil", s) + } + s.Stop(nil) + s.Stop(errors.New("x")) + + if got := outBuf.String(); got != "" { + t.Errorf("out = %q, want empty", got) + } + + Printf("hello %d\n", 1) + if got, want := stdoutBuf.String(), "hello 1\n"; got != want { + t.Errorf("stdout after Printf = %q, want %q", got, want) + } + if got := outBuf.String(); got != "" { + t.Errorf("out after Printf = %q, want empty", got) + } +} + +func TestBarKnownTotal(t *testing.T) { + outBuf, _, clock := setup(t, true) + + b := NewBar("Title", 100, 0) + if b == nil { + t.Fatal("NewBar() = nil, want non-nil") + } + if got, want := outBuf.String(), "Title (100 B)\n"+barLine(0, 0, "0 B", "100 B", "00:00"); got != want { + t.Errorf("initial output = %q, want %q", got, want) + } + + outBuf.Reset() + clock.advance(redrawInterval) + if n, err := b.Write(make([]byte, 50)); n != 50 || err != nil { + t.Errorf("Write() = %d, %v, want 50, nil", n, err) + } + if got, want := outBuf.String(), barLine(17, 50, "50 B", "100 B", "00:00"); got != want { + t.Errorf("output at 50%% = %q, want %q", got, want) + } + if !strings.Contains(outBuf.String(), " 50%") { + t.Errorf("output at 50%% = %q, want it to contain %q", outBuf.String(), " 50%") + } + + outBuf.Reset() + clock.advance(time.Second) + b.Finish() + if got, want := outBuf.String(), barLine(35, 100, "100 B", "100 B", "00:01")+"\n"; got != want { + t.Errorf("output after Finish = %q, want %q", got, want) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen after Finish = %d, want 0", got) + } +} + +func TestBarFinishCapsOverrun(t *testing.T) { + outBuf, _, clock := setup(t, true) + + b := NewBar("Title", 100, 0) + clock.advance(redrawInterval) + outBuf.Reset() + b.Write(make([]byte, 150)) + if got, want := outBuf.String(), barLine(35, 100, "100 B", "100 B", "00:00"); got != want { + t.Errorf("output after overrun Write = %q, want %q", got, want) + } + outBuf.Reset() + b.Finish() + // The final frame is already on the line, so Finish only terminates it. + if got, want := outBuf.String(), "\n"; got != want { + t.Errorf("output after Finish = %q, want %q", got, want) + } +} + +func TestBarResumed(t *testing.T) { + outBuf, _, _ := setup(t, true) + + if b := NewBar("Title", 100, 40); b == nil { + t.Fatal("NewBar() = nil, want non-nil") + } + if got, want := outBuf.String(), "Title (100 B)\n"+barLine(14, 40, "40 B", "100 B", "00:00"); got != want { + t.Errorf("initial output = %q, want %q", got, want) + } +} + +func TestBarUnknownTotal(t *testing.T) { + for _, total := range []int64{0, -1} { + t.Run(fmt.Sprint(total), func(t *testing.T) { + outBuf, _, clock := setup(t, true) + + b := NewBar("Title", total, 0) + if b == nil { + t.Fatal("NewBar() = nil, want non-nil") + } + if got, want := outBuf.String(), "Title\n\r 0 B [00:00]"; got != want { + t.Errorf("initial output = %q, want %q", got, want) + } + + outBuf.Reset() + clock.advance(2 * time.Second) + b.Write(make([]byte, 2048)) + if got, want := outBuf.String(), "\r 2.0 KiB [00:02]"; got != want { + t.Errorf("output after Write = %q, want %q", got, want) + } + + outBuf.Reset() + clock.advance(time.Second) + b.Finish() + got := outBuf.String() + if want := "\r 2.0 KiB [00:03]\n"; got != want { + t.Errorf("output after Finish = %q, want %q", got, want) + } + if strings.Contains(got, "|") { + t.Errorf("output after Finish = %q, want no bar cells", got) + } + }) + } +} + +func TestBarThrottle(t *testing.T) { + outBuf, _, clock := setup(t, true) + + b := NewBar("Title", 100, 0) + outBuf.Reset() + + // A write immediately after the initial draw is within the interval. + b.Write(make([]byte, 10)) + if got := strings.Count(outBuf.String(), "\r"); got != 0 { + t.Errorf("redraws immediately after NewBar = %d, want 0", got) + } + + // Once the interval has passed, only the first of two writes redraws. + clock.advance(redrawInterval) + b.Write(make([]byte, 10)) + b.Write(make([]byte, 10)) + if got := strings.Count(outBuf.String(), "\r"); got != 1 { + t.Errorf("redraws within one interval = %d, want 1", got) + } + // The throttled write still counted its bytes. + if got, want := outBuf.String(), barLine(7, 20, "20 B", "100 B", "00:00"); got != want { + t.Errorf("output = %q, want %q", got, want) + } + + // After another interval, the accumulated total is drawn. + outBuf.Reset() + clock.advance(redrawInterval) + b.Write(make([]byte, 10)) + if got, want := outBuf.String(), barLine(14, 40, "40 B", "100 B", "00:00"); got != want { + t.Errorf("output = %q, want %q", got, want) + } +} + +func TestBarAbort(t *testing.T) { + outBuf, _, clock := setup(t, true) + + b := NewBar("Title", 100, 0) + clock.advance(redrawInterval) + b.Write(make([]byte, 50)) + outBuf.Reset() + + b.Abort() + if got, want := outBuf.String(), "\n"; got != want { + t.Errorf("output after Abort = %q, want %q", got, want) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen after Abort = %d, want 0", got) + } +} + +func TestRedrawPadding(t *testing.T) { + outBuf, _, _ := setup(t, true) + + mu.Lock() + redrawLocked("abcdef") + redrawLocked("ab") + mu.Unlock() + if got, want := outBuf.String(), "\rabcdef\rab "; got != want { + t.Errorf("output = %q, want %q", got, want) + } + if got := currentLastLen(); got != 2 { + t.Errorf("lastLen = %d, want 2", got) + } + + // A longer line needs no padding. + outBuf.Reset() + mu.Lock() + redrawLocked("abcd") + mu.Unlock() + if got, want := outBuf.String(), "\rabcd"; got != want { + t.Errorf("output = %q, want %q", got, want) + } + + // Redrawing the same content writes nothing; serial console loggers + // record every frame, so identical frames are suppressed. + outBuf.Reset() + mu.Lock() + redrawLocked("abcd") + mu.Unlock() + if got := outBuf.String(); got != "" { + t.Errorf("output after identical redraw = %q, want empty", got) + } +} + +func TestClearLocked(t *testing.T) { + outBuf, _, _ := setup(t, true) + + mu.Lock() + clearLocked() + mu.Unlock() + if got := outBuf.String(); got != "" { + t.Errorf("output with nothing rendered = %q, want empty", got) + } + + mu.Lock() + redrawLocked("abc") + clearLocked() + mu.Unlock() + if got, want := outBuf.String(), "\rabc\r \r"; got != want { + t.Errorf("output = %q, want %q", got, want) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen = %d, want 0", got) + } +} + +func TestSpinnerDone(t *testing.T) { + outBuf, _, _ := setup(t, true) + + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + if got := currentActive(); got != s { + t.Errorf("active = %p, want %p", got, s) + } + if got, want := snapshot(outBuf), "\rTitle... - [00:00]"; !strings.HasPrefix(got, want) { + t.Errorf("initial output = %q, want prefix %q", got, want) + } + + s.Stop(nil) + if got, want := outBuf.String(), "\rTitle... done [00:00]\n"; !strings.HasSuffix(got, want) { + t.Errorf("output after Stop(nil) = %q, want suffix %q", got, want) + } + if got := currentActive(); got != nil { + t.Errorf("active after Stop = %p, want nil", got) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen after Stop = %d, want 0", got) + } + + // Stop is idempotent. + outBuf.Reset() + s.Stop(nil) + s.Stop(errors.New("x")) + if got := outBuf.String(); got != "" { + t.Errorf("output after repeated Stop = %q, want empty", got) + } +} + +// redirectChild points the spinner's pass-through writers at fresh buffers +// and returns them. +func redirectChild(s *Spinner) (childOut, childErr *bytes.Buffer) { + childOut, childErr = &bytes.Buffer{}, &bytes.Buffer{} + mu.Lock() + defer mu.Unlock() + s.stdout.dst = childOut + s.stderr.dst = childErr + return childOut, childErr +} + +func TestSpinnerPassesChildOutputThrough(t *testing.T) { + for _, tc := range []struct { + name string + err error + wantEnd string + }{ + {name: "success", err: nil, wantEnd: "\rTitle... done [00:00]\n"}, + {name: "failure", err: errors.New("x"), wantEnd: "\rTitle... failed [00:00]\n"}, + } { + t.Run(tc.name, func(t *testing.T) { + outBuf, _, _ := setup(t, true) + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + childOut, childErr := redirectChild(s) + + // An unterminated line is held so a spinner frame cannot split it. + io.WriteString(Stdout(), "installing") + if got := snapshot(childOut); got != "" { + t.Errorf("child stdout after partial write = %q, want empty", got) + } + // Completing the line clears the spinner and writes it immediately. + n := currentLastLen() + io.WriteString(Stdout(), " step 1\nstep 2") + if got, want := snapshot(childOut), "installing step 1\n"; got != want { + t.Errorf("child stdout after newline = %q, want %q", got, want) + } + if n > 0 { + if got, want := snapshot(outBuf), "\r"+strings.Repeat(" ", n)+"\r"; !strings.Contains(got, want) { + t.Errorf("progress output = %q, want it to contain clear sequence %q", got, want) + } + } + // Stderr stays on stderr. + io.WriteString(Stderr(), "warning: something\n") + if got, want := snapshot(childErr), "warning: something\n"; got != want { + t.Errorf("child stderr = %q, want %q", got, want) + } + if strings.Contains(snapshot(outBuf), "step") || strings.Contains(snapshot(outBuf), "warning") { + t.Errorf("progress output = %q, want no child output on the progress stream", snapshot(outBuf)) + } + + // Stop writes the held tail before the final status, success or not. + s.Stop(tc.err) + if got, want := childOut.String(), "installing step 1\nstep 2"; got != want { + t.Errorf("child stdout after Stop = %q, want %q", got, want) + } + if got := outBuf.String(); !strings.HasSuffix(got, "\n"+tc.wantEnd) { + t.Errorf("progress output after Stop = %q, want suffix %q", got, "\n"+tc.wantEnd) + } + + // A late write from a writer handed out earlier goes straight through. + io.WriteString(s.stdout, "late") + if got, want := childOut.String(), "installing step 1\nstep 2late"; got != want { + t.Errorf("child stdout after late write = %q, want %q", got, want) + } + }) + } +} + +func TestSpinnerLongPartialLine(t *testing.T) { + setup(t, true) + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + defer s.Stop(nil) + childOut, _ := redirectChild(s) + + long := strings.Repeat("x", maxPartialLine) + io.WriteString(Stdout(), long) + if got := snapshot(childOut); got != long { + t.Errorf("child stdout after %d unterminated bytes has %d bytes, want all of them", len(long), len(got)) + } +} + +type errWriter struct{} + +func (errWriter) Write([]byte) (int, error) { return 0, errors.New("closed") } + +func TestSpinnerPassThroughWriteError(t *testing.T) { + setup(t, true) + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + defer s.Stop(nil) + mu.Lock() + s.stdout.dst = errWriter{} + mu.Unlock() + + if _, err := io.WriteString(Stdout(), "line\n"); err == nil { + t.Error("Write() to a failing destination = nil error, want error") + } +} + +func TestSpinnerExclusive(t *testing.T) { + setup(t, true) + + s := NewSpinner("First") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + defer s.Stop(nil) + if s2 := NewSpinner("Second"); s2 != nil { + s2.Stop(nil) + t.Errorf("NewSpinner() while another is active = %v, want nil", s2) + } + s.Stop(nil) + s3 := NewSpinner("Third") + if s3 == nil { + t.Fatal("NewSpinner() after Stop = nil, want non-nil") + } + s3.Stop(nil) +} + +func TestStdoutStderr(t *testing.T) { + setup(t, true) + + if w := Stdout(); w != io.Writer(os.Stdout) { + t.Errorf("Stdout() with no spinner = %v, want os.Stdout", w) + } + if w := Stderr(); w != io.Writer(os.Stderr) { + t.Errorf("Stderr() with no spinner = %v, want os.Stderr", w) + } + + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + if w := Stdout(); w != io.Writer(s.stdout) { + t.Errorf("Stdout() while active = %v, want the spinner's stdout pass-through", w) + } + if w := Stderr(); w != io.Writer(s.stderr) { + t.Errorf("Stderr() while active = %v, want the spinner's stderr pass-through", w) + } + s.Stop(nil) + + if w := Stdout(); w != io.Writer(os.Stdout) { + t.Errorf("Stdout() after Stop = %v, want os.Stdout", w) + } + if w := Stderr(); w != io.Writer(os.Stderr) { + t.Errorf("Stderr() after Stop = %v, want os.Stderr", w) + } +} + +func TestPrintfClearsBar(t *testing.T) { + outBuf, stdoutBuf, _ := setup(t, true) + + NewBar("Title", 100, 0) + n := currentLastLen() + if n == 0 { + t.Fatal("lastLen after NewBar = 0, want > 0") + } + outBuf.Reset() + + Printf("hello %d\n", 1) + if got, want := outBuf.String(), "\r"+strings.Repeat(" ", n)+"\r"; got != want { + t.Errorf("out after Printf = %q, want %q", got, want) + } + if got, want := stdoutBuf.String(), "hello 1\n"; got != want { + t.Errorf("stdout after Printf = %q, want %q", got, want) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen after Printf = %d, want 0", got) + } +} + +func TestPrintfClearsSpinner(t *testing.T) { + outBuf, stdoutBuf, _ := setup(t, true) + + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + n := currentLastLen() + if n == 0 { + t.Fatal("lastLen after NewSpinner = 0, want > 0") + } + reset(outBuf) + + // The spinner goroutine may redraw before or after Printf, so only assert + // that the clear sequence was emitted and the text went to stdout. The + // exact clear-then-lastLen==0 behavior is checked deterministically by + // TestPrintfClearsBar and TestClearLocked. + Printf("hello %d\n", 2) + got := snapshot(outBuf) + if want := "\r" + strings.Repeat(" ", n) + "\r"; !strings.Contains(got, want) { + t.Errorf("out after Printf = %q, want it to contain %q", got, want) + } + if strings.Contains(got, "hello") { + t.Errorf("out after Printf = %q, want no console text", got) + } + if got, want := stdoutBuf.String(), "hello 2\n"; got != want { + t.Errorf("stdout after Printf = %q, want %q", got, want) + } + + s.Stop(nil) + if got, want := outBuf.String(), "\rTitle... done [00:00]\n"; !strings.HasSuffix(got, want) { + t.Errorf("output after Stop = %q, want suffix %q", got, want) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen after Stop = %d, want 0", got) + } +} + +func TestElapsed(t *testing.T) { + start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + for _, tc := range []struct { + d time.Duration + want string + }{ + {0, "00:00"}, + {400 * time.Millisecond, "00:00"}, + {600 * time.Millisecond, "00:01"}, + {61 * time.Second, "01:01"}, + {-5 * time.Second, "00:00"}, + {3599 * time.Second, "59:59"}, + {3600 * time.Second, "60:00"}, + } { + if got := elapsed(start, start.Add(tc.d)); got != tc.want { + t.Errorf("elapsed(%v) = %q, want %q", tc.d, got, tc.want) + } + } +} + +func TestFitTitle(t *testing.T) { + orig := termWidth + t.Cleanup(func() { termWidth = orig }) + for _, tc := range []struct { + width int + title string + want string + }{ + {0, "unknown width keeps the title", "unknown width keeps the title"}, + {80, "short", "short"}, + {30, "googet-package-with-a-long-name", "googet-pack"}, + {30, "ünïcödé-päckägé-nämé", "ünïcödé-päc"}, + {10, "too narrow to help", "too narrow to help"}, + } { + termWidth = func() int { return tc.width } + if got := fitTitle(tc.title); got != tc.want { + t.Errorf("fitTitle(%q) at width %d = %q, want %q", tc.title, tc.width, got, tc.want) + } + } +} + +func TestSpinnerFitsTerminalWidth(t *testing.T) { + outBuf, _, _ := setup(t, true) + const width = 30 + mu.Lock() + termWidth = func() int { return width } + mu.Unlock() + + NewSpinner("Installing googet-package-with-a-long-name").Stop(errors.New("boom")) + + got := outBuf.String() + for _, line := range strings.FieldsFunc(got, func(r rune) bool { return r == '\r' || r == '\n' }) { + if len(line) >= width { + t.Errorf("rendered line %q is %d columns, want fewer than %d", line, len(line), width) + } + } + if want := "\rInstalling... failed [00:00]\n"; !strings.HasSuffix(got, want) { + t.Errorf("output = %q, want suffix %q", got, want) + } +} + +func TestFitBarWidth(t *testing.T) { + orig := termWidth + t.Cleanup(func() { termWidth = orig }) + for _, tc := range []struct{ width, want int }{ + {0, barWidth}, + {200, barWidth}, + {72, barWidth}, + {60, 23}, + {47, minBarWidth}, + {46, 0}, + {20, 0}, + } { + termWidth = func() int { return tc.width } + if got := fitBarWidth(); got != tc.want { + t.Errorf("fitBarWidth() at width %d = %d, want %d", tc.width, got, tc.want) + } + } +} + +func TestBarFitsTerminalWidth(t *testing.T) { + for _, tc := range []struct { + width int + want string + }{ + {60, "\r|=======================| 100% 1023 KiB / 1023 KiB [00:00]\n"}, + {40, "\r 100% 1023 KiB / 1023 KiB [00:00]\n"}, + } { + t.Run(fmt.Sprint(tc.width), func(t *testing.T) { + outBuf, _, _ := setup(t, true) + mu.Lock() + termWidth = func() int { return tc.width } + mu.Unlock() + + b := NewBar("Downloading pkg", 1023<<10, 0) + b.Write(make([]byte, 512<<10)) + b.Finish() + + got := outBuf.String() + for _, line := range strings.FieldsFunc(got, func(r rune) bool { return r == '\r' || r == '\n' }) { + if len(line) >= tc.width { + t.Errorf("rendered line %q is %d columns, want fewer than %d", line, len(line), tc.width) + } + } + if !strings.HasSuffix(got, tc.want) { + t.Errorf("output = %q, want suffix %q", got, tc.want) + } + }) + } +} + +func TestSpinnerFlushesStalePartialLine(t *testing.T) { + outBuf, _, clock := setup(t, true) + s := NewSpinner("Title") + if s == nil { + t.Fatal("NewSpinner() = nil, want non-nil") + } + defer s.Stop(nil) + childOut, _ := redirectChild(s) + flushStale := func() { + mu.Lock() + defer mu.Unlock() + if err := s.stdout.flushStaleLocked(now()); err != nil { + t.Errorf("flushStaleLocked() = %v, want nil", err) + } + } + + io.WriteString(Stdout(), "Extracting files... ") + clock.advance(partialLineDelay - time.Millisecond) + flushStale() + if got := snapshot(childOut); got != "" { + t.Errorf("child stdout before partialLineDelay = %q, want empty", got) + } + + // Appending more text does not restart the delay. + io.WriteString(Stdout(), "50%") + clock.advance(time.Millisecond) + flushStale() + if got, want := snapshot(childOut), "Extracting files... 50%"; got != want { + t.Errorf("child stdout after partialLineDelay = %q, want %q", got, want) + } + if got := currentLastLen(); got != 0 { + t.Errorf("lastLen after stale flush = %d, want 0 so the next frame starts a new line", got) + } + + // The rest of the line is still passed through byte-exact. + io.WriteString(Stdout(), " done\n") + if got, want := snapshot(childOut), "Extracting files... 50% done\n"; got != want { + t.Errorf("child stdout after newline = %q, want %q", got, want) + } + if strings.Contains(snapshot(outBuf), "Extracting") { + t.Errorf("progress output = %q, want no child output on the progress stream", snapshot(outBuf)) + } +} diff --git a/settings/settings.go b/settings/settings.go index 59a0bec..9e87ca5 100644 --- a/settings/settings.go +++ b/settings/settings.go @@ -31,6 +31,10 @@ var ( AllowUnsafeURL bool // StrictConflicts enables strict enforcement of file ownership conflicts. StrictConflicts bool + // Progress enables the download progress bar and install spinner on + // interactive terminals; set from googet.conf (default true) and + // overridden by an explicit -progress flag. + Progress = true ) // Initialize reads the initial settings. @@ -84,6 +88,8 @@ type conf struct { ProxyServer string AllowUnsafeURL bool StrictConflicts bool + // Progress is a pointer so an absent key keeps the default of true. + Progress *bool } // unmarshalConfFile returns a conf from a YAML configuration file. @@ -142,4 +148,8 @@ func readConf(filename string) { AllowUnsafeURL = gc.AllowUnsafeURL StrictConflicts = gc.StrictConflicts + Progress = true + if gc.Progress != nil { + Progress = *gc.Progress + } } diff --git a/settings/settings_test.go b/settings/settings_test.go index 9942131..505137c 100644 --- a/settings/settings_test.go +++ b/settings/settings_test.go @@ -16,7 +16,7 @@ func TestInitialize(t *testing.T) { if err != nil { t.Fatalf("error creating conf file: %v", err) } - content := []byte("archs: [noarch, x86_64, arm64]\ncachelife: 10m\nlockfilemaxage: 1x\nallowunsafeurl: true") + content := []byte("archs: [noarch, x86_64, arm64]\ncachelife: 10m\nlockfilemaxage: 1x\nallowunsafeurl: true\nprogress: false") if _, err := f.Write(content); err != nil { t.Fatalf("error writing conf file: %v", err) } @@ -57,4 +57,31 @@ func TestInitialize(t *testing.T) { t.Errorf("settings.AllowUnsafeURL got: %v, want: %v", got, wantAllowUnsafeURL) } }) + + t.Run("Parsing Progress", func(t *testing.T) { + if got, want := settings.Progress, false; got != want { + t.Errorf("settings.Progress got: %v, want: %v", got, want) + } + }) +} + +func TestProgressDefault(t *testing.T) { + rootDir := t.TempDir() + conf := filepath.Join(rootDir, "googet.conf") + if err := os.WriteFile(conf, []byte("progress: false"), 0644); err != nil { + t.Fatalf("error writing conf file: %v", err) + } + settings.Initialize(rootDir, true) + if settings.Progress { + t.Fatalf("settings.Progress with progress: false = true, want false") + } + + // An absent key restores the default rather than keeping the prior value. + if err := os.WriteFile(conf, []byte("cachelife: 10m"), 0644); err != nil { + t.Fatalf("error writing conf file: %v", err) + } + settings.Initialize(rootDir, true) + if !settings.Progress { + t.Errorf("settings.Progress with no progress key = false, want true") + } }