repository: fix race condition for blobSaver shutdown

wg.Go() may not be called after wg.Wait(). This prevents connecting two
errgroups such that the errors are propagated between them if the child
errgroup dynamically starts goroutines. Instead use just a single errgroup,
and sequence the shutdown using a sync.WaitGroup. This is far simpler
and does not require any "clever" tricks.
This commit is contained in:
Michael Eischer
2025-11-26 21:18:22 +01:00
parent 9f87e9096a
commit 5607fd759f
2 changed files with 32 additions and 52 deletions
+12 -17
View File
@@ -2104,8 +2104,10 @@ type failSaveRepo struct {
} }
func (f *failSaveRepo) WithBlobUploader(ctx context.Context, fn func(ctx context.Context, uploader restic.BlobSaverWithAsync) error) error { func (f *failSaveRepo) WithBlobUploader(ctx context.Context, fn func(ctx context.Context, uploader restic.BlobSaverWithAsync) error) error {
return f.archiverRepo.WithBlobUploader(ctx, func(ctx context.Context, uploader restic.BlobSaverWithAsync) error { outerCtx, outerCancel := context.WithCancelCause(ctx)
return fn(ctx, &failSaveSaver{saver: uploader, failSaveRepo: f, semaphore: make(chan struct{}, 1)}) defer outerCancel(f.err)
return f.archiverRepo.WithBlobUploader(outerCtx, func(ctx context.Context, uploader restic.BlobSaverWithAsync) error {
return fn(ctx, &failSaveSaver{saver: uploader, failSaveRepo: f, semaphore: make(chan struct{}, 1), outerCancel: outerCancel})
}) })
} }
@@ -2113,6 +2115,7 @@ type failSaveSaver struct {
saver restic.BlobSaverWithAsync saver restic.BlobSaverWithAsync
failSaveRepo *failSaveRepo failSaveRepo *failSaveRepo
semaphore chan struct{} semaphore chan struct{}
outerCancel context.CancelCauseFunc
} }
func (f *failSaveSaver) SaveBlob(ctx context.Context, t restic.BlobType, buf []byte, id restic.ID, storeDuplicate bool) (restic.ID, bool, int, error) { func (f *failSaveSaver) SaveBlob(ctx context.Context, t restic.BlobType, buf []byte, id restic.ID, storeDuplicate bool) (restic.ID, bool, int, error) {
@@ -2130,10 +2133,9 @@ func (f *failSaveSaver) SaveBlobAsync(ctx context.Context, t restic.BlobType, bu
val := f.failSaveRepo.cnt.Add(1) val := f.failSaveRepo.cnt.Add(1)
if val >= f.failSaveRepo.failAfter { if val >= f.failSaveRepo.failAfter {
// use a canceled context to make SaveBlobAsync fail // kill the outer context to make SaveBlobAsync fail
var cancel context.CancelCauseFunc // precisely injecting a specific error into the repository is not possible, so just cancel the context
ctx, cancel = context.WithCancelCause(ctx) f.outerCancel(f.failSaveRepo.err)
cancel(f.failSaveRepo.err)
} }
f.saver.SaveBlobAsync(ctx, t, buf, id, storeDuplicate, func(newID restic.ID, known bool, size int, err error) { f.saver.SaveBlobAsync(ctx, t, buf, id, storeDuplicate, func(newID restic.ID, known bool, size int, err error) {
@@ -2141,7 +2143,6 @@ func (f *failSaveSaver) SaveBlobAsync(ctx context.Context, t restic.BlobType, bu
if err == nil { if err == nil {
panic("expected error") panic("expected error")
} }
err = f.failSaveRepo.err
} }
cb(newID, known, size, err) cb(newID, known, size, err)
<-f.semaphore <-f.semaphore
@@ -2149,13 +2150,10 @@ func (f *failSaveSaver) SaveBlobAsync(ctx context.Context, t restic.BlobType, bu
} }
func TestArchiverAbortEarlyOnError(t *testing.T) { func TestArchiverAbortEarlyOnError(t *testing.T) {
var testErr = errors.New("test error")
var tests = []struct { var tests = []struct {
src TestDir src TestDir
wantOpen map[string]uint wantOpen map[string]uint
failAfter uint // error after so many blobs have been saved to the repo failAfter uint // error after so many blobs have been saved to the repo
err error
}{ }{
{ {
src: TestDir{ src: TestDir{
@@ -2167,10 +2165,7 @@ func TestArchiverAbortEarlyOnError(t *testing.T) {
}, },
wantOpen: map[string]uint{ wantOpen: map[string]uint{
filepath.FromSlash("dir/bar"): 1, filepath.FromSlash("dir/bar"): 1,
filepath.FromSlash("dir/baz"): 1,
filepath.FromSlash("dir/foo"): 1,
}, },
err: testErr,
}, },
{ {
src: TestDir{ src: TestDir{
@@ -2198,7 +2193,6 @@ func TestArchiverAbortEarlyOnError(t *testing.T) {
// fails after four to seven files were opened, as the ReadConcurrency allows for // fails after four to seven files were opened, as the ReadConcurrency allows for
// two queued files and one blob queued for saving. // two queued files and one blob queued for saving.
failAfter: 4, failAfter: 4,
err: testErr,
}, },
} }
@@ -2217,10 +2211,11 @@ func TestArchiverAbortEarlyOnError(t *testing.T) {
opened: make(map[string]uint), opened: make(map[string]uint),
} }
testErr := context.Canceled
testRepo := &failSaveRepo{ testRepo := &failSaveRepo{
archiverRepo: repo, archiverRepo: repo,
failAfter: int32(test.failAfter), failAfter: int32(test.failAfter),
err: test.err, err: testErr,
} }
// at most two files may be queued // at most two files may be queued
@@ -2233,8 +2228,8 @@ func TestArchiverAbortEarlyOnError(t *testing.T) {
} }
_, _, _, err := arch.Snapshot(ctx, []string{"."}, SnapshotOptions{Time: time.Now()}) _, _, _, err := arch.Snapshot(ctx, []string{"."}, SnapshotOptions{Time: time.Now()})
if !errors.Is(err, test.err) { if !errors.Is(err, testErr) {
t.Errorf("expected error (%v) not found, got %v", test.err, err) t.Errorf("expected error (%v) not found, got %v", testErr, err)
} }
t.Logf("Snapshot return error: %v", err) t.Logf("Snapshot return error: %v", err)
+20 -35
View File
@@ -42,7 +42,8 @@ type Repository struct {
opts Options opts Options
packerWg *errgroup.Group packerWg *errgroup.Group
blobWg *errgroup.Group mainWg *errgroup.Group
blobSaver *sync.WaitGroup
uploader *packerUploader uploader *packerUploader
treePM *packerManager treePM *packerManager
dataPM *packerManager dataPM *packerManager
@@ -562,12 +563,14 @@ func (r *Repository) removeUnpacked(ctx context.Context, t restic.FileType, id r
func (r *Repository) WithBlobUploader(ctx context.Context, fn func(ctx context.Context, uploader restic.BlobSaverWithAsync) error) error { func (r *Repository) WithBlobUploader(ctx context.Context, fn func(ctx context.Context, uploader restic.BlobSaverWithAsync) error) error {
wg, ctx := errgroup.WithContext(ctx) wg, ctx := errgroup.WithContext(ctx)
// pack uploader + wg.Go below + blob saver (CPU bound)
wg.SetLimit(2 + runtime.GOMAXPROCS(0))
r.mainWg = wg
r.startPackUploader(ctx, wg) r.startPackUploader(ctx, wg)
saverCtx := r.startBlobSaver(ctx, wg) // blob saver are spawned on demand, use wait group to keep track of them
r.blobSaver = &sync.WaitGroup{}
wg.Go(func() error { wg.Go(func() error {
// must use saverCtx to ensure that the ctx used for saveBlob calls is bound to it if err := fn(ctx, &blobSaverRepo{repo: r}); err != nil {
// otherwise the blob saver could deadlock in case of an error.
if err := fn(saverCtx, &blobSaverRepo{repo: r}); err != nil {
return err return err
} }
if err := r.flush(ctx); err != nil { if err := r.flush(ctx); err != nil {
@@ -594,22 +597,6 @@ func (r *Repository) startPackUploader(ctx context.Context, wg *errgroup.Group)
}) })
} }
func (r *Repository) startBlobSaver(ctx context.Context, wg *errgroup.Group) context.Context {
// blob upload computations are CPU bound
blobWg, blobCtx := errgroup.WithContext(ctx)
blobWg.SetLimit(runtime.GOMAXPROCS(0))
r.blobWg = blobWg
wg.Go(func() error {
// As the goroutines are only spawned on demand, wait until the context is canceled.
// This will either happen on an error while saving a blob or when blobWg.Wait() is called
// by flushBlobUploader().
<-blobCtx.Done()
return blobWg.Wait()
})
return blobCtx
}
type blobSaverRepo struct { type blobSaverRepo struct {
repo *Repository repo *Repository
} }
@@ -624,28 +611,26 @@ func (r *blobSaverRepo) SaveBlobAsync(ctx context.Context, t restic.BlobType, bu
// Flush saves all remaining packs and the index // Flush saves all remaining packs and the index
func (r *Repository) flush(ctx context.Context) error { func (r *Repository) flush(ctx context.Context) error {
if err := r.flushBlobUploader(); err != nil { r.flushBlobSaver()
return err r.mainWg = nil
}
if err := r.flushPacks(ctx); err != nil { if err := r.flushPackUploader(ctx); err != nil {
return err return err
} }
return r.idx.Flush(ctx, &internalRepository{r}) return r.idx.Flush(ctx, &internalRepository{r})
} }
func (r *Repository) flushBlobUploader() error { func (r *Repository) flushBlobSaver() {
if r.blobWg == nil { if r.blobSaver == nil {
return nil return
} }
err := r.blobWg.Wait() r.blobSaver.Wait()
r.blobWg = nil r.blobSaver = nil
return err
} }
// FlushPacks saves all remaining packs. // FlushPacks saves all remaining packs.
func (r *Repository) flushPacks(ctx context.Context) error { func (r *Repository) flushPackUploader(ctx context.Context) error {
if r.packerWg == nil { if r.packerWg == nil {
return nil return nil
} }
@@ -1032,11 +1017,11 @@ func (r *Repository) saveBlob(ctx context.Context, t restic.BlobType, buf []byte
} }
func (r *Repository) saveBlobAsync(ctx context.Context, t restic.BlobType, buf []byte, id restic.ID, storeDuplicate bool, cb func(newID restic.ID, known bool, size int, err error)) { func (r *Repository) saveBlobAsync(ctx context.Context, t restic.BlobType, buf []byte, id restic.ID, storeDuplicate bool, cb func(newID restic.ID, known bool, size int, err error)) {
r.blobWg.Go(func() error { r.mainWg.Go(func() error {
if ctx.Err() != nil { if ctx.Err() != nil {
// fail fast if the context is cancelled // fail fast if the context is cancelled
cb(restic.ID{}, false, 0, context.Cause(ctx)) cb(restic.ID{}, false, 0, ctx.Err())
return context.Cause(ctx) return ctx.Err()
} }
newID, known, size, err := r.saveBlob(ctx, t, buf, id, storeDuplicate) newID, known, size, err := r.saveBlob(ctx, t, buf, id, storeDuplicate)
cb(newID, known, size, err) cb(newID, known, size, err)