diff --git a/appview/config/config.go b/appview/config/config.go index 202f81b2..73e29d0b 100644 --- a/appview/config/config.go +++ b/appview/config/config.go @@ -69,7 +69,8 @@ type PlcConfig struct { } type KnotMirrorConfig struct { - Url string `env:"URL, default=https://mirror.tangled.network"` + Url string `env:"URL, default=https://mirror.tangled.network"` + ArchiveHeaderTimeout time.Duration `env:"ARCHIVE_HEADER_TIMEOUT, default=60s"` } type JetstreamConfig struct { diff --git a/appview/pages/templates/repo/fragments/cloneDropdown.html b/appview/pages/templates/repo/fragments/cloneDropdown.html index 20abdc8c..5dd22e5b 100644 --- a/appview/pages/templates/repo/fragments/cloneDropdown.html +++ b/appview/pages/templates/repo/fragments/cloneDropdown.html @@ -68,14 +68,14 @@
{{ i "download" "w-4 h-4" }} Download tar.gz {{ i "download" "w-4 h-4" }} diff --git a/appview/repo/archive.go b/appview/repo/archive.go index c60d2444..a035a359 100644 --- a/appview/repo/archive.go +++ b/appview/repo/archive.go @@ -6,63 +6,77 @@ import ( "net/http" "net/url" "strings" + "time" "github.com/go-chi/chi/v5" + "github.com/samber/lo" "tangled.org/core/api/tangled" + "tangled.org/core/appview/models" + "tangled.org/core/gitutil" ) +const archiveRoute = "/archive/*" + +func newArchiveClient(headerTimeout time.Duration) *http.Client { + transport := http.DefaultTransport.(*http.Transport).Clone() + transport.ResponseHeaderTimeout = headerTimeout + return &http.Client{Transport: transport} +} + func (rp *Repo) DownloadArchive(w http.ResponseWriter, r *http.Request) { l := rp.logger.With("handler", "DownloadArchive") - ref := chi.URLParam(r, "ref") - ref, _ = url.PathUnescape(ref) - format := r.URL.Query().Get("format") - ref, format = archiveRefAndFormat(ref, format, r.UserAgent()) + fail := func(status int) { + w.WriteHeader(status) + lo.Ternary(status == http.StatusServiceUnavailable, rp.pages.Error503, rp.pages.Error404)(w) + } + + params, err := parseArchiveRequest(r) + if err != nil { + l.Warn("rejecting archive request", "err", err) + fail(http.StatusNotFound) + return + } + f, err := rp.repoResolver.Resolve(r) if err != nil { l.Error("failed to get repo and knot", "err", err) + fail(http.StatusNotFound) return } + name := gitutil.RepoName(f.Slug()) + params.Prefix = params.Prefix.OrDefault(name, params.Rev) + // build the xrpc url - query := url.Values{} - query.Set("repo", f.RepoDid) - query.Set("ref", ref) - query.Set("format", format) - query.Set("prefix", r.URL.Query().Get("prefix")) - xrpcURL := fmt.Sprintf( - "%s/xrpc/%s?%s", - rp.config.KnotMirror.Url, - tangled.GitTempGetArchiveNSID, - query.Encode(), - ) + xrpcURL := fmt.Sprintf("%s/xrpc/%s?%s", + rp.config.KnotMirror.Url, tangled.GitTempGetArchiveNSID, params.Query(f.RepoDid).Encode()) // make the get request - resp, err := http.Get(xrpcURL) + req, err := http.NewRequestWithContext(r.Context(), http.MethodGet, xrpcURL, nil) + if err != nil { + l.Error("failed to build XRPC repo.archive request", "err", err) + fail(http.StatusServiceUnavailable) + return + } + resp, err := rp.archiveClient.Do(req) if err != nil { l.Error("failed to call XRPC repo.archive", "err", err) - rp.pages.Error503(w) + fail(http.StatusServiceUnavailable) return } defer resp.Body.Close() - w.Header().Set("Content-Type", archiveContentType(format)) - - filename := "" - if cd := resp.Header.Get("Content-Disposition"); strings.HasPrefix(cd, "attachment;") { - filename = cd // knot has already set the attachment CD - } - if filename == "" { - filename = fmt.Sprintf("attachment; filename=\"%s-%s.%s\"", f.Name, ref, format) + if resp.StatusCode != http.StatusOK { + l.Error("XRPC repo.archive failed", "status", resp.StatusCode, "ref", params.Rev) + overloaded := resp.StatusCode >= http.StatusInternalServerError || resp.StatusCode == http.StatusTooManyRequests + fail(lo.Ternary(overloaded, http.StatusServiceUnavailable, http.StatusNotFound)) + return } - w.Header().Set("Content-Disposition", filename) - w.Header().Set("X-Content-Type-Options", "nosniff") - - if link := resp.Header.Get("Link"); link != "" { - if resolvedRef, err := extractImmutableLink(link); err == nil { - newLink := fmt.Sprintf("<%s/%s/archive/%s.%s>; rel=\"immutable\"", - rp.config.Core.BaseUrl(), f.RepoIdentifier(), resolvedRef, format) - w.Header().Set("Link", newLink) - } + + params.SetHeaders(w.Header(), name) + + if resolvedRev, err := gitutil.ParseImmutableLink(resp.Header.Get("Link")); err == nil { + w.Header().Set("Link", gitutil.ImmutableLink(rp.immutableArchiveURL(f, params.WithRev(resolvedRev)))) } // stream the archive data directly @@ -71,56 +85,49 @@ func (rp *Repo) DownloadArchive(w http.ResponseWriter, r *http.Request) { } } -func archiveRefAndFormat(ref string, requestedFormat string, userAgent string) (string, string) { - switch { - case strings.HasSuffix(ref, ".tar.gz"): - ref = strings.TrimSuffix(ref, ".tar.gz") - if requestedFormat == "" { - requestedFormat = "tar.gz" - } - case strings.HasSuffix(ref, ".zip"): - ref = strings.TrimSuffix(ref, ".zip") - if requestedFormat == "" { - requestedFormat = "zip" - } +func parseArchiveRequest(r *http.Request) (gitutil.ArchiveParams, error) { + ref := chi.URLParam(r, "*") + if unescaped, err := url.PathUnescape(ref); err == nil && r.URL.RawPath != "" { + ref = unescaped } - switch requestedFormat { - case "zip", "tar.gz": - return ref, requestedFormat - default: - if prefersZipArchive(userAgent) { - return ref, "zip" - } - return ref, "tar.gz" + suffix, found := lo.Find(gitutil.ArchiveFormats, func(f gitutil.ArchiveFormat) bool { + return strings.HasSuffix(ref, "."+f.String()) + }) + if found { + ref = strings.TrimSuffix(ref, "."+suffix.String()) } -} - -func prefersZipArchive(userAgent string) bool { - ua := strings.ToLower(userAgent) - return strings.Contains(ua, "windows") || strings.Contains(ua, "win64") || strings.Contains(ua, "win32") -} -func archiveContentType(format string) string { - if format == "zip" { - return "application/zip" + rev, err := gitutil.ParseRev(ref) + if err != nil { + return gitutil.ArchiveParams{}, err } - return "application/gzip" -} -func extractImmutableLink(linkHeader string) (string, error) { - trimmed := strings.TrimPrefix(linkHeader, "<") - trimmed = strings.TrimSuffix(trimmed, ">; rel=\"immutable\"") - - parsedLink, err := url.Parse(trimmed) + query := r.URL.Query() + query.Del("ref") + query.Set("format", archiveFormat(query.Get("format"), suffix, r.UserAgent()).String()) + params, err := gitutil.ParseArchiveParams(query) if err != nil { - return "", err + return gitutil.ArchiveParams{}, err } + return params.WithRev(rev), nil +} - resolvedRef := parsedLink.Query().Get("ref") - if resolvedRef == "" { - return "", fmt.Errorf("no ref found in link header") +func archiveFormat(requested string, suffix gitutil.ArchiveFormat, userAgent string) gitutil.ArchiveFormat { + if format, err := gitutil.ParseArchiveFormat(requested); err == nil { + return format } + if suffix != "" { + return suffix + } + ua := strings.ToLower(userAgent) + windows := lo.SomeBy([]string{"windows", "win64", "win32"}, func(s string) bool { return strings.Contains(ua, s) }) + return lo.Ternary(windows, gitutil.ArchiveZip, gitutil.ArchiveTarGz) +} - return resolvedRef, nil +func (rp *Repo) immutableArchiveURL(f *models.Repo, params gitutil.ArchiveParams) string { + return fmt.Sprintf("%s/%s/archive/%s.%s?%s", + rp.config.Core.BaseUrl(), f.RepoIdentifier(), + url.PathEscape(params.Rev.String()), params.Format, + url.Values{"prefix": {params.Prefix.String()}}.Encode()) } diff --git a/appview/repo/archive_test.go b/appview/repo/archive_test.go new file mode 100644 index 00000000..1f79a3c4 --- /dev/null +++ b/appview/repo/archive_test.go @@ -0,0 +1,94 @@ +package repo + +import ( + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + + "github.com/go-chi/chi/v5" + "tangled.org/core/appview/config" + "tangled.org/core/appview/models" + "tangled.org/core/gitutil" +) + +const ( + testRepoDid = "did:plc:limpet" + testRepoOwner = "did:plc:boltless" + testRepoRkey = "3kzabcdefghij" + testRepoPath = "/boltless.dev/squid" +) + +func TestParseArchiveRequest(t *testing.T) { + var ( + got gitutil.ArchiveParams + gotErr error + ) + router := chi.NewRouter() + router.Get("/{user}/{repo}"+archiveRoute, func(w http.ResponseWriter, r *http.Request) { + got, gotErr = parseArchiveRequest(r) + }) + + resolvedParams := gitutil.ArchiveParams{ + Rev: "6f1d3a2b4c5d6e7f8091a2b3c4d5e6f708192a3b", + Format: gitutil.ArchiveZip, + Prefix: gitutil.ArchivePrefix("").OrDefault("squid", "refs/heads/feat/uni"), + } + rp := &Repo{config: &config.Config{Core: config.CoreConfig{Dev: true, AppviewHost: "tangled.org"}}} + immutable, err := url.Parse(rp.immutableArchiveURL( + &models.Repo{Did: testRepoOwner, Rkey: testRepoRkey, Name: "squid", RepoDid: testRepoDid}, + resolvedParams, + )) + if err != nil { + t.Fatalf("the immutable URL must parse: %v", err) + } + + windows := "Mozilla/5.0 (Windows NT 10.0; Win64; x64)" + targz, zip := gitutil.ArchiveTarGz, gitutil.ArchiveZip + cases := []struct { + name string + path string + userAgent string + want gitutil.ArchiveParams + wantErr bool + }{ + {"short ref", testRepoPath + "/archive/v1.0.0?format=tar.gz", "", gitutil.ArchiveParams{Rev: "v1.0.0", Format: targz}, false}, + {"full ref unescaped", testRepoPath + "/archive/refs/tags/v1.0.0?prefix=did:plc:boltless", "", gitutil.ArchiveParams{Rev: "refs/tags/v1.0.0", Format: targz, Prefix: "did:plc:boltless"}, false}, + {"full ref escaped", testRepoPath + "/archive/refs%2Ftags%2Fv1.0.0?prefix=did:plc:boltless", "", gitutil.ArchiveParams{Rev: "refs/tags/v1.0.0", Format: targz, Prefix: "did:plc:boltless"}, false}, + {"format from suffix", testRepoPath + "/archive/refs/tags/v1.0.0.zip", "", gitutil.ArchiveParams{Rev: "refs/tags/v1.0.0", Format: zip}, false}, + {"unknown format query with a zip suffix", testRepoPath + "/archive/refs/tags/v1.0.0.zip?format=tar.xz", "", gitutil.ArchiveParams{Rev: "refs/tags/v1.0.0", Format: zip}, false}, + {"zip for a windows user agent", testRepoPath + "/archive/main?format=tar.xz", windows, gitutil.ArchiveParams{Rev: "main", Format: zip}, false}, + {"percent in the ref itself", testRepoPath + "/archive/refs/tags/a%252Fb", "", gitutil.ArchiveParams{Rev: "refs/tags/a%2Fb", Format: targz}, false}, + {"prefix wrapped in slashes", testRepoPath + "/archive/main?prefix=/kelp/", "", gitutil.ArchiveParams{Rev: "main", Format: targz, Prefix: "kelp"}, false}, + {"traversal escaped", testRepoPath + "/archive/..%2F..%2Fetc", "", gitutil.ArchiveParams{Rev: "../../etc", Format: targz}, false}, + {"parse deletes a ref query", testRepoPath + "/archive/main?ref=other", "", gitutil.ArchiveParams{Rev: "main", Format: targz}, false}, + {"our own immutable URL", immutable.RequestURI(), "", resolvedParams, false}, + + {"empty ref", testRepoPath + "/archive/", "", gitutil.ArchiveParams{}, true}, + {"escaped space", testRepoPath + "/archive/refs/tags/a%20b", "", gitutil.ArchiveParams{}, true}, + {"ref that git would read as an option", testRepoPath + "/archive/--output=%2Ftmp%2Fevil", "", gitutil.ArchiveParams{}, true}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, gotErr = gitutil.ArchiveParams{}, nil + path := tc.path + if rest, isRepoDid := strings.CutPrefix(path, "/"+testRepoDid); isRepoDid { + path = "/" + testRepoOwner + "/" + testRepoRkey + rest + } + + req := httptest.NewRequest(http.MethodGet, path, nil) + req.Header.Set("User-Agent", tc.userAgent) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("%s: status = %d, want the archive route to match", path, rec.Code) + } + if got != tc.want || (gotErr != nil) != tc.wantErr { + t.Errorf("params = %+v with err %v, want %+v and rejected = %v", got, gotErr, tc.want, tc.wantErr) + } + }) + } +} diff --git a/appview/repo/repo.go b/appview/repo/repo.go index 414835ec..f18037ab 100644 --- a/appview/repo/repo.go +++ b/appview/repo/repo.go @@ -61,6 +61,7 @@ type Repo struct { codesearch *codesearch.CodeSearch knotMirrorXRPC *indigoxrpc.Client + archiveClient *http.Client } func New( @@ -93,6 +94,7 @@ func New( codesearch: codesearch, knotMirrorXRPC: newKnotMirrorXRPCClient(config.KnotMirror.Url), + archiveClient: newArchiveClient(config.KnotMirror.ArchiveHeaderTimeout), } } diff --git a/appview/repo/router.go b/appview/repo/router.go index ea17ad28..c3b053ba 100644 --- a/appview/repo/router.go +++ b/appview/repo/router.go @@ -45,9 +45,7 @@ func (rp *Repo) Router(mw *middleware.Middleware) http.Handler { r.Get("/blob/{ref}/*", rp.Blob) r.Get("/raw/{ref}/*", rp.RepoBlobRaw) - // intentionally doesn't use /* as this isn't - // a file path - r.Get("/archive/{ref}", rp.DownloadArchive) + r.Get(archiveRoute, rp.DownloadArchive) r.With(middleware.Paginate).Get("/stars", rp.Stars) r.With(middleware.Paginate).Get("/forks", rp.Forks) diff --git a/appview/reporesolver/resolver.go b/appview/reporesolver/resolver.go index 1d0e91ee..5f5ef240 100644 --- a/appview/reporesolver/resolver.go +++ b/appview/reporesolver/resolver.go @@ -40,7 +40,7 @@ func CanonicalRepoPath(handle string, repo *models.Repo) string { } func CanonicalRedirectTarget(req *http.Request, canonical string) string { - parts := strings.SplitN(strings.TrimPrefix(req.URL.Path, "/"), "/", 3) + parts := strings.SplitN(strings.TrimPrefix(req.URL.EscapedPath(), "/"), "/", 3) target := "/" + canonical if len(parts) == 3 { target += "/" + parts[2] diff --git a/appview/reporesolver/resolver_test.go b/appview/reporesolver/resolver_test.go index 42079d70..31a8e9a2 100644 --- a/appview/reporesolver/resolver_test.go +++ b/appview/reporesolver/resolver_test.go @@ -51,6 +51,29 @@ func TestCanonicalRepoPath(t *testing.T) { } } +func TestCanonicalRedirectTargetKeepsTheTailEscaped(t *testing.T) { + cases := []struct { + name string + path string + want string + }{ + {"plain tail", "/boltless.dev/limpet/tree/main", "/akshay.dev/anemone/tree/main"}, + {"space in a blob path", "/boltless.dev/limpet/blob/main/a%20b.txt", "/akshay.dev/anemone/blob/main/a%20b.txt"}, + {"escaped slash in a ref", "/boltless.dev/limpet/archive/refs%2Fheads%2Fmain", "/akshay.dev/anemone/archive/refs%2Fheads%2Fmain"}, + {"hash in a filename", "/boltless.dev/limpet/raw/main/c%23.cs", "/akshay.dev/anemone/raw/main/c%23.cs"}, + {"repo root", "/boltless.dev/limpet", "/akshay.dev/anemone"}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + req := httptest.NewRequest("GET", c.path, nil) + if got := CanonicalRedirectTarget(req, "akshay.dev/anemone"); got != c.want { + t.Errorf("CanonicalRedirectTarget = %q, want %q", got, c.want) + } + }) + } +} + func reqWithChiParams(user, repo string) *http.Request { r := httptest.NewRequest("GET", "/", nil) rctx := chi.NewRouteContext() diff --git a/gitutil/archive.go b/gitutil/archive.go new file mode 100644 index 00000000..dd30ab32 --- /dev/null +++ b/gitutil/archive.go @@ -0,0 +1,198 @@ +package gitutil + +import ( + "bytes" + "context" + "fmt" + "io" + "mime" + "net/http" + "net/url" + "os/exec" + "path" + "slices" + "strings" + "syscall" + "time" + "unicode" + + "github.com/go-git/go-git/v5/plumbing" + "github.com/samber/lo" +) + +type ArchiveFormat string + +const ( + ArchiveTarGz ArchiveFormat = "tar.gz" + ArchiveZip ArchiveFormat = "zip" +) + +var ArchiveFormats = []ArchiveFormat{ArchiveTarGz, ArchiveZip} + +func ParseArchiveFormat(raw string) (ArchiveFormat, error) { + if format := ArchiveFormat(raw); slices.Contains(ArchiveFormats, format) { + return format, nil + } + return "", fmt.Errorf("only tar.gz and zip formats are supported, got %q", raw) +} + +func (f ArchiveFormat) String() string { return string(f) } + +func (f ArchiveFormat) contentType() string { + return lo.Ternary(f == ArchiveZip, "application/zip", "application/gzip") +} + +var pathSeparators = strings.NewReplacer("/", "-", `\`, "-") + +type RepoName string + +type Rev string + +const RevHead Rev = "HEAD" + +func ParseRev(raw string) (Rev, error) { + switch { + case raw == "": + return "", fmt.Errorf("ref is empty") + case strings.ContainsFunc(raw, func(c rune) bool { return unicode.IsSpace(c) || unicode.IsControl(c) }): + return "", fmt.Errorf("ref contains whitespace or a control character: %q", raw) + case strings.HasPrefix(raw, "-"): + return "", fmt.Errorf("ref starts with a dash: %q", raw) + } + return Rev(raw), nil +} + +func RevFromHash(h plumbing.Hash) Rev { return Rev(h.String()) } + +func (r Rev) String() string { return string(r) } + +func (r Rev) Slug() string { return pathSeparators.Replace(plumbing.ReferenceName(r).Short()) } + +func (r Rev) Or(fallback Rev) Rev { return lo.Ternary(r == "", fallback, r) } + +func (r Rev) OrHash(h plumbing.Hash) Rev { return r.Or(RevFromHash(h)) } + +type ArchivePrefix string + +const MaxArchivePrefixLen = 255 + +func ParseArchivePrefix(raw string) (ArchivePrefix, error) { + switch { + case len(raw) > MaxArchivePrefixLen: + return "", fmt.Errorf("prefix is %d bytes, over the %d byte limit", len(raw), MaxArchivePrefixLen) + case strings.ContainsFunc(raw, unicode.IsControl): + return "", fmt.Errorf("prefix contains a control character: %q", raw) + case strings.Contains(raw, `\`): + return "", fmt.Errorf("prefix contains a backslash: %q", raw) + } + trimmed := strings.Trim(raw, "/") + if trimmed == "" { + return "", nil + } + cleaned := path.Clean(trimmed) + if cleaned == "." || cleaned == ".." || strings.HasPrefix(cleaned, "../") { + return "", fmt.Errorf("prefix escapes the archive root: %q", raw) + } + return ArchivePrefix(cleaned), nil +} + +func (p ArchivePrefix) String() string { return string(p) } + +func (p ArchivePrefix) OrDefault(repo RepoName, rev Rev) ArchivePrefix { + if p == "" { + return ArchivePrefix(archiveStem(repo, rev)) + } + return p +} + +func archiveStem(repo RepoName, rev Rev) string { + stem := pathSeparators.Replace(string(repo)) + "-" + rev.Slug() + if len(stem) <= MaxArchivePrefixLen { + return stem + } + return strings.ToValidUTF8(stem[:MaxArchivePrefixLen], "") +} + +type ArchiveParams struct { + Rev Rev + Format ArchiveFormat + Prefix ArchivePrefix +} + +func ParseArchiveParams(q url.Values) (ArchiveParams, error) { + p := ArchiveParams{Format: ArchiveTarGz} + var err error + if raw := q.Get("ref"); raw != "" { + if p.Rev, err = ParseRev(raw); err != nil { + return ArchiveParams{}, err + } + } + if raw := q.Get("format"); raw != "" { + if p.Format, err = ParseArchiveFormat(raw); err != nil { + return ArchiveParams{}, err + } + } + if p.Prefix, err = ParseArchivePrefix(q.Get("prefix")); err != nil { + return ArchiveParams{}, err + } + return p, nil +} + +func (p ArchiveParams) WithRev(rev Rev) ArchiveParams { + p.Rev = rev + return p +} + +func (p ArchiveParams) Query(repo string) url.Values { + return url.Values{ + "repo": {repo}, + "ref": {p.Rev.String()}, + "format": {p.Format.String()}, + "prefix": {p.Prefix.String()}, + } +} + +func (p ArchiveParams) SetHeaders(h http.Header, repo RepoName) { + h.Set("Content-Type", p.Format.contentType()) + h.Set("Content-Disposition", mime.FormatMediaType("attachment", map[string]string{ + "filename": archiveStem(repo, p.Rev) + "." + p.Format.String(), + })) + h.Set("X-Content-Type-Options", "nosniff") +} + +func ImmutableLink(target string) string { return `<` + target + `>; rel="immutable"` } + +func ParseImmutableLink(header string) (Rev, error) { + target := strings.TrimSuffix(strings.TrimPrefix(header, "<"), `>; rel="immutable"`) + parsed, err := url.Parse(target) + if err != nil { + return "", err + } + return ParseRev(parsed.Query().Get("ref")) +} + +const archiveWaitDelay = 10 * time.Second + +func WriteArchive(ctx context.Context, w io.Writer, repoPath string, rev Rev, format ArchiveFormat, prefix ArchivePrefix) error { + args := []string{"archive", "--format=" + format.String()} + if prefix != "" { + args = append(args, "--prefix="+prefix.String()+"/") + } + + cmd := exec.CommandContext(ctx, "git", append(args, "--", rev.String())...) + cmd.Dir = repoPath + cmd.Stdout = w + stderr := new(bytes.Buffer) + cmd.Stderr = stderr + cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} + cmd.WaitDelay = archiveWaitDelay + cmd.Cancel = func() error { + err := syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL) + return lo.Ternary(err == syscall.ESRCH, nil, err) + } + + if err := cmd.Run(); err != nil { + return fmt.Errorf("%w, stderr: %s", err, stderr.String()) + } + return nil +} diff --git a/gitutil/archive_test.go b/gitutil/archive_test.go new file mode 100644 index 00000000..c5cda8e3 --- /dev/null +++ b/gitutil/archive_test.go @@ -0,0 +1,187 @@ +package gitutil + +import ( + "archive/zip" + "bytes" + "context" + "mime" + "net/http" + "net/url" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "unicode/utf8" + + "github.com/go-git/go-git/v5/plumbing" + "github.com/samber/lo" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const testArchiveEndpoint = "https://knot.nel.pet/xrpc/sh.tangled.repo.archive" + +func TestParseArchiveParams(t *testing.T) { + cases := []struct { + name string + query url.Values + repo RepoName + want ArchiveParams + wantStem ArchivePrefix + wantFilename string + wantType string + wantErr string + }{ + {"all empty", url.Values{}, "", ArchiveParams{Format: ArchiveTarGz}, "", "", "", ""}, + {"tar.gz", url.Values{"format": {"tar.gz"}}, "", ArchiveParams{Format: ArchiveTarGz}, "", "", "", ""}, + {"zip", url.Values{"format": {"zip"}}, "", ArchiveParams{Format: ArchiveZip}, "", "", "", ""}, + { + "every param", + url.Values{"ref": {"refs/tags/v1.0.0"}, "format": {"zip"}, "prefix": {"/kelp/"}}, + "", ArchiveParams{Rev: "refs/tags/v1.0.0", Format: ArchiveZip, Prefix: "kelp"}, + "kelp", "squid-v1.0.0.zip", "application/zip", "", + }, + + {"branch", url.Values{"ref": {"main"}}, "", ArchiveParams{Rev: "main", Format: ArchiveTarGz}, "squid-main", "squid-main.tar.gz", "application/gzip", ""}, + {"full ref", url.Values{"ref": {"refs/heads/feat/uni"}}, "", ArchiveParams{Rev: "refs/heads/feat/uni", Format: ArchiveTarGz}, "squid-feat-uni", "squid-feat-uni.tar.gz", "application/gzip", ""}, + {"head", url.Values{"ref": {"HEAD"}}, "", ArchiveParams{Rev: "HEAD", Format: ArchiveTarGz}, "squid-HEAD", "", "", ""}, + {"trailing slash", url.Values{"ref": {"refs/heads/main/"}}, "", ArchiveParams{Rev: "refs/heads/main/", Format: ArchiveTarGz}, "squid-main-", "", "", ""}, + {"traversal in a ref", url.Values{"ref": {"../../etc"}}, "", ArchiveParams{Rev: "../../etc", Format: ArchiveTarGz}, "squid-..-..-etc", "", "", ""}, + {"windows separator in a ref", url.Values{"ref": {`feat\uni`}}, "", ArchiveParams{Rev: `feat\uni`, Format: ArchiveTarGz}, "squid-feat-uni", "", "", ""}, + {"slash in a repo name", url.Values{"ref": {"main"}}, "kelp/limpet", ArchiveParams{Rev: "main", Format: ArchiveTarGz}, "kelp-limpet-main", "kelp-limpet-main.tar.gz", "application/gzip", ""}, + {"quote in a repo name", url.Values{"ref": {"main"}}, `squid-a"b`, ArchiveParams{Rev: "main", Format: ArchiveTarGz}, `squid-a"b-main`, `squid-a"b-main.tar.gz`, "application/gzip", ""}, + {"non-ascii repo name", url.Values{"ref": {"main"}, "format": {"zip"}}, "squid-über", ArchiveParams{Rev: "main", Format: ArchiveZip}, "squid-über-main", "squid-über-main.zip", "application/zip", ""}, + + {"bare slash prefix", url.Values{"prefix": {"/"}}, "", ArchiveParams{Format: ArchiveTarGz}, "", "", "", ""}, + {"did prefix", url.Values{"prefix": {"did:plc:boltless"}}, "", ArchiveParams{Format: ArchiveTarGz, Prefix: "did:plc:boltless"}, "", "", "", ""}, + {"nested prefix", url.Values{"prefix": {"squid/main"}}, "", ArchiveParams{Format: ArchiveTarGz, Prefix: "squid/main"}, "", "", "", ""}, + {"prefix wrapped in slashes", url.Values{"prefix": {"/squid/main/"}}, "", ArchiveParams{Format: ArchiveTarGz, Prefix: "squid/main"}, "", "", "", ""}, + {"redundant prefix segments", url.Values{"prefix": {"squid/../limpet"}}, "", ArchiveParams{Format: ArchiveTarGz, Prefix: "limpet"}, "", "", "", ""}, + {"space in a prefix", url.Values{"prefix": {"squid main"}}, "", ArchiveParams{Format: ArchiveTarGz, Prefix: "squid main"}, "", "", "", ""}, + + {"unsupported format", url.Values{"format": {"tar"}}, "", ArchiveParams{}, "", "", "", "only tar.gz and zip formats are supported"}, + {"space in a ref", url.Values{"ref": {"refs/tags/a b"}}, "", ArchiveParams{}, "", "", "", "ref contains whitespace"}, + {"control character in a ref", url.Values{"ref": {"refs/tags/a\nb"}}, "", ArchiveParams{}, "", "", "", "ref contains whitespace"}, + {"ref that git would read as an option", url.Values{"ref": {"--output=/tmp/evil"}}, "", ArchiveParams{}, "", "", "", "ref starts with a dash"}, + {"prefix escaping the root", url.Values{"prefix": {"../../evil"}}, "", ArchiveParams{}, "", "", "", "prefix escapes the archive root"}, + {"prefix escaping after cleaning", url.Values{"prefix": {"squid/../../evil"}}, "", ArchiveParams{}, "", "", "", "prefix escapes the archive root"}, + {"bare dot prefix", url.Values{"prefix": {"."}}, "", ArchiveParams{}, "", "", "", "prefix escapes the archive root"}, + {"control character in a prefix", url.Values{"prefix": {"squid\nmain"}}, "", ArchiveParams{}, "", "", "", "prefix contains a control character"}, + {"windows separator in a prefix", url.Values{"prefix": {`..\..\evil`}}, "", ArchiveParams{}, "", "", "", "prefix contains a backslash"}, + {"prefix over the length limit", url.Values{"prefix": {strings.Repeat("a", MaxArchivePrefixLen+1)}}, "", ArchiveParams{}, "", "", "", "over the 255 byte limit"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ParseArchiveParams(tc.query) + rejected := err != nil + if got != tc.want || rejected != (tc.wantErr != "") || (rejected && !strings.Contains(err.Error(), tc.wantErr)) { + t.Fatalf("params = %+v with err %v, want %+v and an error mentioning %q", got, err, tc.want, tc.wantErr) + } + if rejected { + return + } + + repo := RepoName("squid") + if tc.repo != "" { + repo = tc.repo + } + if stem := got.Prefix.OrDefault(repo, got.Rev); tc.wantStem != "" && stem != tc.wantStem { + t.Errorf("default prefix = %q, want %q", stem, tc.wantStem) + } + if tc.wantFilename != "" { + header := http.Header{} + got.SetHeaders(header, repo) + mediatype, fields, err := mime.ParseMediaType(header.Get("Content-Disposition")) + if err != nil || mediatype != "attachment" || fields["filename"] != tc.wantFilename { + t.Errorf("Content-Disposition = %q (err %v), want an attachment with filename %q", header.Get("Content-Disposition"), err, tc.wantFilename) + } + if ct, sniff := header.Get("Content-Type"), header.Get("X-Content-Type-Options"); ct != tc.wantType || sniff != "nosniff" { + t.Errorf("Content-Type = %q with X-Content-Type-Options %q, want %q and nosniff", ct, sniff, tc.wantType) + } + } + + query := got.Query("did:plc:limpet") + if back, err := ParseArchiveParams(query); err != nil || back != got { + t.Errorf("query round trip = %+v (err %v), want %+v", back, err, got) + } + back, err := ParseImmutableLink(ImmutableLink(testArchiveEndpoint + "?" + query.Encode())) + if got.Rev != "" && (err != nil || back != got.Rev) { + t.Errorf("Link round trip = %q (err %v), want %q", back, err, got.Rev) + } + }) + } +} + +func TestArchiveFallbacks(t *testing.T) { + hash := plumbing.NewHash("6f1d3a2b4c5d6e7f8091a2b3c4d5e6f708192a3b") + if kept, filled := Rev("refs/heads/main").OrHash(hash), Rev("").OrHash(hash); kept != "refs/heads/main" || filled != RevFromHash(hash) { + t.Errorf("OrHash kept %q and filled %q, want refs/heads/main and the hash %q", kept, filled, hash) + } + if got := Rev("").Or(RevHead); got != RevHead { + t.Errorf("empty rev = %q, want HEAD", got) + } + if _, err := ParseImmutableLink(""); err == nil { + t.Error("ParseImmutableLink must reject an empty header") + } + + params := ArchiveParams{Rev: "main", Format: ArchiveZip, Prefix: "kelp"} + if got := params.WithRev("6f1d3a2"); got != (ArchiveParams{Rev: "6f1d3a2", Format: ArchiveZip, Prefix: "kelp"}) || params.Rev != "main" { + t.Errorf("WithRev = %+v leaving the receiver at %q, want only the rev replaced", got, params.Rev) + } + + long := ArchivePrefix("").OrDefault("squid", Rev("refs/heads/"+strings.Repeat("ü", 400))) + if _, err := ParseArchivePrefix(long.String()); len(long) > MaxArchivePrefixLen || !utf8.ValidString(long.String()) || err != nil { + t.Errorf("default prefix is %d bytes %q (err %v), want at most %d bytes ending on a rune boundary", len(long), long, err, MaxArchivePrefixLen) + } +} + +func TestWriteArchive(t *testing.T) { + repoPath := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(repoPath, "README.md"), []byte("# squid\n"), 0644)) + for _, args := range [][]string{ + {"init", "-q", "-b", "main"}, + {"add", "README.md"}, + {"-c", "user.name=nel", "-c", "user.email=nel@nel.pet", "commit", "-qm", "Initial commit"}, + } { + cmd := exec.Command("git", args...) + cmd.Dir = repoPath + require.NoError(t, cmd.Run(), "git %v", args) + } + + canceled, cancel := context.WithCancel(context.Background()) + cancel() + + cases := []struct { + name string + ctx context.Context + prefix ArchivePrefix + want []string + }{ + { + "prefix on every entry", + context.Background(), + ArchivePrefix("").OrDefault("squid", "refs/heads/feat/uni"), + []string{"squid-feat-uni/", "squid-feat-uni/README.md"}, + }, + {"empty prefix", context.Background(), "", []string{"README.md"}}, + {"canceled context", canceled, "squid-main", nil}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + var body bytes.Buffer + err := WriteArchive(tc.ctx, &body, repoPath, RevHead, ArchiveZip, tc.prefix) + if tc.want == nil { + assert.Error(t, err) + return + } + require.NoError(t, err) + + entries, err := zip.NewReader(bytes.NewReader(body.Bytes()), int64(body.Len())) + require.NoError(t, err) + assert.Equal(t, tc.want, lo.Map(entries.File, func(f *zip.File, _ int) string { return f.Name })) + }) + } +} diff --git a/knotmirror/xrpc/git_get_archive.go b/knotmirror/xrpc/git_get_archive.go index 62f789e8..198c82a9 100644 --- a/knotmirror/xrpc/git_get_archive.go +++ b/knotmirror/xrpc/git_get_archive.go @@ -1,145 +1,79 @@ package xrpc import ( - "bytes" - "context" "fmt" - "io" "net/http" - "net/url" - "os/exec" - "strings" "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/syntax" - "github.com/go-git/go-git/v5/plumbing" "tangled.org/core/api/tangled" + "tangled.org/core/gitutil" "tangled.org/core/knotmirror/db" "tangled.org/core/knotmirror/xrpc/gitea" ) func (x *Xrpc) GetArchive(w http.ResponseWriter, r *http.Request) { - var ( - repoQuery = r.URL.Query().Get("repo") - ref = r.URL.Query().Get("ref") - format = r.URL.Query().Get("format") - prefix = r.URL.Query().Get("prefix") - ) + invalid := func(err error) { + writeJson(w, http.StatusBadRequest, atclient.ErrorBody{Name: "InvalidRequest", Message: err.Error()}) + } + repoQuery := r.URL.Query().Get("repo") repo, err := syntax.ParseDID(repoQuery) if err != nil { - writeJson(w, http.StatusBadRequest, atclient.ErrorBody{Name: "BadRequest", Message: fmt.Sprintf("repo parameter invalid: %s", repoQuery)}) + invalid(fmt.Errorf("repo parameter invalid: %s", repoQuery)) return } - if format == "" { - format = "tar.gz" - } - if format != "tar.gz" && format != "zip" { - writeJson(w, http.StatusBadRequest, atclient.ErrorBody{Name: "BadRequest", Message: "only tar.gz and zip formats are supported"}) + params, err := gitutil.ParseArchiveParams(r.URL.Query()) + if err != nil { + invalid(err) return } - l := x.logger.With("repo", repo, "ref", ref, "format", format, "prefix", prefix) + l := x.logger.With("repo", repo, "ref", params.Rev, "format", params.Format, "prefix", params.Prefix) l.Debug("request") ctx := r.Context() - - repoPath, err := x.makeRepoPath(ctx, repo) - if err != nil { + proxy := func(err error, message string) { l.Warn("local mirror failed, trying proxy", "err", err) - if x.proxyToKnot(w, r, repo) { - return + if !x.proxyToKnot(w, r, repo) { + writeJson(w, http.StatusInternalServerError, atclient.ErrorBody{Name: "InternalServerError", Message: message}) } - writeJson(w, http.StatusInternalServerError, atclient.ErrorBody{Name: "InternalServerError", Message: "failed to resolve repo"}) - return } - rev := ref - if rev == "" { - rev = "HEAD" - } - commit, err := gitea.GetCommit(ctx, repoPath, rev) + repoPath, err := x.makeRepoPath(ctx, repo) if err != nil { - l.Warn("local mirror failed, trying proxy", "err", err) - if x.proxyToKnot(w, r, repo) { - return - } - writeJson(w, http.StatusInternalServerError, atclient.ErrorBody{Name: "InternalServerError", Message: "failed to resolve ref"}) + proxy(err, "failed to resolve repo") return } - repoName, err := func() (string, error) { - r, err := db.GetRepoByRepoDid(ctx, x.db, repo) - if err != nil { - return "", err - } - if r == nil { - return "", fmt.Errorf("repo not found: %s", repo) - } - return r.Name, nil - }() + commit, err := gitea.GetCommit(ctx, repoPath, params.Rev.Or(gitutil.RevHead).String()) if err != nil { - l.Warn("local mirror failed, trying proxy", "err", err) - if x.proxyToKnot(w, r, repo) { - return - } - writeJson(w, http.StatusInternalServerError, atclient.ErrorBody{Name: "InternalServerError", Message: "failed to retrieve repo name"}) + proxy(err, "failed to resolve ref") return } - safeRefFilename := strings.ReplaceAll(plumbing.ReferenceName(ref).Short(), "/", "-") - if safeRefFilename == "" { - safeRefFilename = commit.Hash.String() - } - immutableLink := func() string { - params := url.Values{} - params.Set("repo", repo.String()) - params.Set("ref", commit.Hash.String()) - params.Set("format", format) - params.Set("prefix", prefix) - return fmt.Sprintf("%s/xrpc/%s?%s", x.cfg.BaseUrl(), tangled.GitTempGetArchiveNSID, params.Encode()) - }() - - var archivePrefix string - if prefix != "" { - archivePrefix = prefix - } else { - archivePrefix = fmt.Sprintf("%s-%s", repoName, safeRefFilename) - } - - filename := fmt.Sprintf("%s-%s.%s", repoName, safeRefFilename, format) - w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=\"%s\"", filename)) - w.Header().Set("Content-Type", archiveContentType(format)) - w.Header().Set("Link", fmt.Sprintf("<%s>; rel=\"immutable\"", immutableLink)) - - if err := writeLocalArchive(ctx, w, repoPath, commit.Hash.String(), format, archivePrefix); err != nil { - l.Error("writing archive", "err", err.Error(), "format", format) - w.WriteHeader(http.StatusInternalServerError) + mirrored, err := db.GetRepoByRepoDid(ctx, x.db, repo) + if err == nil && mirrored == nil { + err = fmt.Errorf("repo not found: %s", repo) } -} - -func archiveContentType(format string) string { - if format == "zip" { - return "application/zip" + if err != nil { + proxy(err, "failed to retrieve repo name") + return } - return "application/gzip" -} -func writeLocalArchive(ctx context.Context, w io.Writer, repoPath, rev, format, prefix string) error { - args := []string{"-C", repoPath, "archive", "--format=" + format} - if prefix != "" { - args = append(args, "--prefix="+strings.TrimRight(prefix, "/")+"/") - } - args = append(args, rev) + name := gitutil.RepoName(mirrored.Name) + resolvedRev := gitutil.RevFromHash(commit.Hash) + params.Rev = params.Rev.OrHash(commit.Hash) + params.Prefix = params.Prefix.OrDefault(name, params.Rev) - cmd := exec.CommandContext(ctx, "git", args...) - cmd.Stdout = w - stderr := new(bytes.Buffer) - cmd.Stderr = stderr + params.SetHeaders(w.Header(), name) + w.Header().Set("Link", gitutil.ImmutableLink(fmt.Sprintf("%s/xrpc/%s?%s", + x.cfg.BaseUrl(), tangled.GitTempGetArchiveNSID, params.WithRev(resolvedRev).Query(repo.String()).Encode(), + ))) - if err := cmd.Run(); err != nil { - return fmt.Errorf("%w, stderr: %s", err, stderr.String()) + if err := gitutil.WriteArchive(ctx, w, repoPath, resolvedRev, params.Format, params.Prefix); err != nil { + l.Error("writing archive", "err", err.Error(), "format", params.Format) + w.WriteHeader(http.StatusInternalServerError) } - return nil } diff --git a/knotmirror/xrpc/git_get_archive_test.go b/knotmirror/xrpc/git_get_archive_test.go new file mode 100644 index 00000000..7e76d795 --- /dev/null +++ b/knotmirror/xrpc/git_get_archive_test.go @@ -0,0 +1,26 @@ +package xrpc + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestGetArchiveRejectsBadParams(t *testing.T) { + for query, wantMessage := range map[string]string{ + "": "repo parameter invalid", + "repo=oyster.cafe%2Fsquid": "repo parameter invalid", + "repo=did:plc:boltless&format=tar.xz": "only tar.gz and zip formats are supported", + "repo=did:plc:boltless&ref=--output=/x": "ref starts with a dash", + "repo=did:plc:boltless&prefix=../../evil": "prefix escapes the archive root", + } { + rec := httptest.NewRecorder() + (&Xrpc{}).GetArchive(rec, httptest.NewRequest(http.MethodGet, "/xrpc/sh.tangled.git.temp.getArchive?"+query, nil)) + + body := rec.Body.String() + if rec.Code != http.StatusBadRequest || !strings.Contains(body, "InvalidRequest") || !strings.Contains(body, wantMessage) { + t.Errorf("%s: status %d with body %s, want 400 InvalidRequest mentioning %q", query, rec.Code, body, wantMessage) + } + } +} diff --git a/knotserver/git/cmd.go b/knotserver/git/cmd.go index e58d951a..21eabe1c 100644 --- a/knotserver/git/cmd.go +++ b/knotserver/git/cmd.go @@ -1,11 +1,8 @@ package git import ( - "bytes" "fmt" - "io" "os/exec" - "strings" "syscall" ) @@ -61,23 +58,3 @@ func (g *GitRepo) revParse(extraArgs ...string) ([]byte, error) { func (g *GitRepo) mergeBase(extraArgs ...string) ([]byte, error) { return g.runGitCmd("merge-base", extraArgs...) } - -func (g *GitRepo) WriteArchive(w io.Writer, format string, prefix string) error { - args := []string{"archive", "--format=" + format} - if prefix != "" { - args = append(args, "--prefix="+strings.TrimRight(prefix, "/")+"/") - } - args = append(args, g.h.String()) - - cmd := exec.Command("git", args...) - cmd.Dir = g.path - cmd.Stdout = w - stderr := new(bytes.Buffer) - cmd.Stderr = stderr - - if err := cmd.Run(); err != nil { - return fmt.Errorf("%w, stderr: %s", err, stderr.String()) - } - - return nil -} diff --git a/knotserver/xrpc/repo_archive.go b/knotserver/xrpc/repo_archive.go index fd22a2fd..a1225f09 100644 --- a/knotserver/xrpc/repo_archive.go +++ b/knotserver/xrpc/repo_archive.go @@ -3,109 +3,58 @@ package xrpc import ( "fmt" "net/http" - "net/url" - "strings" - - "github.com/go-git/go-git/v5/plumbing" "tangled.org/core/api/tangled" + "tangled.org/core/gitutil" "tangled.org/core/knotserver/git" xrpcerr "tangled.org/core/xrpc/errors" ) func (x *Xrpc) RepoArchive(w http.ResponseWriter, r *http.Request) { - repo := r.URL.Query().Get("repo") - repoPath, err := x.parseRepoParam(repo) + params, err := gitutil.ParseArchiveParams(r.URL.Query()) if err != nil { - writeError(w, err.(xrpcerr.XrpcError), http.StatusBadRequest) - return - } - - ref := r.URL.Query().Get("ref") - // ref can be empty (git.Open handles this) - - format := r.URL.Query().Get("format") - if format == "" { - format = "tar.gz" // default - } - - prefix := r.URL.Query().Get("prefix") - - if format != "tar.gz" && format != "zip" { writeError(w, xrpcerr.NewXrpcError( xrpcerr.WithTag("InvalidRequest"), - xrpcerr.WithMessage("only tar.gz and zip formats are supported"), + xrpcerr.WithMessage(err.Error()), ), http.StatusBadRequest) return } - gr, err := git.Open(repoPath, ref) + repo := r.URL.Query().Get("repo") + resolved, err := x.resolveRepo(repo) if err != nil { - writeError(w, xrpcerr.RefNotFoundError, http.StatusNotFound) + writeError(w, err.(xrpcerr.XrpcError), http.StatusBadRequest) return } - repoParts := strings.Split(repo, "/") - repoName := repoParts[len(repoParts)-1] - - immutableLink, err := x.buildImmutableLink(repo, format, gr.Hash().String(), prefix) + // ref can be empty (git.Open handles this) + gr, err := git.Open(resolved.path, params.Rev.String()) if err != nil { - x.Logger.Error( - "failed to build immutable link", - "err", err.Error(), - "repo", repo, - "format", format, - "ref", gr.Hash().String(), - "prefix", prefix, - ) + writeError(w, xrpcerr.RefNotFoundError, http.StatusNotFound) + return } - safeRefFilename := strings.ReplaceAll(plumbing.ReferenceName(ref).Short(), "/", "-") - - var archivePrefix string - if prefix != "" { - archivePrefix = prefix - } else { - archivePrefix = fmt.Sprintf("%s-%s", repoName, safeRefFilename) - } + hash := gr.Hash() + params.Rev = params.Rev.OrHash(hash) + params.Prefix = params.Prefix.OrDefault(resolved.name, params.Rev) - filename := fmt.Sprintf("%s-%s.%s", repoName, safeRefFilename, format) - w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=\"%s\"", filename)) - w.Header().Set("Content-Type", archiveContentType(format)) - w.Header().Set("Link", fmt.Sprintf("<%s>; rel=\"immutable\"", immutableLink)) + params.SetHeaders(w.Header(), resolved.name) + w.Header().Set("Link", gitutil.ImmutableLink( + x.archiveURL(repo, params.WithRev(gitutil.RevFromHash(hash))), + )) - err = gr.WriteArchive(w, format, archivePrefix) - if err != nil { + if err := gitutil.WriteArchive(r.Context(), w, resolved.path, gitutil.RevFromHash(hash), params.Format, params.Prefix); err != nil { // once we start writing to the body we can't report error anymore // so we are only left with logging the error - x.Logger.Error("writing archive", "error", err.Error(), "format", format) - return + x.Logger.Error("writing archive", "error", err.Error(), "format", params.Format) } } -func archiveContentType(format string) string { - if format == "zip" { - return "application/zip" - } - return "application/gzip" -} - -func (x *Xrpc) buildImmutableLink(repo string, format string, ref string, prefix string) (string, error) { +func (x *Xrpc) archiveURL(repo string, params gitutil.ArchiveParams) string { scheme := "https" if x.Config.Server.Dev { scheme = "http" } - - u, err := url.Parse(scheme + "://" + x.Config.Server.Hostname + "/xrpc/" + tangled.RepoArchiveNSID) - if err != nil { - return "", err - } - - params := url.Values{} - params.Set("repo", repo) - params.Set("format", format) - params.Set("ref", ref) - params.Set("prefix", prefix) - - return fmt.Sprintf("%s?%s", u.String(), params.Encode()), nil + return fmt.Sprintf("%s://%s/xrpc/%s?%s", + scheme, x.Config.Server.Hostname, tangled.RepoArchiveNSID, params.Query(repo).Encode()) } diff --git a/knotserver/xrpc/repo_archive_test.go b/knotserver/xrpc/repo_archive_test.go new file mode 100644 index 00000000..7a555af5 --- /dev/null +++ b/knotserver/xrpc/repo_archive_test.go @@ -0,0 +1,25 @@ +package xrpc + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestRepoArchiveChecksParamsBeforeResolvingTheRepo(t *testing.T) { + for query, wantMessage := range map[string]string{ + "format=tar.xz": "only tar.gz and zip formats are supported", + "ref=--output=/tmp/evil": "ref starts with a dash", + "prefix=../../evil": "prefix escapes the archive root", + "ref=refs/heads/main&format=zip": "repo parameter", + } { + rec := httptest.NewRecorder() + (&Xrpc{}).RepoArchive(rec, httptest.NewRequest(http.MethodGet, "/xrpc/sh.tangled.repo.archive?"+query, nil)) + + body := rec.Body.String() + if rec.Code != http.StatusBadRequest || !strings.Contains(body, "InvalidRequest") || !strings.Contains(body, wantMessage) { + t.Errorf("%s: status %d with body %s, want 400 InvalidRequest mentioning %q", query, rec.Code, body, wantMessage) + } + } +} diff --git a/knotserver/xrpc/repo_get_default_branch.go b/knotserver/xrpc/repo_get_default_branch.go index c16206d2..a53d15f3 100644 --- a/knotserver/xrpc/repo_get_default_branch.go +++ b/knotserver/xrpc/repo_get_default_branch.go @@ -18,6 +18,11 @@ func (x *Xrpc) RepoGetDefaultBranch(w http.ResponseWriter, r *http.Request) { } gr, err := git.PlainOpen(repoPath) + if err != nil { + x.Logger.Error("failed to open", "error", err.Error()) + writeError(w, xrpcerr.RepoNotFoundError, http.StatusNotFound) + return + } branch, err := gr.FindMainBranch() if err != nil { diff --git a/knotserver/xrpc/xrpc.go b/knotserver/xrpc/xrpc.go index ad9794d6..b87be0d4 100644 --- a/knotserver/xrpc/xrpc.go +++ b/knotserver/xrpc/xrpc.go @@ -13,6 +13,7 @@ import ( securejoin "github.com/cyphar/filepath-securejoin" "github.com/go-chi/chi/v5" "tangled.org/core/api/tangled" + "tangled.org/core/gitutil" "tangled.org/core/idresolver" "tangled.org/core/knotserver/config" "tangled.org/core/knotserver/db" @@ -97,20 +98,30 @@ func (x *Xrpc) Router() http.Handler { return r } +type resolvedRepo struct { + path string + name gitutil.RepoName +} + func (x *Xrpc) parseRepoParam(repo string) (string, error) { + resolved, err := x.resolveRepo(repo) + return resolved.path, err +} + +func (x *Xrpc) resolveRepo(repo string) (resolvedRepo, error) { if repo == "" || !strings.HasPrefix(repo, "did:") { - return "", xrpcerr.NewXrpcError( + return resolvedRepo{}, xrpcerr.NewXrpcError( xrpcerr.WithTag("InvalidRequest"), xrpcerr.WithMessage("missing or invalid repo parameter, expected a repo DID"), ) } if !strings.Contains(repo, "/") { - repoPath, _, _, err := x.Db.ResolveRepoDIDOnDisk(x.Config.Repo.ScanPath, repo) + repoPath, _, repoName, err := x.Db.ResolveRepoDIDOnDisk(x.Config.Repo.ScanPath, repo) if err != nil { - return "", xrpcerr.RepoNotFoundError + return resolvedRepo{}, xrpcerr.RepoNotFoundError } - return repoPath, nil + return resolvedRepo{path: repoPath, name: gitutil.RepoName(repoName)}, nil } parts := strings.SplitN(repo, "/", 2) @@ -120,18 +131,18 @@ func (x *Xrpc) parseRepoParam(repo string) (string, error) { if err == nil { repoPath, _, _, resolveErr := x.Db.ResolveRepoDIDOnDisk(x.Config.Repo.ScanPath, repoDid) if resolveErr == nil { - return repoPath, nil + return resolvedRepo{path: repoPath, name: gitutil.RepoName(repoName)}, nil } } repoPath, joinErr := securejoin.SecureJoin(x.Config.Repo.ScanPath, filepath.Join(ownerDid, repoName)) if joinErr != nil { - return "", xrpcerr.RepoNotFoundError + return resolvedRepo{}, xrpcerr.RepoNotFoundError } if _, statErr := os.Stat(repoPath); statErr != nil { - return "", xrpcerr.RepoNotFoundError + return resolvedRepo{}, xrpcerr.RepoNotFoundError } - return repoPath, nil + return resolvedRepo{path: repoPath, name: gitutil.RepoName(repoName)}, nil } func (x *Xrpc) resolveRepoDID(repo *string, ownerDid, name string) (repoident.RepoDid, string, error) {