From d9db46b80d96b2bedf8b2fb19fa77f0096d96415 Mon Sep 17 00:00:00 2001 From: Simon Chrzanowski Date: Wed, 7 Oct 2026 10:04:43 +0200 Subject: [PATCH] Cut memory use of npm metadata rewriting Rewriting a packument decoded the whole document into map[string]any and encoded it again. For a full packument of tens of megabytes that costs many times the document size per request, and with cooldown enabled (which skips the rewrite cache and forces full packuments) it happens on every request. - Rewrite npm metadata in place: walk the raw JSON for member offsets, decode only the version names, time and dist-tags for the existing filters, and assemble the response from slices of the original bytes with only tarball URLs (and filtered time/dist-tags) written fresh. Untouched values now reach clients byte for byte instead of round-tripping through encoding/json. - Look up versions[v].dist.tarball and time[v] directly on the tarball download path instead of decoding every version. - Record the stored metadata's SHA-256 in metadata_cache.content_digest and key the rewrite cache by it, so a repeat npm or Composer request inside the TTL is served from the rewrite cache without reading or hashing the stored document. On a synthetic 25 MB packument with 5000 versions, rewriteMetadata drops from ~150-174 MB / 706k allocations per call to ~31 MB / 90k, and runs in about two thirds of the time. npmVersionTarball drops from ~1 MB / 10k allocations to one allocation. --- internal/handler/composer.go | 6 + internal/handler/handler.go | 11 +- internal/handler/jsonscan.go | 213 ++++++++++++++ internal/handler/jsonscan_test.go | 98 +++++++ internal/handler/npm.go | 306 +++++++++++++++++---- internal/handler/npm_rewrite_bench_test.go | 61 ++++ internal/handler/npm_rewrite_test.go | 272 ++++++++++++++++++ internal/handler/rewrite_cache.go | 54 +++- 8 files changed, 956 insertions(+), 65 deletions(-) create mode 100644 internal/handler/jsonscan.go create mode 100644 internal/handler/jsonscan_test.go create mode 100644 internal/handler/npm_rewrite_bench_test.go create mode 100644 internal/handler/npm_rewrite_test.go 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) {