diff --git a/Makefile b/Makefile index 8236d4e1..df9dcbac 100644 --- a/Makefile +++ b/Makefile @@ -38,6 +38,8 @@ API_DOWNLOAD_FROM_DENY_PRIVATE_IPS=false 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_DISABLE_DOWNLOAD_FROM=false API_DISABLE_HEALTH_CHECK_ROUTE_TELEMETRY=true API_DISABLE_ROOT_ROUTE_TELEMETRY=true diff --git a/compose.yaml b/compose.yaml index 9d2ead6c..1f5b61aa 100644 --- a/compose.yaml +++ b/compose.yaml @@ -35,6 +35,8 @@ services: - "--api-download-from-deny-public-ips=${API_DOWNLOAD_FROM_DENY_PUBLIC_IPS}" - "--api-download-from-enable-environment-proxy=${API_DOWNLOAD_FROM_ENABLE_ENVIRONMENT_PROXY}" - "--api-download-from-max-retry=${API_DOWNLOAD_FROM_MAX_RETRY}" + - "--api-download-from-max-concurrency=${API_DOWNLOAD_FROM_MAX_CONCURRENCY}" + - "--api-download-from-max-entries=${API_DOWNLOAD_FROM_MAX_ENTRIES}" - "--api-disable-download-from=${API_DISABLE_DOWNLOAD_FROM}" - "--api-disable-health-check-route-telemetry=${API_DISABLE_HEALTH_CHECK_ROUTE_TELEMETRY}" - "--api-disable-root-route-telemetry=${API_DISABLE_ROOT_ROUTE_TELEMETRY}" diff --git a/pkg/modules/api/api.go b/pkg/modules/api/api.go index e9cd4048..c5a1b590 100644 --- a/pkg/modules/api/api.go +++ b/pkg/modules/api/api.go @@ -67,6 +67,8 @@ type downloadFromConfig struct { denyPublicIPs bool enableEnvironmentProxy bool maxRetry int + maxConcurrency int + maxEntries int disable bool } @@ -212,6 +214,8 @@ func (a *Api) Descriptor() gotenberg.ModuleDescriptor { fs.Bool("api-download-from-deny-public-ips", false, "Reject downloadFrom URLs whose host resolves to a public IP address. Enable on air-gapped or data-governed deployments to prevent downloads from reaching the public internet") 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.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") @@ -255,6 +259,8 @@ func (a *Api) Provision(ctx *gotenberg.Context) error { denyPublicIPs: flags.MustBool("api-download-from-deny-public-ips"), enableEnvironmentProxy: flags.MustBool("api-download-from-enable-environment-proxy"), maxRetry: flags.MustInt("api-download-from-max-retry"), + maxConcurrency: flags.MustInt("api-download-from-max-concurrency"), + maxEntries: flags.MustInt("api-download-from-max-entries"), disable: flags.MustBool("api-disable-download-from"), } a.disableHealthCheckRouteTelemetry = flags.MustDeprecatedBool("api-disable-health-check-logging", "api-disable-health-check-route-telemetry") @@ -404,6 +410,18 @@ func (a *Api) Validate() error { } } + if a.downloadFromCfg.maxConcurrency < 0 { + err = errors.Join(err, + fmt.Errorf("download from max concurrency must not be negative, got %d; set --api-download-from-max-concurrency (env API_DOWNLOAD_FROM_MAX_CONCURRENCY) to 0 to disable the limit", a.downloadFromCfg.maxConcurrency), + ) + } + + if a.downloadFromCfg.maxEntries < 0 { + err = errors.Join(err, + fmt.Errorf("download from max entries must not be negative, got %d; set --api-download-from-max-entries (env API_DOWNLOAD_FROM_MAX_ENTRIES) to 0 to disable the limit", a.downloadFromCfg.maxEntries), + ) + } + if (a.tlsCertFile != "" && a.tlsKeyFile == "") || (a.tlsCertFile == "" && a.tlsKeyFile != "") { err = errors.Join(err, errors.New("both TLS certificate and key files must be set"), diff --git a/pkg/modules/api/context.go b/pkg/modules/api/context.go index 64bb3acc..87a053ec 100644 --- a/pkg/modules/api/context.go +++ b/pkg/modules/api/context.go @@ -222,6 +222,16 @@ 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 @@ -232,6 +242,13 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys results := make([]downloadFromResult, len(dls)) eg, _ := errgroup.WithContext(ctx) + // Bound the number of in-flight downloads. Each entry allocates a + // retryable client, an outbound transport, a span, and logger state, + // so an unbounded array would otherwise exhaust process memory. A + // value of 0 keeps the fan-out unbounded. + if downloadFromCfg.maxConcurrency > 0 { + eg.SetLimit(downloadFromCfg.maxConcurrency) + } for i, dl := range dls { eg.Go(func() error { deadline, ok := ctx.Deadline() diff --git a/pkg/modules/api/context_test.go b/pkg/modules/api/context_test.go index 96a8a900..308d5108 100644 --- a/pkg/modules/api/context_test.go +++ b/pkg/modules/api/context_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "log/slog" "mime/multipart" @@ -11,6 +12,7 @@ import ( "net/http/httptest" "os" "sync" + "sync/atomic" "testing" "time" @@ -209,6 +211,137 @@ func TestNewContext_DownloadFromConcurrentMapWrites(t *testing.T) { } } +// An oversized downloadFrom array must be rejected at the trust boundary with +// a 400, before any download goroutine is spawned. +// https://github.com/gotenberg/gotenberg/security/advisories/GHSA-6vqw-2jgm-4x88 +func TestNewContext_DownloadFromMaxEntries(t *testing.T) { + var hits atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + w.Header().Set("Content-Disposition", `attachment; filename="download.txt"`) + _, _ = w.Write([]byte("downloaded")) + })) + defer server.Close() + + dls := make([]downloadFrom, 3) + for i := range dls { + dls[i] = downloadFrom{Url: fmt.Sprintf("%s/file?i=%d", server.URL, i)} + } + + payload, err := json.Marshal(dls) + if err != nil { + t.Fatalf("marshal downloadFrom payload: %v", err) + } + + body := new(bytes.Buffer) + writer := multipart.NewWriter(body) + err = writer.WriteField("downloadFrom", string(payload)) + if err != nil { + t.Fatalf("write downloadFrom field: %v", err) + } + err = writer.Close() + if err != nil { + t.Fatalf("close multipart writer: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/forms/libreoffice/convert", body) + req.Header.Set("Content-Type", writer.FormDataContentType()) + + echoCtx := echo.New().NewContext(req, httptest.NewRecorder()) + logger := slog.New(slog.DiscardHandler) + fs := gotenberg.NewFileSystem(new(gotenberg.OsMkdirAll)) + downloadFromCfg := downloadFromConfig{maxEntries: 2} + + _, cancel, err := newContext(echoCtx, logger, fs, 10*time.Second, 0, downloadFromCfg) + if cancel != nil { + defer cancel() + } + if err == nil { + t.Fatal("newContext returned no error, want a 400 for too many entries") + } + + var httpErr HttpError + if !errors.As(err, &httpErr) { + t.Fatalf("error %v is not an HttpError", err) + } + if status, _ := httpErr.HttpError(); status != http.StatusBadRequest { + t.Fatalf("HTTP status = %d, want %d", status, http.StatusBadRequest) + } + if got := hits.Load(); got != 0 { + t.Fatalf("server hits = %d, want 0 (rejected before any download)", got) + } +} + +// The number of in-flight downloadFrom fetches must never exceed the +// configured concurrency limit, regardless of array length. +// https://github.com/gotenberg/gotenberg/security/advisories/GHSA-6vqw-2jgm-4x88 +func TestNewContext_DownloadFromMaxConcurrency(t *testing.T) { + const ( + downloads = 8 + maxConcurrency = 2 + ) + + var current, peak atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + inFlight := current.Add(1) + for { + observed := peak.Load() + if inFlight <= observed || peak.CompareAndSwap(observed, inFlight) { + break + } + } + time.Sleep(20 * time.Millisecond) + current.Add(-1) + + filename := fmt.Sprintf("download-%s.txt", r.URL.Query().Get("i")) + w.Header().Set("Content-Disposition", fmt.Sprintf(`attachment; filename="%s"`, filename)) + _, _ = w.Write([]byte("downloaded")) + })) + defer server.Close() + + dls := make([]downloadFrom, downloads) + for i := range dls { + dls[i] = downloadFrom{Url: fmt.Sprintf("%s/file?i=%d", server.URL, i)} + } + + payload, err := json.Marshal(dls) + if err != nil { + t.Fatalf("marshal downloadFrom payload: %v", err) + } + + body := new(bytes.Buffer) + writer := multipart.NewWriter(body) + err = writer.WriteField("downloadFrom", string(payload)) + if err != nil { + t.Fatalf("write downloadFrom field: %v", err) + } + err = writer.Close() + if err != nil { + t.Fatalf("close multipart writer: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/forms/libreoffice/convert", body) + req.Header.Set("Content-Type", writer.FormDataContentType()) + + echoCtx := echo.New().NewContext(req, httptest.NewRecorder()) + logger := slog.New(slog.DiscardHandler) + fs := gotenberg.NewFileSystem(new(gotenberg.OsMkdirAll)) + downloadFromCfg := downloadFromConfig{maxConcurrency: maxConcurrency} + + ctx, cancel, err := newContext(echoCtx, logger, fs, 10*time.Second, 0, downloadFromCfg) + if err != nil { + t.Fatalf("newContext returned error: %v", err) + } + defer cancel() + + if got := len(ctx.files); got != downloads { + t.Fatalf("downloaded files = %d, want %d", got, downloads) + } + if got := peak.Load(); got > maxConcurrency { + t.Fatalf("peak concurrency = %d, want <= %d", got, maxConcurrency) + } +} + func TestSanitizeFilename(t *testing.T) { for _, tc := range []struct { scenario string