Skip to content
Merged
49 changes: 39 additions & 10 deletions download/download.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down Expand Up @@ -84,32 +85,55 @@ 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
}
if _, err := f.Seek(0, 0); err != nil {
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 {
return err
}
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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand Down
140 changes: 132 additions & 8 deletions download/download_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand All @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand All @@ -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)
}
Expand Down
1 change: 1 addition & 0 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -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=
Expand Down
27 changes: 27 additions & 0 deletions googet.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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 {
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
}
Loading
Loading