diff --git a/gomodfs.go b/gomodfs.go index 4a412f0..c426dcd 100644 --- a/gomodfs.go +++ b/gomodfs.go @@ -57,8 +57,16 @@ type FS struct { // ModuleProxyURL is the URL of the Go module proxy to use. // If empty, "https://proxy.golang.org" is used. // It should not have a trailing slash. + // It is ignored if ModuleProxyURLs is non-empty. ModuleProxyURL string + // ModuleProxyURLs optionally specifies an ordered list of Go + // module proxy URLs. Each download is attempted against each + // proxy in order, moving on to the next after any error, whether + // a network error or a non-200 HTTP status. The URLs should not + // have trailing slashes. If empty, ModuleProxyURL is used. + ModuleProxyURLs []string + Logf func(format string, args ...any) // if non-nil, alternate logger to use Verbose bool @@ -117,74 +125,77 @@ func (fs *FS) client() *http.Client { return cmp.Or(fs.Client, http.DefaultClient) } -func (fs *FS) moduleProxyURL() string { +func (fs *FS) moduleProxyURLs() []string { + if len(fs.ModuleProxyURLs) > 0 { + urls := make([]string, len(fs.ModuleProxyURLs)) + for i, u := range fs.ModuleProxyURLs { + urls[i] = strings.TrimSuffix(u, "/") + } + return urls + } if fs.ModuleProxyURL != "" { - return strings.TrimSuffix(fs.ModuleProxyURL, "/") + return []string{strings.TrimSuffix(fs.ModuleProxyURL, "/")} } - return "https://proxy.golang.org" + return []string{"https://proxy.golang.org"} } -func (fs *FS) modURLBase(mv store.ModuleVersion) (string, error) { +// modURLBases returns the "//@v/" URL prefix +// of the given module version for each configured module proxy, in +// the order they should be tried. +func (fs *FS) modURLBases(mv store.ModuleVersion) ([]string, error) { escMod, err := module.EscapePath(mv.Module) if err != nil { - return "", fmt.Errorf("failed to escape module name %q: %w", mv.Module, err) + return nil, fmt.Errorf("failed to escape module name %q: %w", mv.Module, err) } escVer, err := module.EscapeVersion(mv.Version) if err != nil { - return "", fmt.Errorf("failed to escape version %q: %w", mv.Version, err) + return nil, fmt.Errorf("failed to escape version %q: %w", mv.Version, err) + } + proxies := fs.moduleProxyURLs() + bases := make([]string, len(proxies)) + for i, p := range proxies { + bases[i] = p + "/" + escMod + "/@v/" + escVer } - return fs.moduleProxyURL() + "/" + escMod + "/@v/" + escVer, nil + return bases, nil } -func (fs *FS) downloadModFile(ctx context.Context, mv store.ModuleVersion) (_ []byte, err error) { - sp := fs.Stats.StartSpan("download-mod-file") - defer func() { sp.End(err) }() - - ctx = context.Background() // TODO(bradfitz): make a singleflight variant that refcounts context lifetime - - vi, err, _ := fs.sf.Do("download-mod:"+mv.Module+"@"+mv.Version, func() (any, error) { - urlBase, err := fs.modURLBase(mv) - if err != nil { - return nil, err - } - urlStr := urlBase + ".mod" +func (fs *FS) downloadModFile(ctx context.Context, mv store.ModuleVersion) ([]byte, error) { + return fs.downloadMetaFile(ctx, mv, "mod", fs.Store.PutModFile) +} - data, err := fs.netSlurp(ctx, urlStr) - if err != nil { - return nil, fmt.Errorf("failed to download %q: %w", urlStr, err) - } - if err := fs.Store.PutModFile(ctx, mv, data); err != nil { - return nil, fmt.Errorf("failed to store mod file for %q: %w", mv, err) - } - return data, nil - }) - if err != nil { - return nil, err - } - return vi.([]byte), nil +func (fs *FS) downloadInfoFile(ctx context.Context, mv store.ModuleVersion) ([]byte, error) { + return fs.downloadMetaFile(ctx, mv, "info", fs.Store.PutInfoFile) } -func (fs *FS) downloadInfoFile(ctx context.Context, mv store.ModuleVersion) (_ []byte, err error) { - sp := fs.Stats.StartSpan("download-info-file") +// downloadMetaFile downloads the ".mod" or ".info" file (per ext) of +// the given module version from the first configured module proxy +// that can serve it and stores it with put. +func (fs *FS) downloadMetaFile(ctx context.Context, mv store.ModuleVersion, ext string, put func(context.Context, store.ModuleVersion, []byte) error) (_ []byte, err error) { + sp := fs.Stats.StartSpan("download-" + ext + "-file") defer func() { sp.End(err) }() ctx = context.Background() // TODO(bradfitz): make a singleflight variant that refcounts context lifetime - vi, err, _ := fs.sf.Do("download-info:"+mv.Module+"@"+mv.Version, func() (any, error) { - urlBase, err := fs.modURLBase(mv) + vi, err, _ := fs.sf.Do("download-"+ext+":"+mv.Module+"@"+mv.Version, func() (any, error) { + urlBases, err := fs.modURLBases(mv) if err != nil { return nil, err } - urlStr := urlBase + ".info" + var errs []error + for _, urlBase := range urlBases { + urlStr := urlBase + "." + ext - data, err := fs.netSlurp(ctx, urlStr) - if err != nil { - return nil, fmt.Errorf("failed to download %q: %w", urlStr, err) - } - if err := fs.Store.PutInfoFile(ctx, mv, data); err != nil { - return nil, fmt.Errorf("failed to store info file for %q: %w", mv, err) + data, err := fs.netSlurp(ctx, urlStr) + if err != nil { + errs = append(errs, fmt.Errorf("failed to download %q: %w", urlStr, err)) + continue + } + if err := put(ctx, mv, data); err != nil { + return nil, fmt.Errorf("failed to store %s file for %q: %w", ext, mv, err) + } + return data, nil } - return data, nil + return nil, errors.Join(errs...) }) if err != nil { return nil, err @@ -208,21 +219,26 @@ func (fs *FS) logf(format string, arg ...any) { } func (fs *FS) downloadZip(ctx context.Context, mv store.ModuleVersion) (store.ModHandle, error) { - baseURL, err := fs.modURLBase(mv) + baseURLs, err := fs.modURLBases(mv) if err != nil { return nil, err } - download := map[string][]byte{} // extension (zip, info, mod) -> data - for _, ext := range []string{"zip", "info", "mod"} { - urlStr := baseURL + "." + ext - sp := fs.Stats.StartSpan("net-downloadZip-ext-" + ext) - data, err := fs.netSlurp(ctx, urlStr) - sp.End(err) + // A module version's zip, info, and mod files must all come from + // the same proxy, so any error moves the whole set to the next + // proxy rather than mixing sources. + var download map[string][]byte // extension (zip, info, mod) -> data + var errs []error + for _, baseURL := range baseURLs { + download, err = fs.downloadZipSet(ctx, baseURL) if err != nil { - return nil, fmt.Errorf("failed to download %q: %w", urlStr, err) + errs = append(errs, err) + continue } - download[ext] = data + break + } + if download == nil { + return nil, errors.Join(errs...) } zr, err := zip.NewReader(bytes.NewReader(download["zip"]), int64(len(download["zip"]))) @@ -265,6 +281,23 @@ func (fs *FS) downloadZip(ctx context.Context, mv store.ModuleVersion) (store.Mo return fs.Store.PutModule(ctx, mv, put) } +// downloadZipSet downloads a module version's zip, info, and mod +// files from the module proxy URL prefix baseURL. +func (fs *FS) downloadZipSet(ctx context.Context, baseURL string) (map[string][]byte, error) { + download := map[string][]byte{} // extension (zip, info, mod) -> data + for _, ext := range []string{"zip", "info", "mod"} { + urlStr := baseURL + "." + ext + sp := fs.Stats.StartSpan("net-downloadZip-ext-" + ext) + data, err := fs.netSlurp(ctx, urlStr) + sp.End(err) + if err != nil { + return nil, fmt.Errorf("failed to download %q: %w", urlStr, err) + } + download[ext] = data + } + return download, nil +} + type putFile struct { path string zf *zip.File diff --git a/gomodfs_test.go b/gomodfs_test.go index b2528a8..0bc8e0e 100644 --- a/gomodfs_test.go +++ b/gomodfs_test.go @@ -70,3 +70,115 @@ entry[2]: upload_file.txt, -rw-r--r--, size=9 t.Fatalf("bad directory entries; got:\n%s\nwant:\n%s", got, want) } } + +// hostRecordingTransport wraps an http.RoundTripper, recording the +// hosts of attempted requests. +type hostRecordingTransport struct { + inner http.RoundTripper + hosts []string +} + +func (t *hostRecordingTransport) RoundTrip(r *http.Request) (*http.Response, error) { + t.hosts = append(t.hosts, r.URL.Host) + return t.inner.RoundTrip(r) +} + +func TestModuleProxyFallback(t *testing.T) { + mv := store.ModuleVersion{ + Module: "github.com/bramvdbogaerde/go-scp", + Version: "v1.4.0", + } + const wantZipHash = "h1:jKMwpwCbcX1KyvDbm/PDJuXcMuNVlLGi0Q0reuzjyKY=" + + newFS := func(t *testing.T, proxies ...string) (*FS, *hostRecordingTransport) { + st := &gitstore.Storage{GitRepo: testGitDir(t)} + addStopGitStoreCleanup(t, st) + tr := &hostRecordingTransport{inner: testDataTransport{}} + return &FS{ + Store: st, + Client: &http.Client{Transport: tr}, + ModuleProxyURLs: proxies, + Logf: t.Logf, + }, tr + } + + t.Run("fall-through-to-second", func(t *testing.T) { + // testDataTransport errors on any URL it doesn't know, + // simulating an unreachable first proxy. + fs, tr := newFS(t, "https://bad.proxy.example.com", "https://proxy.golang.org") + ctx := t.Context() + + mh, err := fs.downloadZip(ctx, mv) + if err != nil { + t.Fatalf("downloadZip: %v", err) + } + zipHash, err := fs.Store.GetZipHash(ctx, mh) + if err != nil { + t.Fatalf("GetZipHash: %v", err) + } + if g := string(zipHash); g != wantZipHash { + t.Fatalf("zip hash = %q; want %q", g, wantZipHash) + } + if len(tr.hosts) < 2 || tr.hosts[0] != "bad.proxy.example.com" || tr.hosts[1] != "proxy.golang.org" { + t.Fatalf("request hosts = %q; want the bad proxy attempted first, then proxy.golang.org", tr.hosts) + } + + if _, err := fs.downloadModFile(ctx, mv); err != nil { + t.Fatalf("downloadModFile: %v", err) + } + if _, err := fs.downloadInfoFile(ctx, mv); err != nil { + t.Fatalf("downloadInfoFile: %v", err) + } + }) + + t.Run("first-success-stops", func(t *testing.T) { + fs, tr := newFS(t, "https://proxy.golang.org", "https://bad.proxy.example.com") + if _, err := fs.downloadZip(t.Context(), mv); err != nil { + t.Fatalf("downloadZip: %v", err) + } + for _, h := range tr.hosts { + if h != "proxy.golang.org" { + t.Fatalf("request to %q; the second proxy should never be attempted", h) + } + } + }) + + t.Run("all-fail", func(t *testing.T) { + fs, _ := newFS(t, "https://bad1.example.com", "https://bad2.example.com") + _, err := fs.downloadZip(t.Context(), mv) + if err == nil { + t.Fatal("downloadZip succeeded; want error") + } + for _, want := range []string{"bad1.example.com", "bad2.example.com"} { + if !strings.Contains(err.Error(), want) { + t.Errorf("error %q does not mention %q", err, want) + } + } + }) +} + +func TestModuleProxyURLs(t *testing.T) { + tests := []struct { + name string + fs *FS + want []string + }{ + {"default", &FS{}, []string{"https://proxy.golang.org"}}, + {"single", &FS{ModuleProxyURL: "https://a/"}, []string{"https://a"}}, + {"list", &FS{ModuleProxyURLs: []string{"https://a/", "https://b"}}, []string{"https://a", "https://b"}}, + {"list-wins", &FS{ModuleProxyURL: "https://c", ModuleProxyURLs: []string{"https://a"}}, []string{"https://a"}}, + } + for _, tt := range tests { + got := tt.fs.moduleProxyURLs() + if len(got) != len(tt.want) { + t.Errorf("%s: got %q; want %q", tt.name, got, tt.want) + continue + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("%s: got %q; want %q", tt.name, got, tt.want) + break + } + } + } +}