diff --git a/internal/handler/composer.go b/internal/handler/composer.go index 2f475ac..3c10f6c 100644 --- a/internal/handler/composer.go +++ b/internal/handler/composer.go @@ -98,6 +98,12 @@ func (h *ComposerHandler) handlePackageMetadata(w http.ResponseWriter, r *http.R upstreamURL := fmt.Sprintf("%s/p2/%s/%s.json", h.repoURL, vendor, pkg) + if rewritten, ok := h.proxy.storedRewrite("composer", packageName, h.proxyURL, packageName); ok { + w.Header().Set(headerContentType, "application/json") + _, _ = w.Write(rewritten) + return + } + body, _, err := h.proxy.FetchOrCacheMetadata(r.Context(), "composer", packageName, upstreamURL) if err != nil { if errors.Is(err, ErrUpstreamNotFound) { diff --git a/internal/handler/handler.go b/internal/handler/handler.go index 20dce9a..449fcef 100644 --- a/internal/handler/handler.go +++ b/internal/handler/handler.go @@ -1383,7 +1383,7 @@ func (p *Proxy) cacheMetadataBlob(ctx context.Context, ecosystem, cacheKey, stor return } - size, _, err := p.Storage.Store(ctx, storagePath, bytes.NewReader(meta.body)) + size, hash, err := p.Storage.Store(ctx, storagePath, bytes.NewReader(meta.body)) if err != nil { p.Logger.Warn("failed to cache metadata", "ecosystem", ecosystem, "key", cacheKey, "error", err) return @@ -1396,9 +1396,12 @@ func (p *Proxy) cacheMetadataBlob(ctx context.Context, ecosystem, cacheKey, stor ETag: sql.NullString{String: meta.etag, Valid: meta.etag != ""}, ContentType: sql.NullString{String: meta.contentType, Valid: meta.contentType != ""}, ContentEncoding: sql.NullString{String: meta.contentEncoding, Valid: meta.contentEncoding != ""}, - Size: sql.NullInt64{Int64: size, Valid: true}, - LastModified: sql.NullTime{Time: meta.lastModified, Valid: !meta.lastModified.IsZero()}, - FetchedAt: sql.NullTime{Time: time.Now(), Valid: true}, + // The digest identifies the stored bytes, so a rewrite cached for them + // can be found without reading them back (see storedRewrite). + ContentDigest: sql.NullString{String: "sha256:" + hash, Valid: hash != ""}, + Size: sql.NullInt64{Int64: size, Valid: true}, + LastModified: sql.NullTime{Time: meta.lastModified, Valid: !meta.lastModified.IsZero()}, + FetchedAt: sql.NullTime{Time: time.Now(), Valid: true}, }) if err != nil { // The blob is written but the row describing it is not, so a later diff --git a/internal/handler/jsonscan.go b/internal/handler/jsonscan.go new file mode 100644 index 0000000..b1b1050 --- /dev/null +++ b/internal/handler/jsonscan.go @@ -0,0 +1,213 @@ +package handler + +import ( + "bytes" + "encoding/json" + "errors" +) + +// The helpers in this file walk a JSON document in place, reporting where +// object members sit as offsets into the original bytes. They let a handler +// read or replace one value in a large metadata document without decoding the +// rest of it into Go values, which for an npm packument costs many times the +// document's size. +// +// They check structure only as far as they need to find their way through +// the document. Callers that copy unvisited bytes into a response validate the +// whole document first with json.Valid. + +var errMalformedJSON = errors.New("malformed JSON") + +// errNotJSONObject is returned when a value expected to be an object is not. +var errNotJSONObject = errors.New("JSON value is not an object") + +// jsonMember is one member of a JSON object as offsets into the bytes that +// were scanned. The key span includes its quotes. +type jsonMember struct { + keyStart, keyEnd int + valStart, valEnd int +} + +func (m jsonMember) key(data []byte) []byte { return data[m.keyStart:m.keyEnd] } +func (m jsonMember) value(data []byte) []byte { return data[m.valStart:m.valEnd] } + +func skipJSONSpace(data []byte, i int) int { + for i < len(data) { + switch data[i] { + case ' ', '\t', '\n', '\r': + i++ + default: + return i + } + } + return i +} + +// skipJSONString returns the index just past the string starting at data[i]. +func skipJSONString(data []byte, i int) (int, error) { + for j := i + 1; j < len(data); j++ { + switch data[j] { + case '\\': + j++ + case '"': + return j + 1, nil + } + } + return 0, errMalformedJSON +} + +// skipJSONValue returns the index just past the value starting at data[i]. +func skipJSONValue(data []byte, i int) (int, error) { + if i >= len(data) { + return 0, errMalformedJSON + } + switch data[i] { + case '"': + return skipJSONString(data, i) + case '{', '[': + depth := 0 + for j := i; j < len(data); j++ { + switch data[j] { + case '"': + end, err := skipJSONString(data, j) + if err != nil { + return 0, err + } + j = end - 1 + case '{', '[': + depth++ + case '}', ']': + depth-- + if depth == 0 { + return j + 1, nil + } + } + } + return 0, errMalformedJSON + case '}', ']', ',', ':': + return 0, errMalformedJSON + default: + // A number, true, false or null runs to the next delimiter. + j := i + for j < len(data) { + switch data[j] { + case ',', '}', ']', ' ', '\t', '\n', '\r': + return j, nil + } + j++ + } + return j, nil + } +} + +// forEachJSONMember calls fn for each member of the object obj holds, in +// document order. It returns errNotJSONObject if obj is not an object, and +// stops early with fn's error if fn returns one. +func forEachJSONMember(obj []byte, fn func(jsonMember) error) error { + i := skipJSONSpace(obj, 0) + if i >= len(obj) || obj[i] != '{' { + return errNotJSONObject + } + i = skipJSONSpace(obj, i+1) + if i < len(obj) && obj[i] == '}' { + return nil + } + for { + if i >= len(obj) || obj[i] != '"' { + return errMalformedJSON + } + keyEnd, err := skipJSONString(obj, i) + if err != nil { + return err + } + colon := skipJSONSpace(obj, keyEnd) + if colon >= len(obj) || obj[colon] != ':' { + return errMalformedJSON + } + valStart := skipJSONSpace(obj, colon+1) + valEnd, err := skipJSONValue(obj, valStart) + if err != nil { + return err + } + if err := fn(jsonMember{keyStart: i, keyEnd: keyEnd, valStart: valStart, valEnd: valEnd}); err != nil { + return err + } + next := skipJSONSpace(obj, valEnd) + if next >= len(obj) { + return errMalformedJSON + } + switch obj[next] { + case ',': + i = skipJSONSpace(obj, next+1) + case '}': + return nil + default: + return errMalformedJSON + } + } +} + +// jsonKey decodes a quoted key as forEachJSONMember reports it. +func jsonKey(raw []byte) (string, error) { + if bytes.IndexByte(raw, '\\') < 0 { + return string(raw[1 : len(raw)-1]), nil + } + var key string + err := json.Unmarshal(raw, &key) + return key, err +} + +// jsonKeyIs reports whether a quoted key decodes to name. +func jsonKeyIs(raw []byte, name string) bool { + if bytes.IndexByte(raw, '\\') < 0 { + return string(raw[1:len(raw)-1]) == name + } + key, err := jsonKey(raw) + return err == nil && key == name +} + +// findJSONMember returns the member of obj named name. When the key repeats, +// the last one wins, as it does for JSON.parse and encoding/json. +func findJSONMember(obj []byte, name string) (jsonMember, bool, error) { + var found jsonMember + ok := false + err := forEachJSONMember(obj, func(m jsonMember) error { + if jsonKeyIs(m.key(obj), name) { + found, ok = m, true + } + return nil + }) + return found, ok, err +} + +// lookupJSON follows path through nested objects in doc and returns the +// value at its end. It returns nil and no error when a key along the path is +// missing or names something other than an object. +func lookupJSON(doc []byte, path ...string) ([]byte, error) { + value := doc + for _, name := range path { + m, ok, err := findJSONMember(value, name) + if errors.Is(err, errNotJSONObject) { + return nil, nil + } + if err != nil || !ok { + return nil, err + } + value = m.value(value) + } + return value, nil +} + +// lookupJSONString is lookupJSON for a string value. ok is false when the +// value is missing or is not a string. +func lookupJSONString(doc []byte, path ...string) (string, bool, error) { + raw, err := lookupJSON(doc, path...) + if err != nil || len(raw) == 0 || raw[0] != '"' { + return "", false, err + } + var s string + if err := json.Unmarshal(raw, &s); err != nil { + return "", false, err + } + return s, true, nil +} diff --git a/internal/handler/jsonscan_test.go b/internal/handler/jsonscan_test.go new file mode 100644 index 0000000..5003809 --- /dev/null +++ b/internal/handler/jsonscan_test.go @@ -0,0 +1,98 @@ +package handler + +import ( + "bytes" + "errors" + "testing" +) + +func TestLookupJSONString(t *testing.T) { + doc := []byte(`{ + "name": "demo", + "versions": { + "1.0.0": {"dist": {"tarball": "https://example.com/a.tgz", "shasum": "x"}}, + "2.0.0": {"readme": "a } tricky \" string ] with {brackets}", "dist": {"tarball": "https://example.com/b.tgz"}}, + "3.0.0": "not an object", + "4.0.0": {"dist": {"tarball": 42}} + }, + "time": {"1.0.0": "2020-01-01T00:00:00Z", "1.0.0": "2021-01-01T00:00:00Z"}, + "n": -1.5e3, "t": true, "z": null, "list": [1, {"a": []}, "]"] + }`) + + cases := []struct { + path []string + want string + wantOK bool + }{ + {[]string{"versions", "1.0.0", "dist", "tarball"}, "https://example.com/a.tgz", true}, + {[]string{"versions", "2.0.0", "dist", "tarball"}, "https://example.com/b.tgz", true}, + {[]string{"versions", "3.0.0", "dist", "tarball"}, "", false}, + {[]string{"versions", "4.0.0", "dist", "tarball"}, "", false}, + {[]string{"versions", "9.9.9", "dist", "tarball"}, "", false}, + {[]string{"time", "1.0.0"}, "2021-01-01T00:00:00Z", true}, // last duplicate wins + {[]string{"name"}, "demo", true}, + {[]string{"name", "x"}, "", false}, + } + for _, c := range cases { + got, ok, err := lookupJSONString(doc, c.path...) + if err != nil || ok != c.wantOK || got != c.want { + t.Errorf("lookup %v = %q, %v, %v; want %q, %v", c.path, got, ok, err, c.want, c.wantOK) + } + } +} + +func TestLookupJSONEscapedKey(t *testing.T) { + doc := []byte(`{"versions": {"1.0.0": {"dist": {"tarball": "u"}}}}`) + got, ok, err := lookupJSONString(doc, "versions", "1.0.0", "dist", "tarball") + if err != nil || !ok || got != "u" { + t.Errorf("lookup = %q, %v, %v", got, ok, err) + } +} + +func TestLookupJSONMalformed(t *testing.T) { + for _, doc := range []string{ + `{"versions": {"1.0.0": {"dist": `, + `{"versions" {}}`, + `{"versions": {"a": 1 "b": 2}}`, + `{"a": "unterminated}`, + `{"a": [1, 2}`, + } { + _, _, err := lookupJSONString([]byte(doc), "versions", "1.0.0", "dist", "tarball") + if !errors.Is(err, errMalformedJSON) { + t.Errorf("%s: err = %v, want errMalformedJSON", doc, err) + } + } +} + +func TestForEachJSONMemberOffsets(t *testing.T) { + doc := []byte(` { "a" : 1 , "b":{"c":[true,null]} ,"d":"x\"y" } `) + var got [][2]string + err := forEachJSONMember(doc, func(m jsonMember) error { + got = append(got, [2]string{string(m.key(doc)), string(m.value(doc))}) + return nil + }) + want := [][2]string{{`"a"`, `1`}, {`"b"`, `{"c":[true,null]}`}, {`"d"`, `"x\"y"`}} + if err != nil || len(got) != len(want) { + t.Fatalf("members = %v, %v", got, err) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("member %d = %v, want %v", i, got[i], want[i]) + } + } + + if err := forEachJSONMember([]byte(`[1]`), func(jsonMember) error { return nil }); !errors.Is(err, errNotJSONObject) { + t.Errorf("array: err = %v, want errNotJSONObject", err) + } + if err := forEachJSONMember([]byte(`{}`), func(jsonMember) error { t.Error("called for empty object"); return nil }); err != nil { + t.Errorf("empty object: err = %v", err) + } +} + +func TestWriteFilteredJSONObject(t *testing.T) { + var out bytes.Buffer + err := writeFilteredJSONObject(&out, []byte(`{"a": 1, "b": {"x": 2}, "c": 3}`), func(k string) bool { return k != "b" }) + if err != nil || out.String() != `{"a": 1,"c": 3}` { + t.Errorf("filtered = %s, %v", out.String(), err) + } +} diff --git a/internal/handler/npm.go b/internal/handler/npm.go index 98ae10a..31c5341 100644 --- a/internal/handler/npm.go +++ b/internal/handler/npm.go @@ -6,8 +6,10 @@ import ( "errors" "fmt" "io" + "maps" "net/http" "net/url" + "reflect" "sort" "strings" "sync" @@ -178,6 +180,13 @@ func (h *NPMHandler) handlePackageMetadata(w http.ResponseWriter, r *http.Reques accept = contentTypeJSON } + if rewritten, ok := h.proxy.storedRewrite("npm", packageName, h.proxyURL, packageName); ok { + w.Header().Set(headerContentType, contentTypeJSON) + w.WriteHeader(http.StatusOK) + _, _ = w.Write(rewritten) + return + } + body, _, err := h.proxy.FetchOrCacheMetadata(r.Context(), "npm", packageName, upstreamURL, accept) if err != nil { if errors.Is(err, ErrUpstreamNotFound) { @@ -215,26 +224,235 @@ func (h *NPMHandler) handlePackageMetadata(w http.ResponseWriter, r *http.Reques // rewriteMetadata rewrites tarball URLs in npm package metadata to point at this proxy. // If cooldown is enabled, versions published too recently are filtered out. +// +// The document is never decoded as a whole. A full packument can run to tens +// of megabytes, and decoding it into generic maps took many times that in +// memory for every rewrite. Instead the version names, time map and dist-tags +// are decoded to decide what to keep, and the response is assembled from +// slices of the original bytes: only each version's tarball URL, and the time +// map and dist-tags when filtering changed them, are written fresh. Everything +// else reaches the client byte for byte as upstream sent it. func (h *NPMHandler) rewriteMetadata(packageName string, body []byte) ([]byte, error) { - var metadata map[string]any - if err := json.Unmarshal(body, &metadata); err != nil { + doc, err := indexNPMPackument(body) + if err != nil { return nil, err } - // Rewrite tarball URLs in versions - versions, ok := metadata["versions"].(map[string]any) - if !ok { + versionsRaw := doc.value(doc.versionsAt) + if len(versionsRaw) == 0 || versionsRaw[0] != '{' { if len(h.proxy.Denylist.Versions(canonicalPackagePURL("npm", packageName))) != 0 { return nil, errors.New("npm metadata has no versions object") } return body, nil // No versions to rewrite } + // The filters only look at which versions exist, so they run on a map of + // version names alongside the decoded time map and dist-tags. + versions := map[string]any{} + if err := forEachJSONMember(versionsRaw, func(m jsonMember) error { + version, err := jsonKey(m.key(versionsRaw)) + versions[version] = nil + return err + }); err != nil { + return nil, err + } + metadata := map[string]any{} + timeMap := doc.decodeObject(doc.timeAt, metadata, "time") + distTags := doc.decodeObject(doc.tagsAt, metadata, "dist-tags") + timeEntries := len(timeMap) + originalTags := maps.Clone(distTags) + h.applyCooldownFiltering(metadata, versions, packageName) h.applyDenylistFiltering(metadata, versions, packageName) - h.rewriteTarballURLs(versions, packageName) - return json.Marshal(metadata) + // Untouched members are copied as they are; the three the filters + // handled are written from their filtered state, and only when it changed. + return doc.write(func(out *bytes.Buffer, i int) error { + switch { + case i == doc.versionsAt: + return h.writeNPMVersions(out, packageName, versionsRaw, versions) + case i == doc.timeAt && timeMap != nil && len(timeMap) != timeEntries: + return writeFilteredJSONObject(out, doc.value(i), func(key string) bool { + _, ok := timeMap[key] + return ok + }) + case i == doc.tagsAt && distTags != nil && !reflect.DeepEqual(distTags, originalTags): + encoded, err := json.Marshal(distTags) + out.Write(encoded) + return err + default: + out.Write(doc.value(i)) + return nil + } + }) +} + +// npmPackument is a packument's top-level members as offsets into its bytes, +// with the positions of the members rewriteMetadata filters, or -1 for those +// it lacks. A key that repeats is recorded at its last occurrence, the one +// JSON.parse keeps. +type npmPackument struct { + body []byte + members []jsonMember + versionsAt, timeAt, tagsAt int +} + +func indexNPMPackument(body []byte) (*npmPackument, error) { + if !json.Valid(body) { + return nil, errors.New("npm metadata is not valid JSON") + } + doc := &npmPackument{body: body, versionsAt: -1, timeAt: -1, tagsAt: -1} + err := forEachJSONMember(body, func(m jsonMember) error { + switch key := m.key(body); { + case jsonKeyIs(key, "versions"): + doc.versionsAt = len(doc.members) + case jsonKeyIs(key, "time"): + doc.timeAt = len(doc.members) + case jsonKeyIs(key, "dist-tags"): + doc.tagsAt = len(doc.members) + } + doc.members = append(doc.members, m) + return nil + }) + return doc, err +} + +// value returns the raw value of member i, or nil when i is -1. +func (d *npmPackument) value(i int) []byte { + if i < 0 { + return nil + } + return d.members[i].value(d.body) +} + +// decodeObject decodes member i into metadata[name] and returns it when it is +// an object. The filters edit the returned map in place. +func (d *npmPackument) decodeObject(i int, metadata map[string]any, name string) map[string]any { + raw := d.value(i) + if raw == nil { + return nil + } + var value any + if err := json.Unmarshal(raw, &value); err != nil { + return nil + } + metadata[name] = value + object, _ := value.(map[string]any) + return object +} + +// write assembles the document in member order, with writeValue writing each +// member's value. Earlier duplicates of the filtered keys are dropped: they +// would carry unfiltered data, and JSON.parse ignores them anyway. +func (d *npmPackument) write(writeValue func(out *bytes.Buffer, i int) error) ([]byte, error) { + var out bytes.Buffer + out.Grow(len(d.body) + len(d.body)/8) + out.WriteByte('{') + first := true + for i, m := range d.members { + if i != d.versionsAt && i != d.timeAt && i != d.tagsAt { + key := m.key(d.body) + if jsonKeyIs(key, "versions") || jsonKeyIs(key, "time") || jsonKeyIs(key, "dist-tags") { + continue + } + } + if !first { + out.WriteByte(',') + } + first = false + out.Write(m.key(d.body)) + out.WriteByte(':') + if err := writeValue(&out, i); err != nil { + return nil, err + } + } + out.WriteByte('}') + return out.Bytes(), nil +} + +// writeNPMVersions writes the versions object, keeping the versions still in +// keep and pointing each kept version's tarball at this proxy. +func (h *NPMHandler) writeNPMVersions(out *bytes.Buffer, packageName string, versionsRaw []byte, keep map[string]any) error { + out.WriteByte('{') + first := true + err := forEachJSONMember(versionsRaw, func(m jsonMember) error { + version, err := jsonKey(m.key(versionsRaw)) + if err != nil { + return err + } + if _, ok := keep[version]; !ok { + return nil + } + if !first { + out.WriteByte(',') + } + first = false + out.Write(m.key(versionsRaw)) + out.WriteByte(':') + h.writeNPMVersion(out, packageName, version, m.value(versionsRaw)) + return nil + }) + out.WriteByte('}') + return err +} + +// writeNPMVersion writes one version entry with its dist.tarball pointing at +// this proxy. An entry without a string tarball is written unchanged. +func (h *NPMHandler) writeNPMVersion(out *bytes.Buffer, packageName, version string, raw []byte) { + dist, ok, err := findJSONMember(raw, "dist") + if err != nil || !ok { + out.Write(raw) + return + } + distRaw := dist.value(raw) + tarball, ok, err := findJSONMember(distRaw, "tarball") + if err != nil || !ok || distRaw[tarball.valStart] != '"' { + out.Write(raw) + return + } + var oldTarball string + if err := json.Unmarshal(tarball.value(distRaw), &oldTarball); err != nil { + out.Write(raw) + return + } + + newTarball := h.proxyTarballURL(packageName, version, oldTarball) + encoded, err := json.Marshal(newTarball) + if err != nil { + out.Write(raw) + return + } + out.Write(raw[:dist.valStart+tarball.valStart]) + out.Write(encoded) + out.Write(raw[dist.valStart+tarball.valEnd:]) + + h.proxy.Logger.Debug("rewrote tarball URL", + "package", packageName, "version", version, + "old", oldTarball, "new", newTarball) +} + +// writeFilteredJSONObject writes the object obj holds, keeping the members +// whose decoded key keep accepts. +func writeFilteredJSONObject(out *bytes.Buffer, obj []byte, keep func(string) bool) error { + out.WriteByte('{') + first := true + err := forEachJSONMember(obj, func(m jsonMember) error { + key, err := jsonKey(m.key(obj)) + if err != nil { + return err + } + if !keep(key) { + return nil + } + if !first { + out.WriteByte(',') + } + first = false + out.Write(obj[m.keyStart:m.valEnd]) + return nil + }) + out.WriteByte('}') + return err } // applyCooldownFiltering removes versions that are too recently published, @@ -294,44 +512,22 @@ func (h *NPMHandler) updateDistTagsLatest(metadata, versions, timeMap map[string } } -// rewriteTarballURLs rewrites all tarball URLs in version entries to point at this proxy. -func (h *NPMHandler) rewriteTarballURLs(versions map[string]any, packageName string) { - for version, vdata := range versions { - vmap, ok := vdata.(map[string]any) - if !ok { - continue - } - - dist, ok := vmap["dist"].(map[string]any) - if !ok { - continue - } - - tarball, ok := dist["tarball"].(string) - if !ok { - continue - } - - filename := tarball - if idx := strings.LastIndex(tarball, "/"); idx >= 0 { - filename = tarball[idx+1:] +// proxyTarballURL returns the proxy URL that serves the tarball upstream +// lists at tarball for this version. +func (h *NPMHandler) proxyTarballURL(packageName, version, tarball string) string { + filename := tarball + if idx := strings.LastIndex(tarball, "/"); idx >= 0 { + filename = tarball[idx+1:] + } + if h.extractVersionFromFilename(packageName, filename) != version { + _, shortName, scoped := strings.Cut(packageName, "/") + if !scoped { + shortName = packageName } - if h.extractVersionFromFilename(packageName, filename) != version { - _, shortName, scoped := strings.Cut(packageName, "/") - if !scoped { - shortName = packageName - } - filename = shortName + "-" + version + ".tgz" - } - - escapedName := url.PathEscape(packageName) - newTarball := fmt.Sprintf("%s/npm/%s/-/%s", h.proxyURL, escapedName, filename) - dist["tarball"] = newTarball - - h.proxy.Logger.Debug("rewrote tarball URL", - "package", packageName, "version", version, - "old", tarball, "new", newTarball) + filename = shortName + "-" + version + ".tgz" } + + return fmt.Sprintf("%s/npm/%s/-/%s", h.proxyURL, url.PathEscape(packageName), filename) } // findNewestVersion returns the version string with the most recent timestamp @@ -435,18 +631,12 @@ func (h *NPMHandler) getTarball(r *http.Request, packageName, version, filename return h.proxy.GetOrFetchArtifactFromURL(r.Context(), "npm", packageName, version, filename, downloadURL) } +// npmVersionTarball returns versions[version].dist.tarball from a packument, +// or "" if it has none. It reads that one value in place rather than decoding +// every version, since it runs on each tarball the proxy has not cached yet. func npmVersionTarball(body []byte, version string) string { - var metadata struct { - Versions map[string]struct { - Dist struct { - Tarball string `json:"tarball"` - } `json:"dist"` - } `json:"versions"` - } - if err := json.Unmarshal(body, &metadata); err != nil { - return "" - } - return metadata.Versions[version].Dist.Tarball + tarball, _, _ := lookupJSONString(body, "versions", version, "dist", "tarball") + return tarball } func (h *NPMHandler) validateTarballURL(raw string) (string, error) { @@ -500,16 +690,12 @@ func (h *NPMHandler) versionInCooldown(packageName, version string, metadata fun return false } - var document struct { - Time map[string]string `json:"time"` - } - if err := json.Unmarshal(body, &document); err != nil { + published, ok, err := lookupJSONString(body, "time", version) + if err != nil { h.proxy.Logger.Warn("cooldown: could not parse npm metadata for download check", "package", packageName, "version", version, "error", err) return false } - - published, ok := document.Time[version] if !ok { return false } diff --git a/internal/handler/npm_rewrite_bench_test.go b/internal/handler/npm_rewrite_bench_test.go new file mode 100644 index 0000000..b200c5f --- /dev/null +++ b/internal/handler/npm_rewrite_bench_test.go @@ -0,0 +1,61 @@ +package handler + +import ( + "fmt" + "strings" + "testing" +) + +// benchPackument builds a full packument shaped like a large real one: +// thousands of versions, each with dependencies, scripts and a readme. +func benchPackument(versions int) []byte { + var b strings.Builder + readme := strings.Repeat("Some documentation text. ", 160) + b.WriteString(`{"_id":"big","name":"big","dist-tags":{"latest":"1.0.` + fmt.Sprint(versions-1) + `"},"versions":{`) + for i := range versions { + if i > 0 { + b.WriteByte(',') + } + v := fmt.Sprintf("1.0.%d", i) + fmt.Fprintf(&b, `%q:{"name":"big","version":%q,"description":"A big package","main":"index.js",`+ + `"scripts":{"test":"node test.js","build":"tsc -p ."},"dependencies":{`, v, v) + for d := range 25 { + if d > 0 { + b.WriteByte(',') + } + fmt.Fprintf(&b, `"dep-%d":"^%d.0.0"`, d, d) + } + fmt.Fprintf(&b, `},"readme":%q,"dist":{"shasum":"0123456789abcdef0123456789abcdef01234567",`+ + `"integrity":"sha512-AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==",`+ + `"tarball":"https://registry.npmjs.org/big/-/big-%s.tgz"}}`, readme, v) + } + b.WriteString(`},"time":{"created":"2015-01-01T00:00:00.000Z"`) + for i := range versions { + fmt.Fprintf(&b, `,"1.0.%d":"2016-01-01T00:00:00.000Z"`, i) + } + b.WriteString(`}}`) + return []byte(b.String()) +} + +func BenchmarkNPMRewriteMetadata(b *testing.B) { + body := benchPackument(5000) + h := &NPMHandler{proxy: testProxy(), proxyURL: "http://proxy.example"} + b.SetBytes(int64(len(body))) + b.ReportAllocs() + for b.Loop() { + if _, err := h.rewriteMetadata("big", body); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNPMVersionTarball(b *testing.B) { + body := benchPackument(5000) + b.SetBytes(int64(len(body))) + b.ReportAllocs() + for b.Loop() { + if npmVersionTarball(body, "1.0.2500") == "" { + b.Fatal("tarball not found") + } + } +} diff --git a/internal/handler/npm_rewrite_test.go b/internal/handler/npm_rewrite_test.go new file mode 100644 index 0000000..6503ae1 --- /dev/null +++ b/internal/handler/npm_rewrite_test.go @@ -0,0 +1,272 @@ +package handler + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/git-pkgs/cooldown" +) + +// TestNPMRewriteMetadataKeepsUpstreamBytes checks that only tarball URLs +// change: everything else, including numbers encoding/json would round +// through float64 and characters it would HTML-escape, is copied verbatim. +func TestNPMRewriteMetadataKeepsUpstreamBytes(t *testing.T) { + h := &NPMHandler{proxy: testProxy(), proxyURL: "http://proxy"} + input := `{"name":"demo","_big":12345678901234567890,"readme":"&",` + + `"versions":{"1.0.0":{"version":"1.0.0","dist":{"shasum":"s","tarball":"https://registry.npmjs.org/demo/-/demo-1.0.0.tgz","integrity":"i"}}}}` + + out, err := h.rewriteMetadata("demo", []byte(input)) + if err != nil { + t.Fatal(err) + } + want := strings.Replace(input, "https://registry.npmjs.org/demo/-/demo-1.0.0.tgz", "http://proxy/npm/demo/-/demo-1.0.0.tgz", 1) + if string(out) != want { + t.Errorf("rewrite =\n%s\nwant\n%s", out, want) + } +} + +func TestNPMRewriteMetadataUnusualEntries(t *testing.T) { + h := &NPMHandler{proxy: testProxy(), proxyURL: "http://proxy"} + input := `{ + "versions": { + "1.0.0": "not an object", + "2.0.0": {"dist": "no object either"}, + "3.0.0": {"dist": {"tarball": 7}}, + "4.0.0": {"dist": {"tarball": "https://registry.npmjs.org/demo/-/renamed.tgz"}}, + "5.0.0": {"dist": {}} + } + }` + out, err := h.rewriteMetadata("demo", []byte(input)) + if err != nil { + t.Fatal(err) + } + var doc struct { + Versions map[string]json.RawMessage `json:"versions"` + } + if err := json.Unmarshal(out, &doc); err != nil { + t.Fatalf("output is not JSON: %v\n%s", err, out) + } + for v, want := range map[string]string{ + "1.0.0": `"not an object"`, + "2.0.0": `{"dist": "no object either"}`, + "3.0.0": `{"dist": {"tarball": 7}}`, + "4.0.0": `{"dist": {"tarball": "http://proxy/npm/demo/-/demo-4.0.0.tgz"}}`, + "5.0.0": `{"dist": {}}`, + } { + if got := string(doc.Versions[v]); got != want { + t.Errorf("%s = %s, want %s", v, got, want) + } + } +} + +func TestNPMRewriteMetadataInvalidJSON(t *testing.T) { + h := &NPMHandler{proxy: testProxy(), proxyURL: "http://proxy"} + for _, input := range []string{`{"versions": {}`, `[1]`, `{"versions": {}} trailing`} { + if _, err := h.rewriteMetadata("demo", []byte(input)); err == nil { + t.Errorf("%s: expected an error", input) + } + } +} + +func TestNPMRewriteMetadataNoVersions(t *testing.T) { + h := &NPMHandler{proxy: testProxy(), proxyURL: "http://proxy"} + input := `{"name":"demo","versions":[]}` + out, err := h.rewriteMetadata("demo", []byte(input)) + if err != nil || string(out) != input { + t.Errorf("rewrite = %s, %v; want input unchanged", out, err) + } + + setTestDenylist(t, h.proxy, "pkg:npm/demo@1.0.0") + if _, err := h.rewriteMetadata("demo", []byte(input)); err == nil { + t.Error("expected an error when a denylisted package has no versions object") + } +} + +// TestNPMRewriteMetadataCooldownAndDenylist runs both filters together and +// checks that versions, their time entries and the tags pointing at them are +// all removed, with latest moved to the newest version left. +func TestNPMRewriteMetadataCooldownAndDenylist(t *testing.T) { + now := time.Now() + ts := func(age time.Duration) string { return now.Add(-age).UTC().Format(time.RFC3339) } + + proxy := testProxy() + proxy.Cooldown = &cooldown.Config{Default: "3d"} + setTestDenylist(t, proxy, "pkg:npm/demo@2.0.0") + h := &NPMHandler{proxy: proxy, proxyURL: "http://proxy"} + + version := func(v string) string { + return fmt.Sprintf(`%q:{"version":%q,"dist":{"tarball":"https://registry.npmjs.org/demo/-/demo-%s.tgz"}}`, v, v, v) + } + input := `{"name":"demo",` + + `"dist-tags":{"latest":"3.0.0","beta":"2.0.0","next":"1.0.0"},` + + `"versions":{` + version("1.0.0") + `,` + version("2.0.0") + `,` + version("3.0.0") + `},` + + `"time":{"created":"` + ts(1000*time.Hour) + `","1.0.0":"` + ts(900*time.Hour) + `","2.0.0":"` + ts(800*time.Hour) + `","3.0.0":"` + ts(time.Hour) + `"}}` + + out, err := h.rewriteMetadata("demo", []byte(input)) + if err != nil { + t.Fatal(err) + } + var doc struct { + DistTags map[string]string `json:"dist-tags"` + Versions map[string]json.RawMessage `json:"versions"` + Time map[string]string `json:"time"` + } + if err := json.Unmarshal(out, &doc); err != nil { + t.Fatalf("output is not JSON: %v\n%s", err, out) + } + if len(doc.Versions) != 1 || doc.Versions["1.0.0"] == nil { + t.Errorf("versions = %v, want only 1.0.0", keysOf(doc.Versions)) + } + if _, ok := doc.Time["2.0.0"]; ok { + t.Error("time still lists denylisted 2.0.0") + } + if _, ok := doc.Time["3.0.0"]; ok { + t.Error("time still lists 3.0.0, which is inside cooldown") + } + if _, ok := doc.Time["created"]; !ok { + t.Error("time lost its created entry") + } + want := map[string]string{"latest": "1.0.0", "next": "1.0.0"} + if fmt.Sprint(doc.DistTags) != fmt.Sprint(want) { + t.Errorf("dist-tags = %v, want %v", doc.DistTags, want) + } + if !strings.Contains(string(doc.Versions["1.0.0"]), "http://proxy/npm/demo/-/demo-1.0.0.tgz") { + t.Errorf("1.0.0 tarball not rewritten: %s", doc.Versions["1.0.0"]) + } +} + +// TestNPMRewriteMetadataDuplicateKeys checks that a repeated top-level key +// cannot smuggle an unfiltered copy of the versions past the filters. +func TestNPMRewriteMetadataDuplicateKeys(t *testing.T) { + proxy := testProxy() + setTestDenylist(t, proxy, "pkg:npm/demo@2.0.0") + h := &NPMHandler{proxy: proxy, proxyURL: "http://proxy"} + input := `{"versions":{"2.0.0":{"dist":{"tarball":"https://registry.npmjs.org/demo/-/demo-2.0.0.tgz"}}},` + + `"versions":{"1.0.0":{"dist":{"tarball":"https://registry.npmjs.org/demo/-/demo-1.0.0.tgz"}},"2.0.0":{"dist":{"tarball":"https://registry.npmjs.org/demo/-/demo-2.0.0.tgz"}}}}` + out, err := h.rewriteMetadata("demo", []byte(input)) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(out), "2.0.0") || strings.Count(string(out), `"versions"`) != 1 { + t.Errorf("rewrite = %s", out) + } +} + +func keysOf(m map[string]json.RawMessage) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + return keys +} + +// TestMetadataServesRewriteWithoutReadingStorage checks that a repeat +// request inside the metadata TTL is answered from the rewrite cache by the +// stored digest alone, without reading the cached document back. +func TestMetadataServesRewriteWithoutReadingStorage(t *testing.T) { + var upstreamCalls atomic.Int64 + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + upstreamCalls.Add(1) + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/left-pad": + _, _ = w.Write([]byte(`{"name":"left-pad","versions":{"1.3.0":{"dist":{"tarball":"https://registry.npmjs.org/left-pad/-/left-pad-1.3.0.tgz"}}}}`)) + case "/p2/vendor/pkg.json": + _, _ = w.Write([]byte(`{"packages":{"vendor/pkg":[{"version":"1.0.0","dist":{"type":"zip","url":"https://example.com/pkg-1.0.0.zip"}}]}}`)) + default: + http.NotFound(w, r) + } + })) + defer upstream.Close() + + cases := []struct { + ecosystem, name, path string + handler func(*Proxy) http.Handler + }{ + {"npm", "left-pad", "/left-pad", func(p *Proxy) http.Handler { + return NewNPMHandler(p, "http://proxy.example", upstream.URL).Routes() + }}, + {"composer", "vendor/pkg", "/p2/vendor/pkg.json", func(p *Proxy) http.Handler { + return NewComposerHandlerWithUpstreams(p, "http://proxy.example", upstream.URL, upstream.URL).Routes() + }}, + } + for _, c := range cases { + t.Run(c.ecosystem, func(t *testing.T) { + upstreamCalls.Store(0) + proxy, db, store, _ := setupTestProxy(t) + proxy.HTTPClient = upstream.Client() + proxy.CacheMetadata = true + proxy.MetadataTTL = time.Hour + proxy.SetMetadataRewriteCacheSize(1 << 20) + handler := c.handler(proxy) + + get := func() string { + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, c.path, nil)) + if rec.Code != http.StatusOK { + t.Fatalf("status %d: %s", rec.Code, rec.Body.String()) + } + return rec.Body.String() + } + + first := get() + entry, err := db.GetMetadataCache(c.ecosystem, c.name) + if err != nil || entry == nil || !strings.HasPrefix(entry.ContentDigest.String, "sha256:") { + t.Fatalf("metadata row = %+v, %v; want a sha256 content digest", entry, err) + } + + store.mu.Lock() + store.openErr = errors.New("storage must not be read") + store.mu.Unlock() + + if second := get(); second != first { + t.Errorf("second response %q differs from first %q", second, first) + } + if got := upstreamCalls.Load(); got != 1 { + t.Errorf("upstream calls = %d, want 1", got) + } + }) + } +} + +// TestNPMStoredRewriteSkippedWhenStale checks that the digest fast path is +// only taken inside the metadata TTL. +func TestNPMStoredRewriteSkippedWhenStale(t *testing.T) { + proxy, db, _, _ := setupTestProxy(t) + proxy.CacheMetadata = true + proxy.MetadataTTL = time.Hour + proxy.SetMetadataRewriteCacheSize(1 << 20) + + body := []byte(`{}`) + proxy.cacheMetadataBlob(t.Context(), "npm", "demo", metadataStoragePath("npm", "demo"), &upstreamMetadata{body: body}) + key := rewriteCacheKey("npm", "http://proxy", "demo", body) + if _, err := proxy.rewrites.rewrite(t.Context(), key, body, func(b []byte) ([]byte, error) { return b, nil }); err != nil { + t.Fatal(err) + } + if _, ok := proxy.storedRewrite("npm", "demo", "http://proxy", "demo"); !ok { + t.Fatal("fresh row: expected a stored rewrite") + } + + entry, _ := db.GetMetadataCache("npm", "demo") + entry.FetchedAt.Time = time.Now().Add(-2 * time.Hour) + if err := db.UpsertMetadataCache(entry); err != nil { + t.Fatal(err) + } + if _, ok := proxy.storedRewrite("npm", "demo", "http://proxy", "demo"); ok { + t.Error("stale row: expected no stored rewrite") + } + + proxy.Cooldown = &cooldown.Config{Default: "3d"} + entry.FetchedAt.Time = time.Now() + _ = db.UpsertMetadataCache(entry) + if _, ok := proxy.storedRewrite("npm", "demo", "http://proxy", "demo"); ok { + t.Error("cooldown on: expected no stored rewrite") + } +} diff --git a/internal/handler/rewrite_cache.go b/internal/handler/rewrite_cache.go index 7626827..23096d6 100644 --- a/internal/handler/rewrite_cache.go +++ b/internal/handler/rewrite_cache.go @@ -4,9 +4,13 @@ import ( "container/list" "context" "crypto/sha256" + "encoding/hex" "errors" "strings" "sync" + "time" + + "github.com/git-pkgs/proxy/internal/metrics" ) // rewriteCache keeps metadata documents after a handler has rewritten them. @@ -68,7 +72,28 @@ func newRewriteCache(maxBytes int64) *rewriteCache { // handler rewrites for, the package, and the exact upstream bytes. func rewriteCacheKey(ecosystem, proxyURL, name string, in []byte) string { sum := sha256.Sum256(in) - return strings.Join([]string{ecosystem, proxyURL, name, string(sum[:])}, "\x00") + return rewriteCacheKeyForDigest(ecosystem, proxyURL, name, "sha256:"+hex.EncodeToString(sum[:])) +} + +// rewriteCacheKeyForDigest is rewriteCacheKey for upstream bytes known only by +// their digest, in the form the metadata cache records it. +func rewriteCacheKeyForDigest(ecosystem, proxyURL, name, digest string) string { + return strings.Join([]string{ecosystem, proxyURL, name, digest}, "\x00") +} + +// get returns the cached rewrite for key, if there is one. +func (c *rewriteCache) get(key string) ([]byte, bool) { + if c == nil { + return nil, false + } + c.mu.Lock() + defer c.mu.Unlock() + el, ok := c.entries[key] + if !ok { + return nil, false + } + c.order.MoveToFront(el) + return el.Value.(*rewriteEntry).out, true } // rewrite returns rewrite(in), from the cache when it can. Callers must treat @@ -153,6 +178,33 @@ func (p *Proxy) cachedRewrite(ctx context.Context, ecosystem, proxyURL, name str return p.rewrites.rewrite(ctx, rewriteCacheKey(ecosystem, proxyURL, name, in), in, rewrite) } +// storedRewrite returns the cached rewrite of the metadata stored for +// ecosystem and cacheKey, when that metadata is still within its TTL and a +// rewrite of it is cached. It reads only the cache row, never the stored +// document, so a repeated request for a large packument costs a database +// lookup instead of reading and hashing the whole document again. When it +// reports false the caller takes the usual fetch and rewrite path. +func (p *Proxy) storedRewrite(ecosystem, cacheKey, proxyURL, name string) ([]byte, bool) { + if p.rewrites == nil || (p.Cooldown != nil && p.Cooldown.Enabled()) { + return nil, false + } + if !p.CacheMetadata || p.DB == nil || p.MetadataTTL <= 0 { + return nil, false + } + entry, err := p.DB.GetMetadataCache(ecosystem, cacheKey) + if err != nil || entry == nil || !entry.ContentDigest.Valid || entry.ContentEncoding.String != "" { + return nil, false + } + if !entry.FetchedAt.Valid || time.Since(entry.FetchedAt.Time) >= p.MetadataTTL { + return nil, false + } + out, ok := p.rewrites.get(rewriteCacheKeyForDigest(ecosystem, proxyURL, name, entry.ContentDigest.String)) + if ok { + metrics.RecordCacheHit(ecosystem) + } + return out, ok +} + // SetMetadataRewriteCacheSize enables the cache of rewritten metadata with // room for maxBytes of output. Zero or less disables it. func (p *Proxy) SetMetadataRewriteCacheSize(maxBytes int64) {