diff --git a/pkg/modules/api/context.go b/pkg/modules/api/context.go index c91b3551..2d93da14 100644 --- a/pkg/modules/api/context.go +++ b/pkg/modules/api/context.go @@ -216,6 +216,15 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys ) } + // 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 + // writes. + type downloadFromResult struct { + filename, path, formField string + } + results := make([]downloadFromResult, len(dls)) + eg, _ := errgroup.WithContext(ctx) for i, dl := range dls { eg.Go(func() error { @@ -392,18 +401,16 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys dlSpan.SetStatus(codes.Ok, "") dlSpan.End() - ctx.files[filename] = path - ctx.diskToOriginal[path] = filename - - // Route the downloaded file to the appropriate field bucket. + var formField string switch { case dl.Field == "embedded" || dl.Embedded: - ctx.filesByField[EmbedsFormField] = append(ctx.filesByField[EmbedsFormField], path) + formField = EmbedsFormField case dl.Field == "watermark": - ctx.filesByField[WatermarkFormField] = append(ctx.filesByField[WatermarkFormField], path) + formField = WatermarkFormField case dl.Field == "stamp": - ctx.filesByField[StampFormField] = append(ctx.filesByField[StampFormField], path) + formField = StampFormField } + results[i] = downloadFromResult{filename: filename, path: path, formField: formField} return nil }) @@ -413,6 +420,14 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys if err != nil { return ctx, cancel, err } + + for _, r := range results { + ctx.files[r.filename] = r.path + ctx.diskToOriginal[r.path] = r.filename + if r.formField != "" { + ctx.filesByField[r.formField] = append(ctx.filesByField[r.formField], r.path) + } + } } copyToDisk := func(fh *multipart.FileHeader) error { diff --git a/pkg/modules/api/context_test.go b/pkg/modules/api/context_test.go index d9f1543c..8a13ab6e 100644 --- a/pkg/modules/api/context_test.go +++ b/pkg/modules/api/context_test.go @@ -3,10 +3,13 @@ package api import ( "bytes" "context" + "encoding/json" + "fmt" "log/slog" "mime/multipart" "net/http" "net/http/httptest" + "sync" "testing" "time" @@ -70,6 +73,83 @@ func TestNewContext_Cancellation(t *testing.T) { } } +// Concurrent downloadFrom entries must not race on the shared maps +// (ctx.files, ctx.diskToOriginal, ctx.filesByField). Run under -race +// to catch the data race; without -race a sufficient number of entries +// still surfaces "fatal error: concurrent map writes". +func TestNewContext_DownloadFromConcurrentMapWrites(t *testing.T) { + const downloads = 64 + + var ready sync.WaitGroup + ready.Add(downloads) + release := make(chan struct{}) + var releaseOnce sync.Once + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ready.Done() + go func() { + ready.Wait() + releaseOnce.Do(func() { close(release) }) + }() + <-release + + 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), + Field: "embedded", + } + } + + 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{ + maxRetry: 0, + } + + 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 := len(ctx.diskToOriginal); got != downloads { + t.Fatalf("diskToOriginal entries = %d, want %d", got, downloads) + } + if got := len(ctx.filesByField[EmbedsFormField]); got != downloads { + t.Fatalf("filesByField[%q] entries = %d, want %d", EmbedsFormField, got, downloads) + } +} + func TestSanitizeFilename(t *testing.T) { for _, tc := range []struct { scenario string