diff --git a/Makefile b/Makefile index 85dace02..cd7242c1 100644 --- a/Makefile +++ b/Makefile @@ -42,7 +42,7 @@ API_DOWNLOAD_FROM_DENY_PUBLIC_IPS=false API_DOWNLOAD_FROM_ENABLE_ENVIRONMENT_PROXY=false API_DOWNLOAD_FROM_MAX_RETRY=4 API_DOWNLOAD_FROM_MAX_CONCURRENCY=10 -API_DOWNLOAD_FROM_MAX_ENTRIES=0 +API_DOWNLOAD_FROM_MAX_ENTRIES=1000 API_DISABLE_DOWNLOAD_FROM=false API_DISABLE_HEALTH_CHECK_ROUTE_TELEMETRY=true API_DISABLE_ROOT_ROUTE_TELEMETRY=true diff --git a/pkg/modules/api/api.go b/pkg/modules/api/api.go index 0b4506f7..663a24e4 100644 --- a/pkg/modules/api/api.go +++ b/pkg/modules/api/api.go @@ -215,7 +215,7 @@ func (a *Api) Descriptor() gotenberg.ModuleDescriptor { fs.Bool("api-download-from-enable-environment-proxy", false, "Route downloadFrom fetches through the proxy defined by the standard HTTP_PROXY, HTTPS_PROXY, and NO_PROXY variables, including credentials") fs.Int("api-download-from-max-retry", 4, "Set the maximum number of retries for the download from feature") fs.Int("api-download-from-max-concurrency", 10, "Set the maximum number of downloadFrom entries fetched concurrently per request - bounds the outbound fan-out. Set to 0 to disable this feature") - fs.Int("api-download-from-max-entries", 0, "Set the maximum number of downloadFrom entries allowed per request. Set to 0 to disable this feature") + fs.Int("api-download-from-max-entries", 1000, "Set the maximum number of downloadFrom entries allowed per request. Set to 0 to disable this limit, which lets a single request expand into an arbitrarily large array") fs.Bool("api-disable-download-from", false, "Disable the download from feature") fs.Bool("api-disable-health-check-route-telemetry", true, "Disable telemetry for health check route") fs.Bool("api-disable-root-route-telemetry", true, "Disable telemetry for the root route") diff --git a/pkg/modules/api/context.go b/pkg/modules/api/context.go index cc46832b..4f11bfdd 100644 --- a/pkg/modules/api/context.go +++ b/pkg/modules/api/context.go @@ -80,6 +80,48 @@ func (t *trackingReader) Read(p []byte) (int, error) { return n, nil } +// errTooManyDownloadFromEntries is returned by [decodeDownloadFrom] when the +// array holds more entries than the configured maximum. +var errTooManyDownloadFromEntries = errors.New("too many downloadFrom entries") + +// decodeDownloadFrom decodes the downloadFrom form field, refusing to +// accumulate more than maxEntries. A maxEntries of 0 means no limit. +// +// It decodes element by element rather than calling [json.Unmarshal] on the +// whole value. A compact array such as "[{},{},{}]" costs three bytes per +// entry on the wire and expands to roughly seventy times that once +// unmarshalled, so counting the entries afterwards is too late to bound the +// allocation. Streaming keeps the cost proportional to maxEntries no matter +// how long the array is. +func decodeDownloadFrom(raw string, maxEntries int) ([]downloadFrom, error) { + dec := json.NewDecoder(strings.NewReader(raw)) + + token, err := dec.Token() + if err != nil { + return nil, err + } + if delim, ok := token.(json.Delim); !ok || delim != '[' { + return nil, fmt.Errorf("expected a JSON array, got '%v'", token) + } + + var dls []downloadFrom + for dec.More() { + if maxEntries > 0 && len(dls) >= maxEntries { + return nil, errTooManyDownloadFromEntries + } + + var dl downloadFrom + err = dec.Decode(&dl) + if err != nil { + return nil, err + } + + dls = append(dls, dl) + } + + return dls, nil +} + type downloadFrom struct { // Url is the URL to download a file from. Url string `json:"url"` @@ -213,8 +255,13 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys // any. raw, ok := ctx.values["downloadFrom"] if !downloadFromCfg.disable && ok { - var dls []downloadFrom - err = json.Unmarshal([]byte(raw[0]), &dls) + dls, err := decodeDownloadFrom(raw[0], downloadFromCfg.maxEntries) + if errors.Is(err, errTooManyDownloadFromEntries) { + return nil, cancel, WrapError( + fmt.Errorf("decode downloadFrom: %w", err), + NewSentinelHttpError(http.StatusBadRequest, fmt.Sprintf("Invalid 'downloadFrom' form field value: too many entries, the maximum is %d", downloadFromCfg.maxEntries)), + ) + } if err != nil { return nil, cancel, WrapError( fmt.Errorf("unmarshal json: %w", err), @@ -222,16 +269,6 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys ) } - // Reject oversized arrays at the trust boundary, before allocating - // the results slice or spawning any goroutine, so a compact request - // cannot inflate into an unbounded fan-out. - if downloadFromCfg.maxEntries > 0 && len(dls) > downloadFromCfg.maxEntries { - return nil, cancel, WrapError( - fmt.Errorf("too many downloadFrom entries: got %d, max %d", len(dls), downloadFromCfg.maxEntries), - NewSentinelHttpError(http.StatusBadRequest, fmt.Sprintf("Invalid 'downloadFrom' form field value: too many entries, the maximum is %d", downloadFromCfg.maxEntries)), - ) - } - // Each goroutine writes to its own results slot. The main // goroutine merges into ctx.files, ctx.diskToOriginal, and // ctx.filesByField after eg.Wait() to avoid concurrent map diff --git a/pkg/modules/api/context_test.go b/pkg/modules/api/context_test.go index c205a69d..064a739c 100644 --- a/pkg/modules/api/context_test.go +++ b/pkg/modules/api/context_test.go @@ -11,6 +11,8 @@ import ( "net/http" "net/http/httptest" "os" + "runtime" + "strings" "sync" "sync/atomic" "testing" @@ -528,3 +530,75 @@ func TestNewContext_DownloadFromExpiredBudgetFailsClosed(t *testing.T) { t.Fatal("newContext never returned: an entry starting past the deadline built an unbounded client") } } + +func TestDecodeDownloadFrom(t *testing.T) { + for _, tc := range []struct { + scenario string + raw string + maxEntries int + expectErr error + expectLen int + }{ + {"empty array", `[]`, 10, nil, 0}, + {"under the limit", `[{"url":"http://a"},{"url":"http://b"}]`, 10, nil, 2}, + {"exactly the limit", `[{"url":"http://a"},{"url":"http://b"}]`, 2, nil, 2}, + {"over the limit", `[{"url":"http://a"},{"url":"http://b"}]`, 1, errTooManyDownloadFromEntries, 0}, + {"no limit", `[{"url":"http://a"},{"url":"http://b"}]`, 0, nil, 2}, + {"not an array", `{"url":"http://a"}`, 10, nil, 0}, + {"malformed", `[{"url":`, 10, nil, 0}, + {"not json", `nope`, 10, nil, 0}, + } { + t.Run(tc.scenario, func(t *testing.T) { + dls, err := decodeDownloadFrom(tc.raw, tc.maxEntries) + + if tc.expectErr != nil { + if !errors.Is(err, tc.expectErr) { + t.Fatalf("error = %v, want %v", err, tc.expectErr) + } + return + } + if tc.scenario == "not an array" || tc.scenario == "malformed" || tc.scenario == "not json" { + if err == nil { + t.Fatalf("expected an error for %q", tc.raw) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(dls) != tc.expectLen { + t.Fatalf("decoded %d entries, want %d", len(dls), tc.expectLen) + } + }) + } +} + +// A compact array costs three bytes per entry on the wire and expands by +// roughly seventy times once unmarshalled. Decoding must stop at the limit +// rather than materialize the whole array and count afterwards. +func TestDecodeDownloadFrom_StopsBeforeMaterializingTheArray(t *testing.T) { + const entries = 2_000_000 + + raw := "[" + strings.Repeat("{},", entries) + "{}]" + + var before, after runtime.MemStats + runtime.GC() + runtime.ReadMemStats(&before) + + _, err := decodeDownloadFrom(raw, 1000) + + runtime.ReadMemStats(&after) + + if !errors.Is(err, errTooManyDownloadFromEntries) { + t.Fatalf("error = %v, want errTooManyDownloadFromEntries", err) + } + + // json.Unmarshal on the same input allocates hundreds of MiB. Bounded + // decoding should stay in the low single-digit MiB, so this threshold is + // deliberately loose and still fails loudly on a regression. + allocated := after.TotalAlloc - before.TotalAlloc + if allocated > 32<<20 { + t.Fatalf("decoding allocated %d MiB for a %d-entry array, want the limit to bound it", allocated>>20, entries) + } + t.Logf("allocated %d KiB decoding a %d-entry array with a limit of 1000", allocated>>10, entries) +}