fix(api): serialize downloadFrom result merging to avoid concurrent map writes

This commit is contained in:
Julien Neuhart
2026-05-12 19:25:25 +02:00
parent f9a01c9fb3
commit 6671b5e5d3
2 changed files with 102 additions and 7 deletions

View File

@@ -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) eg, _ := errgroup.WithContext(ctx)
for i, dl := range dls { for i, dl := range dls {
eg.Go(func() error { eg.Go(func() error {
@@ -392,18 +401,16 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys
dlSpan.SetStatus(codes.Ok, "") dlSpan.SetStatus(codes.Ok, "")
dlSpan.End() dlSpan.End()
ctx.files[filename] = path var formField string
ctx.diskToOriginal[path] = filename
// Route the downloaded file to the appropriate field bucket.
switch { switch {
case dl.Field == "embedded" || dl.Embedded: case dl.Field == "embedded" || dl.Embedded:
ctx.filesByField[EmbedsFormField] = append(ctx.filesByField[EmbedsFormField], path) formField = EmbedsFormField
case dl.Field == "watermark": case dl.Field == "watermark":
ctx.filesByField[WatermarkFormField] = append(ctx.filesByField[WatermarkFormField], path) formField = WatermarkFormField
case dl.Field == "stamp": 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 return nil
}) })
@@ -413,6 +420,14 @@ func newContext(echoCtx echo.Context, logger *slog.Logger, fs *gotenberg.FileSys
if err != nil { if err != nil {
return ctx, cancel, err 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 { copyToDisk := func(fh *multipart.FileHeader) error {

View File

@@ -3,10 +3,13 @@ package api
import ( import (
"bytes" "bytes"
"context" "context"
"encoding/json"
"fmt"
"log/slog" "log/slog"
"mime/multipart" "mime/multipart"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"sync"
"testing" "testing"
"time" "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) { func TestSanitizeFilename(t *testing.T) {
for _, tc := range []struct { for _, tc := range []struct {
scenario string scenario string