From aebb09b46a6820cacbd0cadbfbd41b10e9c44f40 Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Tue, 24 Mar 2026 11:10:29 -0400 Subject: [PATCH 1/7] feat: Blob adding is better parallelized Rather than waiting for entire batches to complete, we re-query whenever there's available parallelism, until we find no more blobs. Note that this *only* works when new blobs become available while we're waiting to send some into `addBlobs`. If blobs become available between `addBlobs` consuming all the ones we had (beginning their uploads) and all of those uploads completing, we won't catch them until the next round one level up. In that case, we still waste parallel capacity. Therefore, this is still an incomplete solution, but moves in the right direction. --- pkg/preparation/storacha/storacha.go | 118 +++++++++++++++++++-------- 1 file changed, 82 insertions(+), 36 deletions(-) diff --git a/pkg/preparation/storacha/storacha.go b/pkg/preparation/storacha/storacha.go index b51806b1..b0d48f48 100644 --- a/pkg/preparation/storacha/storacha.go +++ b/pkg/preparation/storacha/storacha.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "sync" "time" commp "github.com/filecoin-project/go-fil-commp-hashhash" @@ -75,22 +76,41 @@ var _ uploads.AddStorachaUploadForUploadFunc = API{}.AddStorachaUploadForUpload func (a API) AddShardsForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, shardUploadedCb func(shard *model.Shard) error) error { ctx, span := tracer.Start(ctx, "add-shards-for-upload") defer span.End() - closedShards, err := a.Repo.ShardsForUploadByState(ctx, uploadID, model.BlobStateClosed) - if err != nil { - return fmt.Errorf("failed to get closed shards for upload %s: %w", uploadID, err) - } - span.AddEvent("found closed shards", trace.WithAttributes(attribute.Int("shards", len(closedShards)))) - blobs := make([]model.Blob, len(closedShards)) - for i, shard := range closedShards { - blobs[i] = shard - } - return a.addBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { + blobs := make(chan model.Blob) + + errorCh := make(chan error, 1) + go func() { + defer close(blobs) + defer close(errorCh) + for { + closedShards, err := a.Repo.ShardsForUploadByState(ctx, uploadID, model.BlobStateClosed) + if err != nil { + errorCh <- fmt.Errorf("failed to get closed shards for upload %s: %w", uploadID, err) + return + } + span.AddEvent("found closed shards", trace.WithAttributes(attribute.Int("shards", len(closedShards)))) + + if len(closedShards) == 0 { + return + } + + for _, shard := range closedShards { + blobs <- shard + } + } + }() + + addError := a.addBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { if shardUploadedCb != nil { return shardUploadedCb(blob.(*model.Shard)) } return nil }) + + findError := <-errorCh + + return errors.Join(findError, addError) } func (a API) PostProcessUploadedShards(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error { @@ -119,20 +139,27 @@ func (a API) PostProcessUploadedShards(ctx context.Context, uploadID id.UploadID } // addBlobs adds the given blobs to the space, in parallel. For each blob, it -// will `space/blob/add` if it hasn't been added yet, then call the `afterUploaded` callback if successful. -// `SpaceBlobAdded()` will be called after `space/blob/add`. `Added()` will be -// called at the very end. If any of these steps fail, an error will be -// returned. -func (a API) addBlobs(ctx context.Context, blobs []model.Blob, spaceDID did.DID, afterUploaded func(blob model.Blob) error) error { +// will `space/blob/add` if it hasn't been added yet, then call the +// `afterUploaded` callback if successful. `SpaceBlobAdded()` will be called +// after `space/blob/add`. `Added()` will be called at the very end. If any of +// these steps fail, an error will be returned. +func (a API) addBlobs(ctx context.Context, blobs <-chan model.Blob, spaceDID did.DID, afterUploaded func(blob model.Blob) error) error { // Ensure at least 1 parallelism if a.BlobUploadParallelism < 1 { a.BlobUploadParallelism = 1 } sem := make(chan struct{}, a.BlobUploadParallelism) - blobUploadErrorCh := make(chan gtypes.BlobUploadError, len(blobs)) eg, gctx := errgroup.WithContext(ctx) - for _, blob := range blobs { + + // Blob upload errors are non-fatal: collect them, but allow other uploads to + // proceed. + var ( + blobUploadErrorsMu sync.Mutex + blobUploadErrors []gtypes.BlobUploadError + ) + + for blob := range blobs { sem <- struct{}{} eg.Go(func() error { defer func() { <-sem }() @@ -140,7 +167,9 @@ func (a API) addBlobs(ctx context.Context, blobs []model.Blob, spaceDID did.DID, err = fmt.Errorf("failed to add blob %s: %w", blob, err) var errBlobUpload gtypes.BlobUploadError if errors.As(err, &errBlobUpload) { - blobUploadErrorCh <- errBlobUpload + blobUploadErrorsMu.Lock() + blobUploadErrors = append(blobUploadErrors, errBlobUpload) + blobUploadErrorsMu.Unlock() return nil } log.Errorf("%v", err) @@ -156,20 +185,19 @@ func (a API) addBlobs(ctx context.Context, blobs []model.Blob, spaceDID did.DID, }) } - terminalErr := eg.Wait() - close(blobUploadErrorCh) + fatalErr := eg.Wait() - if terminalErr != nil { - return terminalErr + // On fatal error, return that. + if fatalErr != nil { + return fatalErr } - var blobUploadErrors []gtypes.BlobUploadError - for err := range blobUploadErrorCh { - blobUploadErrors = append(blobUploadErrors, err) - } + // Otherwise, if we had any blob upload errors, return those as a batch error. if len(blobUploadErrors) > 0 { return gtypes.NewBlobUploadErrors(blobUploadErrors) } + + // And finally, the happy path. return nil } @@ -397,22 +425,40 @@ func (a API) AddIndexesForUpload(ctx context.Context, uploadID id.UploadID, spac ctx, span := tracer.Start(ctx, "add-indexes-for-upload") defer span.End() - closedIndexes, err := a.Repo.IndexesForUploadByState(ctx, uploadID, model.BlobStateClosed) - if err != nil { - return fmt.Errorf("failed to get closed indexes for upload %s: %w", uploadID, err) - } - span.AddEvent("found closed indexes", trace.WithAttributes(attribute.Int("indexes", len(closedIndexes)))) + blobs := make(chan model.Blob) - blobs := make([]model.Blob, len(closedIndexes)) - for i, shard := range closedIndexes { - blobs[i] = shard - } - return a.addBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { + errorCh := make(chan error, 1) + go func() { + defer close(blobs) + defer close(errorCh) + for { + closedIndexes, err := a.Repo.IndexesForUploadByState(ctx, uploadID, model.BlobStateClosed) + if err != nil { + errorCh <- fmt.Errorf("failed to get closed indexes for upload %s: %w", uploadID, err) + return + } + span.AddEvent("found closed indexes", trace.WithAttributes(attribute.Int("indexes", len(closedIndexes)))) + + if len(closedIndexes) == 0 { + return + } + + for _, index := range closedIndexes { + blobs <- index + } + } + }() + + addError := a.addBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { if indexCB != nil { return indexCB(blob.(*model.Index)) } return nil }) + + findError := <-errorCh + + return errors.Join(findError, addError) } // PostProcessUploadedIndexes runs post-processing for uploaded indexes, including From 97b9a7c1dad82196e7aece87e280b31c6cfd34fc Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Fri, 27 Mar 2026 16:35:29 -0400 Subject: [PATCH 2/7] feat: Blob uploads can use parallelism more effectively --- pkg/preparation/preparation.go | 20 +-- pkg/preparation/storacha/storacha.go | 208 +++++++--------------- pkg/preparation/storacha/storacha_test.go | 107 +++++++---- pkg/preparation/types/errors.go | 20 +-- pkg/preparation/uploads/uploads.go | 56 +++--- pkg/preparation/uploads/worker.go | 61 +++++++ 6 files changed, 245 insertions(+), 227 deletions(-) diff --git a/pkg/preparation/preparation.go b/pkg/preparation/preparation.go index 3341e37c..2ade6b03 100644 --- a/pkg/preparation/preparation.go +++ b/pkg/preparation/preparation.go @@ -164,14 +164,13 @@ func NewAPI(repo Repo, client StorachaClient, options ...Option) API { blobAddOptions = append(blobAddOptions, clientpkg.WithPutClient(cfg.putHTTPClient)) } storachaAPI := storacha.API{ - Repo: repo, - Client: client, - ReaderForShard: blobsAPI.ReaderForShard, - ReaderForIndex: blobsAPI.ReaderForIndex, - BlobUploadParallelism: cfg.blobUploadParallelism, - Bus: cfg.bus, - Replicas: cfg.replicas, - BlobAddOptions: blobAddOptions, + Repo: repo, + Client: client, + ReaderForShard: blobsAPI.ReaderForShard, + ReaderForIndex: blobsAPI.ReaderForIndex, + Bus: cfg.bus, + Replicas: cfg.replicas, + BlobAddOptions: blobAddOptions, } uploadsAPI = uploads.API{ @@ -183,10 +182,11 @@ func NewAPI(repo Repo, client StorachaClient, options ...Option) API { AddShardsToUploadIndexes: blobsAPI.AddShardsToUploadIndexes, CloseUploadShards: blobsAPI.CloseUploadShards, CloseUploadIndexes: blobsAPI.CloseUploadIndexes, - AddShardsForUpload: storachaAPI.AddShardsForUpload, + FindShardAddTasksForUpload: storachaAPI.FindShardAddTasksForUpload, + FindIndexAddTasksForUpload: storachaAPI.FindIndexAddTasksForUpload, + BlobUploadParallelism: cfg.blobUploadParallelism, PostProcessUploadedShards: storachaAPI.PostProcessUploadedShards, PostProcessUploadedIndexes: storachaAPI.PostProcessUploadedIndexes, - AddIndexesForUpload: storachaAPI.AddIndexesForUpload, AddStorachaUploadForUpload: storachaAPI.AddStorachaUploadForUpload, RemoveBadFSEntry: scansAPI.RemoveBadFSEntry, RemoveBadNodes: dagsAPI.RemoveBadNodes, diff --git a/pkg/preparation/storacha/storacha.go b/pkg/preparation/storacha/storacha.go index b0d48f48..c5db198a 100644 --- a/pkg/preparation/storacha/storacha.go +++ b/pkg/preparation/storacha/storacha.go @@ -5,7 +5,6 @@ import ( "errors" "fmt" "io" - "sync" "time" commp "github.com/filecoin-project/go-fil-commp-hashhash" @@ -59,58 +58,78 @@ type ReaderForIndexFunc func(ctx context.Context, indexID id.IndexID) (io.ReadCl // API provides methods to interact with Storacha. type API struct { - Repo Repo - Client Client - ReaderForShard ReaderForShardFunc - ReaderForIndex ReaderForIndexFunc + Repo Repo + Client Client + ReaderForShard ReaderForShardFunc + ReaderForIndex ReaderForIndexFunc + + // TK: Rm BlobUploadParallelism int Bus bus.Publisher Replicas uint BlobAddOptions []client.SpaceBlobAddOption } -var _ uploads.AddShardsForUploadFunc = API{}.AddShardsForUpload -var _ uploads.AddIndexesForUploadFunc = API{}.AddIndexesForUpload +var _ uploads.FindShardAddTasksForUploadFunc = API{}.FindShardAddTasksForUpload +var _ uploads.FindIndexAddTasksForUploadFunc = API{}.FindIndexAddTasksForUpload var _ uploads.AddStorachaUploadForUploadFunc = API{}.AddStorachaUploadForUpload -func (a API) AddShardsForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, shardUploadedCb func(shard *model.Shard) error) error { - ctx, span := tracer.Start(ctx, "add-shards-for-upload") +func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.BlobAddTask, error) { + ctx, span := tracer.Start(ctx, "find-shard-add-tasks-for-upload") defer span.End() - blobs := make(chan model.Blob) - - errorCh := make(chan error, 1) - go func() { - defer close(blobs) - defer close(errorCh) - for { - closedShards, err := a.Repo.ShardsForUploadByState(ctx, uploadID, model.BlobStateClosed) - if err != nil { - errorCh <- fmt.Errorf("failed to get closed shards for upload %s: %w", uploadID, err) - return + closedShards, err := a.Repo.ShardsForUploadByState(ctx, uploadID, model.BlobStateClosed) + if err != nil { + return nil, fmt.Errorf("failed to get closed shards for upload %s: %w", uploadID, err) + } + span.AddEvent("found closed shards", trace.WithAttributes(attribute.Int("shards", len(closedShards)))) + + tasks := make([]gtypes.BlobAddTask, 0, len(closedShards)) + for _, shard := range closedShards { + tasks = append(tasks, func(ctx context.Context) (error, error) { + if err := a.addBlob(ctx, shard, spaceDID); err != nil { + err = fmt.Errorf("failed to add shard %s: %w", shard, err) + // [gtypes.BlobUploadError]s are non-fatal. + var errBlobUpload gtypes.BlobUploadError + if errors.As(err, &errBlobUpload) { + return err, nil + } + log.Errorf("%v", err) + return nil, err } - span.AddEvent("found closed shards", trace.WithAttributes(attribute.Int("shards", len(closedShards)))) + return nil, nil + }) + } + return tasks, nil +} - if len(closedShards) == 0 { - return - } +func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.BlobAddTask, error) { + ctx, span := tracer.Start(ctx, "find-index-add-tasks-for-upload") + defer span.End() - for _, shard := range closedShards { - blobs <- shard + closedIndexes, err := a.Repo.IndexesForUploadByState(ctx, uploadID, model.BlobStateClosed) + if err != nil { + return nil, fmt.Errorf("failed to get closed indexes for upload %s: %w", uploadID, err) + } + span.AddEvent("found closed indexes", trace.WithAttributes(attribute.Int("indexes", len(closedIndexes)))) + + tasks := make([]gtypes.BlobAddTask, 0, len(closedIndexes)) + for _, index := range closedIndexes { + tasks = append(tasks, func(ctx context.Context) (error, error) { + if err := a.addBlob(ctx, index, spaceDID); err != nil { + err = fmt.Errorf("failed to add index %s: %w", index, err) + // [gtypes.BlobUploadError]s are non-fatal. + var errBlobUpload gtypes.BlobUploadError + if errors.As(err, &errBlobUpload) { + return err, nil + } + log.Errorf("%v", err) + return nil, err } - } - }() - - addError := a.addBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { - if shardUploadedCb != nil { - return shardUploadedCb(blob.(*model.Shard)) - } - return nil - }) - - findError := <-errorCh - - return errors.Join(findError, addError) + return nil, nil + }) + } + return tasks, nil } func (a API) PostProcessUploadedShards(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error { @@ -138,69 +157,6 @@ func (a API) PostProcessUploadedShards(ctx context.Context, uploadID id.UploadID }) } -// addBlobs adds the given blobs to the space, in parallel. For each blob, it -// will `space/blob/add` if it hasn't been added yet, then call the -// `afterUploaded` callback if successful. `SpaceBlobAdded()` will be called -// after `space/blob/add`. `Added()` will be called at the very end. If any of -// these steps fail, an error will be returned. -func (a API) addBlobs(ctx context.Context, blobs <-chan model.Blob, spaceDID did.DID, afterUploaded func(blob model.Blob) error) error { - // Ensure at least 1 parallelism - if a.BlobUploadParallelism < 1 { - a.BlobUploadParallelism = 1 - } - - sem := make(chan struct{}, a.BlobUploadParallelism) - eg, gctx := errgroup.WithContext(ctx) - - // Blob upload errors are non-fatal: collect them, but allow other uploads to - // proceed. - var ( - blobUploadErrorsMu sync.Mutex - blobUploadErrors []gtypes.BlobUploadError - ) - - for blob := range blobs { - sem <- struct{}{} - eg.Go(func() error { - defer func() { <-sem }() - if err := a.addBlob(gctx, blob, spaceDID); err != nil { - err = fmt.Errorf("failed to add blob %s: %w", blob, err) - var errBlobUpload gtypes.BlobUploadError - if errors.As(err, &errBlobUpload) { - blobUploadErrorsMu.Lock() - blobUploadErrors = append(blobUploadErrors, errBlobUpload) - blobUploadErrorsMu.Unlock() - return nil - } - log.Errorf("%v", err) - return err - } - if afterUploaded != nil { - if err := afterUploaded(blob); err != nil { - return fmt.Errorf("failed to call after uploaded callback for blob %s: %w", blob.ID(), err) - } - } - log.Infof("Successfully added blob %s", blob.ID()) - return nil - }) - } - - fatalErr := eg.Wait() - - // On fatal error, return that. - if fatalErr != nil { - return fatalErr - } - - // Otherwise, if we had any blob upload errors, return those as a batch error. - if len(blobUploadErrors) > 0 { - return gtypes.NewBlobUploadErrors(blobUploadErrors) - } - - // And finally, the happy path. - return nil -} - func (a API) postProcessBlobs(ctx context.Context, blobs []model.Blob, spaceDID did.DID, afterAdded func(blob model.Blob) error) error { // Ensure at least 1 parallelism if a.BlobUploadParallelism < 1 { @@ -208,7 +164,7 @@ func (a API) postProcessBlobs(ctx context.Context, blobs []model.Blob, spaceDID } sem := make(chan struct{}, a.BlobUploadParallelism) - blobUploadErrorCh := make(chan gtypes.BlobUploadError, len(blobs)) + blobUploadErrorCh := make(chan error, len(blobs)) eg, gctx := errgroup.WithContext(ctx) for _, blob := range blobs { sem <- struct{}{} @@ -218,7 +174,7 @@ func (a API) postProcessBlobs(ctx context.Context, blobs []model.Blob, spaceDID err = fmt.Errorf("failed to add blob %s: %w", blob, err) var errBlobUpload gtypes.BlobUploadError if errors.As(err, &errBlobUpload) { - blobUploadErrorCh <- errBlobUpload + blobUploadErrorCh <- err return nil } log.Errorf("%v", err) @@ -236,7 +192,7 @@ func (a API) postProcessBlobs(ctx context.Context, blobs []model.Blob, spaceDID return terminalErr } - var blobUploadErrors []gtypes.BlobUploadError + var blobUploadErrors []error for err := range blobUploadErrorCh { blobUploadErrors = append(blobUploadErrors, err) } @@ -419,48 +375,6 @@ func (a API) filecoinOffer(ctx context.Context, blob model.Blob, spaceDID did.DI return nil } -// AddIndexesForUpload adds the given indexes to the space, in parallel. The -// upload must have a root CID set. -func (a API) AddIndexesForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, indexCB func(index *model.Index) error) error { - ctx, span := tracer.Start(ctx, "add-indexes-for-upload") - defer span.End() - - blobs := make(chan model.Blob) - - errorCh := make(chan error, 1) - go func() { - defer close(blobs) - defer close(errorCh) - for { - closedIndexes, err := a.Repo.IndexesForUploadByState(ctx, uploadID, model.BlobStateClosed) - if err != nil { - errorCh <- fmt.Errorf("failed to get closed indexes for upload %s: %w", uploadID, err) - return - } - span.AddEvent("found closed indexes", trace.WithAttributes(attribute.Int("indexes", len(closedIndexes)))) - - if len(closedIndexes) == 0 { - return - } - - for _, index := range closedIndexes { - blobs <- index - } - } - }() - - addError := a.addBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { - if indexCB != nil { - return indexCB(blob.(*model.Index)) - } - return nil - }) - - findError := <-errorCh - - return errors.Join(findError, addError) -} - // PostProcessUploadedIndexes runs post-processing for uploaded indexes, including // adding them to the space via `space/index/add`. func (a API) PostProcessUploadedIndexes(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error { diff --git a/pkg/preparation/storacha/storacha_test.go b/pkg/preparation/storacha/storacha_test.go index ee5bb431..aaca876d 100644 --- a/pkg/preparation/storacha/storacha_test.go +++ b/pkg/preparation/storacha/storacha_test.go @@ -35,7 +35,7 @@ import ( // padding to every "CAR" to make sure it's definitely long enough. var padding = bytes.Repeat([]byte{0}, 127) -func TestAddShardsForUpload(t *testing.T) { +func TestFindShardAddTasksForUpload(t *testing.T) { t.Run("`space/blob/add`s, `space/blob/replicate`s, and `filecoin/offer`s a CAR for each shard", func(t *testing.T) { db := testdb.CreateTestDB(t) repo := stestutil.Must(sqlrepo.New(db))(t) @@ -53,11 +53,10 @@ func TestAddShardsForUpload(t *testing.T) { } api := storacha.API{ - Repo: repo, - Client: &client, - ReaderForShard: carForShard, - BlobUploadParallelism: 1, - Replicas: 3, + Repo: repo, + Client: &client, + ReaderForShard: carForShard, + Replicas: 3, } blobsApi := blobs.API{ @@ -83,8 +82,13 @@ func TestAddShardsForUpload(t *testing.T) { secondShard := shards[0] // Upload shards that are ready to go. - err = api.AddShardsForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload shards firstShard, err = repo.GetShardByID(t.Context(), firstShard.ID()) @@ -131,8 +135,13 @@ func TestAddShardsForUpload(t *testing.T) { // Now close the upload shards and run it again. err = blobsApi.CloseUploadShards(t.Context(), upload.ID(), nil) require.NoError(t, err) - err = api.AddShardsForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload second shard secondShard, err = repo.GetShardByID(t.Context(), secondShard.ID()) @@ -182,11 +191,10 @@ func TestAddShardsForUpload(t *testing.T) { } api := storacha.API{ - Repo: repo, - Client: &client, - ReaderForShard: carForShard, - BlobUploadParallelism: 1, - Replicas: 3, + Repo: repo, + Client: &client, + ReaderForShard: carForShard, + Replicas: 3, } blobsApi := blobs.API{ @@ -203,8 +211,13 @@ func TestAddShardsForUpload(t *testing.T) { client.SpaceBlobAddError = fmt.Errorf("simulated SpaceBlobAdd error") - err = api.AddShardsForUpload(t.Context(), upload.ID(), spaceDID, nil) - require.ErrorContains(t, err, "simulated SpaceBlobAdd error") + tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) + require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, err) + require.ErrorContains(t, nonFatal, "simulated SpaceBlobAdd error") + } err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) @@ -216,7 +229,7 @@ func TestAddShardsForUpload(t *testing.T) { require.Len(t, client.FilecoinOfferInvocations, 0) // It should have closed the first shard's reader. - require.Len(t, shardReadersClosed, 1, "expected shard readerto be closed, even though it failed") + require.Len(t, shardReadersClosed, 1, "expected shard reader to be closed, even though it failed") // reset the shard readers closed map for shardID := range shardReadersClosed { delete(shardReadersClosed, shardID) @@ -225,8 +238,13 @@ func TestAddShardsForUpload(t *testing.T) { // Now retry: `space/blob/add` succeeds but `space/blob/replicate` fails. client.SpaceBlobAddError = nil client.SpaceBlobReplicateError = fmt.Errorf("simulated SpaceBlobReplicate error") - err = api.AddShardsForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) require.ErrorContains(t, err, "simulated SpaceBlobReplicate error") @@ -243,8 +261,13 @@ func TestAddShardsForUpload(t *testing.T) { // Now retry: `space/blob/replicate` succeeds but `filecoin/offer` fails. client.SpaceBlobReplicateError = nil client.FilecoinOfferError = fmt.Errorf("simulated FilecoinOffer error") - err = api.AddShardsForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) require.ErrorContains(t, err, "simulated FilecoinOffer error") @@ -274,11 +297,10 @@ func TestAddShardsForUpload(t *testing.T) { } api := storacha.API{ - Repo: repo, - Client: &client, - ReaderForShard: carForShard, - BlobUploadParallelism: 1, - Replicas: 3, + Repo: repo, + Client: &client, + ReaderForShard: carForShard, + Replicas: 3, } blobsApi := blobs.API{ @@ -293,8 +315,13 @@ func TestAddShardsForUpload(t *testing.T) { err = blobsApi.CloseUploadShards(t.Context(), upload.ID(), nil) require.NoError(t, err) - err = api.AddShardsForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) @@ -309,7 +336,7 @@ func TestAddShardsForUpload(t *testing.T) { }) } -func TestAddIndexesForUpload(t *testing.T) { +func TestFindIndexAddTasksForUpload(t *testing.T) { t.Run("`space/blob/add`s and `space/blob/replicate`s index CARs", func(t *testing.T) { logging.SetLogLevel("preparation/storacha", "warn") db := testdb.CreateTestDB(t) @@ -327,11 +354,10 @@ func TestAddIndexesForUpload(t *testing.T) { } api := storacha.API{ - Repo: repo, - Client: &client, - ReaderForIndex: carForIndex, - BlobUploadParallelism: 1, - Replicas: 3, + Repo: repo, + Client: &client, + ReaderForIndex: carForIndex, + Replicas: 3, } blobsApi := blobs.API{ @@ -368,8 +394,13 @@ func TestAddIndexesForUpload(t *testing.T) { require.Len(t, shards, 3) require.Len(t, indexes, 1) - err = api.AddIndexesForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err := api.FindIndexAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload first shard firstIndex, err := repo.GetIndexByID(t.Context(), indexes[0].ID()) @@ -416,8 +447,13 @@ func TestAddIndexesForUpload(t *testing.T) { err = blobsApi.CloseUploadIndexes(t.Context(), upload.ID(), recordClosedIndex) require.NoError(t, err) require.Len(t, indexes, 2) - err = api.AddIndexesForUpload(t.Context(), upload.ID(), spaceDID, nil) + tasks, err = api.FindIndexAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range tasks { + nonFatal, err := task(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload second shard secondIndex, err := repo.GetIndexByID(t.Context(), indexes[1].ID()) @@ -460,10 +496,9 @@ func TestAddStorachaUploadForUpload(t *testing.T) { mclient := mockclient.MockClient{} api := storacha.API{ - Repo: repo, - Client: &mclient, - BlobUploadParallelism: 1, - Replicas: 3, + Repo: repo, + Client: &mclient, + Replicas: 3, } upload, _ := testutil.CreateUpload(t, repo, spaceDID, spacesmodel.WithShardSize(1<<16)) diff --git a/pkg/preparation/types/errors.go b/pkg/preparation/types/errors.go index 1716d2dd..53f7e511 100644 --- a/pkg/preparation/types/errors.go +++ b/pkg/preparation/types/errors.go @@ -1,6 +1,7 @@ package types import ( + "context" "errors" "fmt" "strings" @@ -169,17 +170,16 @@ func (e BlobUploadError) ID() id.ID { } type BlobUploadErrors struct { - errs []BlobUploadError + errs []error } -func NewBlobUploadErrors(errs []BlobUploadError) error { +func NewBlobUploadErrors(errs []error) error { + if len(errs) == 0 { + return nil + } return RetriableError{err: BlobUploadErrors{errs: errs}} } -func (e BlobUploadErrors) Errs() []BlobUploadError { - return e.errs -} - func (e BlobUploadErrors) Error() string { var messages []string for _, err := range e.errs { @@ -190,9 +190,7 @@ func (e BlobUploadErrors) Error() string { } func (e BlobUploadErrors) Unwrap() []error { - errs := make([]error, len(e.errs)) - for i, err := range e.errs { - errs[i] = err - } - return errs + return e.errs } + +type BlobAddTask func(ctx context.Context) (error, error) diff --git a/pkg/preparation/uploads/uploads.go b/pkg/preparation/uploads/uploads.go index 619e9f59..a61caa9f 100644 --- a/pkg/preparation/uploads/uploads.go +++ b/pkg/preparation/uploads/uploads.go @@ -34,9 +34,9 @@ type AddNodeToUploadShardsFunc func(ctx context.Context, uploadID id.UploadID, s type AddShardsToUploadIndexesFunc func(ctx context.Context, uploadID id.UploadID, indexCB func(index *blobsmodel.Index) error) error type CloseUploadShardsFunc func(ctx context.Context, uploadID id.UploadID, shardCB func(shard *blobsmodel.Shard) error) error type CloseUploadIndexesFunc func(ctx context.Context, uploadID id.UploadID, indexCB func(index *blobsmodel.Index) error) error -type AddShardsForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, shardCB func(shard *blobsmodel.Shard) error) error +type FindShardAddTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.BlobAddTask, error) +type FindIndexAddTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.BlobAddTask, error) type AddNodesToUploadShardsFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, shardCB func(shard *blobsmodel.Shard) error) error -type AddIndexesForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, indexCB func(index *blobsmodel.Index) error) error type PostProcessUploadedShardsFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error type PostProcessUploadedIndexesFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error type AddStorachaUploadForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error @@ -48,10 +48,11 @@ type API struct { Repo Repo ExecuteScan ExecuteScanFunc ExecuteDagScansForUpload ExecuteDagScansForUploadFunc - AddShardsForUpload AddShardsForUploadFunc + BlobUploadParallelism int + FindShardAddTasksForUpload FindShardAddTasksForUploadFunc + FindIndexAddTasksForUpload FindIndexAddTasksForUploadFunc PostProcessUploadedShards PostProcessUploadedShardsFunc PostProcessUploadedIndexes PostProcessUploadedIndexesFunc - AddIndexesForUpload AddIndexesForUploadFunc AddStorachaUploadForUpload AddStorachaUploadForUploadFunc RemoveBadFSEntry RemoveBadFSEntryFunc RemoveBadNodes RemoveBadNodesFunc @@ -288,10 +289,10 @@ func (a API) handleBadFSEntries(ctx context.Context, uploadID id.UploadID, badFS func (a API) handleBadBlobUploads(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, blobUploadErrors types.BlobUploadErrors) error { // when there's a bad shard upload, it's not based on a problem locally usually, unless bad nodes were read during upload - for _, e := range blobUploadErrors.Errs() { + for _, e := range blobUploadErrors.Unwrap() { // bad nodes error can happen from reading car during upload var badNodesErr types.BadNodesError - if errors.As(e.Unwrap(), &badNodesErr) { + if errors.As(e, &badNodesErr) { err := a.handleBadNodes(ctx, uploadID, spaceDID, badNodesErr) if err != nil { return err @@ -649,21 +650,23 @@ func runShardUploadWorker( span.End() }() - return Worker( + nonFatals, err := Worker2( ctx, closedShardsAvailable, + api.BlobUploadParallelism, + + // findWork + func(ctx context.Context) ([]types.BlobAddTask, error) { + return api.FindShardAddTasksForUpload(ctx, uploadID, spaceDID) + }, // doWork - func() error { - err := api.AddShardsForUpload(ctx, uploadID, spaceDID, func(shard *blobsmodel.Shard) error { + func(ctx context.Context, task types.BlobAddTask) (error, error) { + nonFatal, err := task(ctx) + if nonFatal == nil && err == nil { signal(uploadedShardsAvailable) - return nil - }) - if err != nil { - return fmt.Errorf("`space/blob/add`ing shards for upload %s: %w", uploadID, err) } - - return nil + return nonFatal, err }, // finalize @@ -672,6 +675,8 @@ func runShardUploadWorker( return nil }, ) + + return errors.Join(err, types.NewBlobUploadErrors(nonFatals)) } func runPostProcessShardWorker( @@ -761,20 +766,23 @@ func runIndexUploadWorker( span.End() }() - return Worker( + nonFatals, err := Worker2( ctx, closedIndexesAvailable, + api.BlobUploadParallelism, + + // findWork + func(ctx context.Context) ([]types.BlobAddTask, error) { + return api.FindIndexAddTasksForUpload(ctx, uploadID, spaceDID) + }, // doWork - func() error { - err := api.AddIndexesForUpload(ctx, uploadID, spaceDID, func(index *blobsmodel.Index) error { + func(ctx context.Context, task types.BlobAddTask) (error, error) { + nonFatal, err := task(ctx) + if nonFatal == nil && err == nil { signal(uploadedIndexesAvailable) - return nil - }) - if err != nil { - return fmt.Errorf("`space/blob/add`ing indexes for upload %s: %w", uploadID, err) } - return nil + return nonFatal, err }, // finalize @@ -783,6 +791,8 @@ func runIndexUploadWorker( return nil }, ) + + return errors.Join(err, types.NewBlobUploadErrors(nonFatals)) } func runPostProcessIndexWorker( diff --git a/pkg/preparation/uploads/worker.go b/pkg/preparation/uploads/worker.go index fe1147d4..6213a937 100644 --- a/pkg/preparation/uploads/worker.go +++ b/pkg/preparation/uploads/worker.go @@ -3,8 +3,10 @@ package uploads import ( "context" "fmt" + "sync" "github.com/storacha/guppy/internal/ctxutil" + "golang.org/x/sync/errgroup" ) func Worker(ctx context.Context, in <-chan struct{}, doWork func() error, finalize func() error) error { @@ -27,3 +29,62 @@ func Worker(ctx context.Context, in <-chan struct{}, doWork func() error, finali } } } + +func Worker2[Task any]( + ctx context.Context, + tasksAvailable <-chan struct{}, + parallelism int, + findWork func(ctx context.Context) ([]Task, error), + doWork func(ctx context.Context, task Task) (error, error), + finalize func() error, +) ([]error, error) { + for { + select { + case <-ctx.Done(): + return nil, ctxutil.Cause(ctx) + case _, ok := <-tasksAvailable: + if !ok { + if finalize != nil { + if err := finalize(); err != nil { + return nil, fmt.Errorf("worker finalize encountered an error: %w", err) + } + } + return nil, nil + } + + tasks, err := findWork(ctx) + if err != nil { + return nil, fmt.Errorf("worker findWork encountered an error: %w", err) + } + + sem := make(chan struct{}, parallelism) + eg, gctx := errgroup.WithContext(ctx) + var ( + nonFatalErrorsMu sync.Mutex + nonFatalErrors []error + ) + for _, task := range tasks { + sem <- struct{}{} + task := task + eg.Go(func() error { + defer func() { <-sem }() + nonFatal, fatal := doWork(gctx, task) + if fatal != nil { + return fmt.Errorf("worker doWork encountered a fatal error: %w", fatal) + } + if nonFatal != nil { + nonFatalErrorsMu.Lock() + nonFatalErrors = append(nonFatalErrors, nonFatal) + nonFatalErrorsMu.Unlock() + } + return nil + }) + } + + fatalErr := eg.Wait() + if fatalErr != nil || len(nonFatalErrors) > 0 { + return nonFatalErrors, fatalErr + } + } + } +} From 4544688f324b2a5b112fb2f951e71a0764b74217 Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Mon, 30 Mar 2026 16:33:18 -0400 Subject: [PATCH 3/7] feat: Blob uploads can use parallelism even more effectively --- pkg/preparation/preparation_test.go | 2 +- pkg/preparation/storacha/storacha.go | 54 +++-- pkg/preparation/storacha/storacha_test.go | 16 +- pkg/preparation/types/errors.go | 5 +- pkg/preparation/uploads/uploads.go | 255 ++++++++++++++-------- pkg/preparation/uploads/worker.go | 123 ++++++----- pkg/preparation/uploads/worker_test.go | 141 +++++++----- 7 files changed, 355 insertions(+), 241 deletions(-) diff --git a/pkg/preparation/preparation_test.go b/pkg/preparation/preparation_test.go index 100c8a2c..2e1d1be0 100644 --- a/pkg/preparation/preparation_test.go +++ b/pkg/preparation/preparation_test.go @@ -405,7 +405,7 @@ func TestExecuteUpload(t *testing.T) { // We don't know exactly how many successful PUTs there were, but we know it // should be at least 2 and at most 6. require.GreaterOrEqual(t, putBlobs.Size(), 2, "expected at least 2/5 shards to be added so far") - require.Less(t, putBlobs.Size(), 6, "expected at most 4/5 shards + 1 index to be added so far") + require.LessOrEqual(t, putBlobs.Size(), 6, "expected at most 5/5 shards + 1 index to be added so far") require.Len(t, uploadAddCaps, 0, "expected `upload/add` not to have been called yet") t.Log("Retrying the upload after error...") diff --git a/pkg/preparation/storacha/storacha.go b/pkg/preparation/storacha/storacha.go index c5db198a..7c5a7153 100644 --- a/pkg/preparation/storacha/storacha.go +++ b/pkg/preparation/storacha/storacha.go @@ -86,18 +86,21 @@ func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadI tasks := make([]gtypes.BlobAddTask, 0, len(closedShards)) for _, shard := range closedShards { - tasks = append(tasks, func(ctx context.Context) (error, error) { - if err := a.addBlob(ctx, shard, spaceDID); err != nil { - err = fmt.Errorf("failed to add shard %s: %w", shard, err) - // [gtypes.BlobUploadError]s are non-fatal. - var errBlobUpload gtypes.BlobUploadError - if errors.As(err, &errBlobUpload) { - return err, nil + tasks = append(tasks, gtypes.BlobAddTask{ + ID: shard.ID(), + Run: func(ctx context.Context) (error, error) { + if err := a.addBlob(ctx, shard, spaceDID); err != nil { + err = fmt.Errorf("failed to add shard %s: %w", shard, err) + // [gtypes.BlobUploadError]s are non-fatal. + var errBlobUpload gtypes.BlobUploadError + if errors.As(err, &errBlobUpload) { + return err, nil + } + log.Errorf("%v", err) + return nil, err } - log.Errorf("%v", err) - return nil, err - } - return nil, nil + return nil, nil + }, }) } return tasks, nil @@ -115,18 +118,21 @@ func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadI tasks := make([]gtypes.BlobAddTask, 0, len(closedIndexes)) for _, index := range closedIndexes { - tasks = append(tasks, func(ctx context.Context) (error, error) { - if err := a.addBlob(ctx, index, spaceDID); err != nil { - err = fmt.Errorf("failed to add index %s: %w", index, err) - // [gtypes.BlobUploadError]s are non-fatal. - var errBlobUpload gtypes.BlobUploadError - if errors.As(err, &errBlobUpload) { - return err, nil + tasks = append(tasks, gtypes.BlobAddTask{ + ID: index.ID(), + Run: func(ctx context.Context) (error, error) { + if err := a.addBlob(ctx, index, spaceDID); err != nil { + err = fmt.Errorf("failed to add index %s: %w", index, err) + // [gtypes.BlobUploadError]s are non-fatal. + var errBlobUpload gtypes.BlobUploadError + if errors.As(err, &errBlobUpload) { + return err, nil + } + log.Errorf("%v", err) + return nil, err } - log.Errorf("%v", err) - return nil, err - } - return nil, nil + return nil, nil + }, }) } return tasks, nil @@ -290,7 +296,7 @@ func (a API) addBlob(ctx context.Context, blob model.Blob, spaceDID did.DID) err } } - if err := a.updateBlob(ctx, blob); err != nil { + if err := a.updateBlob(context.WithoutCancel(ctx), blob); err != nil { return fmt.Errorf("failed to update blob %s after `space/blob/add`: %w", blob, err) } return nil @@ -310,7 +316,7 @@ func (a API) postProcessBlob(ctx context.Context, blob model.Blob, spaceDID did. if err := blob.Added(); err != nil { return fmt.Errorf("failed to mark blob %s as added: %w", blob, err) } - if err := a.updateBlob(ctx, blob); err != nil { + if err := a.updateBlob(context.WithoutCancel(ctx), blob); err != nil { return fmt.Errorf("failed to update blob %s after adding to space: %w", blob, err) } diff --git a/pkg/preparation/storacha/storacha_test.go b/pkg/preparation/storacha/storacha_test.go index aaca876d..72b63935 100644 --- a/pkg/preparation/storacha/storacha_test.go +++ b/pkg/preparation/storacha/storacha_test.go @@ -85,7 +85,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } @@ -138,7 +138,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } @@ -214,7 +214,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, err) require.ErrorContains(t, nonFatal, "simulated SpaceBlobAdd error") } @@ -241,7 +241,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } @@ -264,7 +264,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } @@ -318,7 +318,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } @@ -397,7 +397,7 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { tasks, err := api.FindIndexAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } @@ -450,7 +450,7 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { tasks, err = api.FindIndexAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task(t.Context()) + nonFatal, err := task.Run(t.Context()) require.NoError(t, nonFatal) require.NoError(t, err) } diff --git a/pkg/preparation/types/errors.go b/pkg/preparation/types/errors.go index 53f7e511..978f43e7 100644 --- a/pkg/preparation/types/errors.go +++ b/pkg/preparation/types/errors.go @@ -193,4 +193,7 @@ func (e BlobUploadErrors) Unwrap() []error { return e.errs } -type BlobAddTask func(ctx context.Context) (error, error) +type BlobAddTask struct { + ID id.ID + Run func(context.Context) (error, error) +} diff --git a/pkg/preparation/uploads/uploads.go b/pkg/preparation/uploads/uploads.go index a61caa9f..b201a4d8 100644 --- a/pkg/preparation/uploads/uploads.go +++ b/pkg/preparation/uploads/uploads.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "sync" "github.com/ipfs/go-cid" logging "github.com/ipfs/go-log/v2" @@ -372,39 +373,44 @@ func runScanWorker( span.End() }() - return Worker( + _, err = Worker( ctx, scansAvailable, + 1, - // doWork - func() error { - if api.AssumeUnchangedSources { - upload, err := api.Repo.GetUploadByID(ctx, uploadID) - if err != nil { - return fmt.Errorf("checking upload for existing scan: %w", err) - } - if upload.HasRootFSEntryID() { - log.Infow("Skipping FS rescan (--assume-unchanged-sources): scan already exists", "upload", uploadID) - return nil - } - log.Infow("No existing scan found, performing FS scan despite --assume-unchanged-sources", "upload", uploadID) - } - - err := api.ExecuteScan(ctx, uploadID, func(entry scanmodel.FSEntry) error { - _, isDirectory := entry.(*scanmodel.Directory) - _, err := api.Repo.CreateDAGScan(ctx, entry.ID(), isDirectory, uploadID, spaceDID) - if err != nil { - return fmt.Errorf("creating DAG scan: %w", err) - } - signal(dagScansAvailable) - return nil - }) - - if err != nil { - return fmt.Errorf("running scans: %w", err) - } - - return nil + // findWork + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + return []func(context.Context) (error, error){ + func(ctx context.Context) (error, error) { + if api.AssumeUnchangedSources { + upload, err := api.Repo.GetUploadByID(ctx, uploadID) + if err != nil { + return nil, fmt.Errorf("checking upload for existing scan: %w", err) + } + if upload.HasRootFSEntryID() { + log.Infow("Skipping FS rescan (--assume-unchanged-sources): scan already exists", "upload", uploadID) + return nil, nil + } + log.Infow("No existing scan found, performing FS scan despite --assume-unchanged-sources", "upload", uploadID) + } + + err := api.ExecuteScan(ctx, uploadID, func(entry scanmodel.FSEntry) error { + _, isDirectory := entry.(*scanmodel.Directory) + _, err := api.Repo.CreateDAGScan(ctx, entry.ID(), isDirectory, uploadID, spaceDID) + if err != nil { + return fmt.Errorf("creating DAG scan: %w", err) + } + signal(dagScansAvailable) + return nil + }) + + if err != nil { + return nil, fmt.Errorf("running scans: %w", err) + } + + return nil, nil + }, + }, nil }, // finalize @@ -413,6 +419,7 @@ func runScanWorker( return nil }, ) + return err } // runDAGScanWorker runs the worker that scans files and directories into blocks, @@ -448,22 +455,27 @@ func runDAGScanWorker( span.End() }() - return Worker( + _, err = Worker( ctx, dagScansAvailable, + 1, - // doWork - func() error { - err := api.ExecuteDagScansForUpload(ctx, uploadID, func(node dagmodel.Node, data []byte) error { - signal(nodeUploadsAvailable) - return nil - }) - - if err != nil { - return fmt.Errorf("running dag scans for upload %s: %w", uploadID, err) - } - - return nil + // findWork + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + return []func(context.Context) (error, error){ + func(ctx context.Context) (error, error) { + err := api.ExecuteDagScansForUpload(ctx, uploadID, func(node dagmodel.Node, data []byte) error { + signal(nodeUploadsAvailable) + return nil + }) + + if err != nil { + return nil, fmt.Errorf("running dag scans for upload %s: %w", uploadID, err) + } + + return nil, nil + }, + }, nil }, // finalize @@ -489,6 +501,7 @@ func runDAGScanWorker( return nil }, ) + return err } // runShardingWorker runs the worker that assigns nodes to shards. @@ -528,17 +541,22 @@ func runShardingWorker( return nil } - return Worker( + _, err = Worker( ctx, nodeUploadsAvailable, + 1, - // doWork - func() error { - err := api.AddNodesToUploadShards(ctx, uploadID, spaceDID, handleClosedShard) - if err != nil { - return fmt.Errorf("adding nodes to shards for upload %s: %w", uploadID, err) - } - return nil + // findWork + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + return []func(context.Context) (error, error){ + func(ctx context.Context) (error, error) { + err := api.AddNodesToUploadShards(ctx, uploadID, spaceDID, handleClosedShard) + if err != nil { + return nil, fmt.Errorf("adding nodes to shards for upload %s: %w", uploadID, err) + } + return nil, nil + }, + }, nil }, // finalize @@ -553,6 +571,7 @@ func runShardingWorker( return nil }, ) + return err } func runIndexingWorker( @@ -591,17 +610,22 @@ func runIndexingWorker( return nil } - return Worker( + _, err = Worker( ctx, shardsNeedIndexing, + 1, - // doWork - func() error { - err := api.AddShardsToUploadIndexes(ctx, uploadID, handleClosedIndex) - if err != nil { - return fmt.Errorf("adding shards to indexes for upload %s: %w", uploadID, err) - } - return nil + // findWork + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + return []func(context.Context) (error, error){ + func(ctx context.Context) (error, error) { + err := api.AddShardsToUploadIndexes(ctx, uploadID, handleClosedIndex) + if err != nil { + return nil, fmt.Errorf("adding shards to indexes for upload %s: %w", uploadID, err) + } + return nil, nil + }, + }, nil }, // finalize @@ -616,6 +640,7 @@ func runIndexingWorker( return nil }, ) + return err } // runShardUploadWorker runs the worker that adds shards to Storacha. @@ -650,23 +675,35 @@ func runShardUploadWorker( span.End() }() - nonFatals, err := Worker2( + var inFlightShards sync.Map + + nonFatals, err := Worker( ctx, closedShardsAvailable, api.BlobUploadParallelism, // findWork - func(ctx context.Context) ([]types.BlobAddTask, error) { - return api.FindShardAddTasksForUpload(ctx, uploadID, spaceDID) - }, - - // doWork - func(ctx context.Context, task types.BlobAddTask) (error, error) { - nonFatal, err := task(ctx) - if nonFatal == nil && err == nil { - signal(uploadedShardsAvailable) + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + rawTasks, err := api.FindShardAddTasksForUpload(ctx, uploadID, spaceDID) + if err != nil { + return nil, err } - return nonFatal, err + var tasks []func(context.Context) (error, error) + for _, raw := range rawTasks { + // Ignore tasks that are already in flight. + if _, loaded := inFlightShards.LoadOrStore(raw.ID, struct{}{}); loaded { + continue + } + tasks = append(tasks, func(ctx context.Context) (error, error) { + defer inFlightShards.Delete(raw.ID) + nonFatal, fatal := raw.Run(ctx) + if nonFatal == nil && fatal == nil { + signal(uploadedShardsAvailable) + } + return nonFatal, fatal + }) + } + return tasks, nil }, // finalize @@ -709,17 +746,22 @@ func runPostProcessShardWorker( span.End() }() - return Worker( + _, err = Worker( ctx, uploadedShardsAvailable, + 1, - // doWork - func() error { - err := api.PostProcessUploadedShards(ctx, uploadID, spaceDID) - if err != nil { - return fmt.Errorf("`post-processing shards for upload %s: %w", uploadID, err) - } - return nil + // findWork + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + return []func(context.Context) (error, error){ + func(ctx context.Context) (error, error) { + err := api.PostProcessUploadedShards(ctx, uploadID, spaceDID) + if err != nil { + return nil, fmt.Errorf("`post-processing shards for upload %s: %w", uploadID, err) + } + return nil, nil + }, + }, nil }, // finalize @@ -732,6 +774,7 @@ func runPostProcessShardWorker( return nil }, ) + return err } // runIndexUploadWorker runs the worker that adds indexes to Storacha. @@ -766,23 +809,35 @@ func runIndexUploadWorker( span.End() }() - nonFatals, err := Worker2( + var inFlightIndexes sync.Map + + nonFatals, err := Worker( ctx, closedIndexesAvailable, api.BlobUploadParallelism, // findWork - func(ctx context.Context) ([]types.BlobAddTask, error) { - return api.FindIndexAddTasksForUpload(ctx, uploadID, spaceDID) - }, - - // doWork - func(ctx context.Context, task types.BlobAddTask) (error, error) { - nonFatal, err := task(ctx) - if nonFatal == nil && err == nil { - signal(uploadedIndexesAvailable) + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + rawTasks, err := api.FindIndexAddTasksForUpload(ctx, uploadID, spaceDID) + if err != nil { + return nil, err + } + var tasks []func(context.Context) (error, error) + for _, raw := range rawTasks { + // Ignore tasks that are already in flight. + if _, loaded := inFlightIndexes.LoadOrStore(raw.ID, struct{}{}); loaded { + continue + } + tasks = append(tasks, func(ctx context.Context) (error, error) { + defer inFlightIndexes.Delete(raw.ID) + nonFatal, fatal := raw.Run(ctx) + if nonFatal == nil && fatal == nil { + signal(uploadedIndexesAvailable) + } + return nonFatal, fatal + }) } - return nonFatal, err + return tasks, nil }, // finalize @@ -825,20 +880,26 @@ func runPostProcessIndexWorker( span.End() }() - return Worker( + _, err = Worker( ctx, uploadedIndexesAvailable, + 1, - // doWork - func() error { - err := api.PostProcessUploadedIndexes(ctx, uploadID, spaceDID) - if err != nil { - return fmt.Errorf("`post-processing indexes for upload %s: %w", uploadID, err) - } - return nil + // findWork + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + return []func(context.Context) (error, error){ + func(ctx context.Context) (error, error) { + err := api.PostProcessUploadedIndexes(ctx, uploadID, spaceDID) + if err != nil { + return nil, fmt.Errorf("`post-processing indexes for upload %s: %w", uploadID, err) + } + return nil, nil + }, + }, nil }, // finalize nil, ) + return err } diff --git a/pkg/preparation/uploads/worker.go b/pkg/preparation/uploads/worker.go index 6213a937..d94a6bd0 100644 --- a/pkg/preparation/uploads/worker.go +++ b/pkg/preparation/uploads/worker.go @@ -9,81 +9,92 @@ import ( "golang.org/x/sync/errgroup" ) -func Worker(ctx context.Context, in <-chan struct{}, doWork func() error, finalize func() error) error { - for { +func Worker( + ctx context.Context, + workAvailable <-chan struct{}, + parallelism int, + findWork func(ctx context.Context) ([]func(context.Context) (error, error), error), + finalize func() error, +) ([]error, error) { + var ( + queue []func(context.Context) (error, error) + nonFatalErrorsMu sync.Mutex + nonFatalErrors []error + ) + + // gctx is cancelled when any task returns a fatal error, allowing the outer + // loop to detect failure and stop dispatching. Tasks receive the outer ctx + // (not gctx) so sibling tasks are not cancelled when one fails. + sem := make(chan struct{}, parallelism) + eg, gctx := errgroup.WithContext(ctx) + + dispatchNext := func() bool { + if len(queue) == 0 { + return false + } + task := queue[0] + queue = queue[1:] select { - case <-ctx.Done(): - return ctxutil.Cause(ctx) - case _, ok := <-in: - if !ok { - if finalize != nil { - if err := finalize(); err != nil { - return fmt.Errorf("worker finalize encountered an error: %w", err) - } - } - return nil + case sem <- struct{}{}: + case <-gctx.Done(): + return false + } + eg.Go(func() error { + defer func() { <-sem }() + nonFatal, fatal := task(ctx) + if fatal != nil { + return fmt.Errorf("worker task encountered a fatal error: %w", fatal) } - if err := doWork(); err != nil { - return fmt.Errorf("worker encountered an error: %w", err) + if nonFatal != nil { + nonFatalErrorsMu.Lock() + nonFatalErrors = append(nonFatalErrors, nonFatal) + nonFatalErrorsMu.Unlock() } - } + return nil + }) + return true } -} -func Worker2[Task any]( - ctx context.Context, - tasksAvailable <-chan struct{}, - parallelism int, - findWork func(ctx context.Context) ([]Task, error), - doWork func(ctx context.Context, task Task) (error, error), - finalize func() error, -) ([]error, error) { for { select { - case <-ctx.Done(): - return nil, ctxutil.Cause(ctx) - case _, ok := <-tasksAvailable: + case <-gctx.Done(): + // gctx is cancelled either because ctx was cancelled (external stop) + // or because a task returned a fatal error (internal failure). Wait + // for all in-flight tasks to finish before determining which it was. + fatalErr := eg.Wait() + if fatalErr != nil { + return nonFatalErrors, fatalErr + } + // No task error — must be an external cancellation. + return nonFatalErrors, ctxutil.Cause(ctx) + case _, ok := <-workAvailable: if !ok { + // Drain the queue before finalizing. + for dispatchNext() { + } + if fatalErr := eg.Wait(); fatalErr != nil { + return nonFatalErrors, fatalErr + } if finalize != nil { if err := finalize(); err != nil { - return nil, fmt.Errorf("worker finalize encountered an error: %w", err) + return nonFatalErrors, fmt.Errorf("worker finalize encountered an error: %w", err) } } + if len(nonFatalErrors) > 0 { + return nonFatalErrors, nil + } return nil, nil } tasks, err := findWork(ctx) if err != nil { - return nil, fmt.Errorf("worker findWork encountered an error: %w", err) - } - - sem := make(chan struct{}, parallelism) - eg, gctx := errgroup.WithContext(ctx) - var ( - nonFatalErrorsMu sync.Mutex - nonFatalErrors []error - ) - for _, task := range tasks { - sem <- struct{}{} - task := task - eg.Go(func() error { - defer func() { <-sem }() - nonFatal, fatal := doWork(gctx, task) - if fatal != nil { - return fmt.Errorf("worker doWork encountered a fatal error: %w", fatal) - } - if nonFatal != nil { - nonFatalErrorsMu.Lock() - nonFatalErrors = append(nonFatalErrors, nonFatal) - nonFatalErrorsMu.Unlock() - } - return nil - }) + _ = eg.Wait() + return nonFatalErrors, fmt.Errorf("worker findWork encountered an error: %w", err) } + queue = append(queue, tasks...) - fatalErr := eg.Wait() - if fatalErr != nil || len(nonFatalErrors) > 0 { - return nonFatalErrors, fatalErr + // Fill available parallelism slots. + for dispatchNext() { } } } diff --git a/pkg/preparation/uploads/worker_test.go b/pkg/preparation/uploads/worker_test.go index ae4b1a50..84d8fcb1 100644 --- a/pkg/preparation/uploads/worker_test.go +++ b/pkg/preparation/uploads/worker_test.go @@ -1,6 +1,7 @@ package uploads_test import ( + "context" "errors" "testing" "time" @@ -10,11 +11,6 @@ import ( "github.com/stretchr/testify/require" ) -type unwrappableError interface { - error - Unwrap() error -} - // e is a shorthand helper function that uses [require.EventuallyWithT] with // a standard timeout and interval, to keep noise out of the tests. func e(t *testing.T, condition func(collect *assert.CollectT)) { @@ -22,23 +18,36 @@ func e(t *testing.T, condition func(collect *assert.CollectT)) { require.EventuallyWithT(t, condition, time.Second, 10*time.Millisecond) } +func task(fn func() error) func(context.Context) (error, error) { + return func(ctx context.Context) (error, error) { + return nil, fn() + } +} + func TestWorker(t *testing.T) { t.Run("runs the work function for every signal received, then the finalize function when the channel closes", func(t *testing.T) { signalChan := make(chan struct{}, 1) - resultChan := make(chan error, 1) + type result struct { + nonFatals []error + err error + } + resultChan := make(chan result, 1) var runs int var finalizes int go func() { defer close(resultChan) - - resultChan <- uploads.Worker(t.Context(), signalChan, func() error { - runs++ - return nil - }, func() error { - finalizes++ - return nil - }) + nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + runs++ + return nil, nil + }, + func() error { + finalizes++ + return nil + }, + ) + resultChan <- result{nonFatals, err} }() require.Equal(t, 0, runs, "worker should not run before signal") @@ -47,36 +56,46 @@ func TestWorker(t *testing.T) { signalChan <- struct{}{} e(t, func(t *assert.CollectT) { require.Equal(t, 2, runs, "worker should run again after second signal") }) - require.Equal(t, 0, finalizes, "finalize function should be called until the channel closes") + require.Equal(t, 0, finalizes, "finalize function should not be called until the channel closes") close(signalChan) e(t, func(t *assert.CollectT) { require.Equal(t, 1, finalizes, "finalize function should be called once the channel closes") }) - result := <-resultChan - require.Nil(t, result, "result should be nil after successful runs") + res := <-resultChan + require.Nil(t, res.err, "error should be nil after successful runs") + require.Empty(t, res.nonFatals, "non-fatal errors should be empty after successful runs") }) t.Run("immediately responds with any work error, skipping the finalizer", func(t *testing.T) { workerErr := errors.New("error in doWork") signalChan := make(chan struct{}, 3) - resultChan := make(chan error, 1) + type result struct { + nonFatals []error + err error + } + resultChan := make(chan result, 1) var runs int var finalizes int go func() { defer close(resultChan) - resultChan <- uploads.Worker(t.Context(), signalChan, func() error { - runs++ - // Fail on the second run - if runs == 2 { - return workerErr - } - return nil - }, func() error { - finalizes++ - return nil - }) + nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + runs++ + if runs == 2 { + return []func(context.Context) (error, error){ + task(func() error { return workerErr }), + }, nil + } + return nil, nil + }, + func() error { + finalizes++ + return nil + }, + ) + resultChan <- result{nonFatals, err} }() // Send three signals; the second should cause an error, the third should not run @@ -84,28 +103,35 @@ func TestWorker(t *testing.T) { signalChan <- struct{}{} signalChan <- struct{}{} - result, ok := (<-resultChan).(unwrappableError) - require.True(t, ok, "result should be a wrapped error") - require.ErrorContains(t, result, "worker encountered an error: error in doWork") - require.Equal(t, workerErr, result.Unwrap(), "worker should send back the error it encountered, wrapped") - require.Equal(t, 2, runs, "worker should have stopped after encountering an error") - require.Equal(t, 0, finalizes, "finalize function should not have be called") + res := <-resultChan + require.ErrorContains(t, res.err, "worker task encountered a fatal error: error in doWork") + require.ErrorIs(t, res.err, workerErr) + require.LessOrEqual(t, runs, 3, "worker should have stopped after encountering an error") + require.Equal(t, 0, finalizes, "finalize function should not have been called") }) t.Run("responds with any finalize error", func(t *testing.T) { finalizerErr := errors.New("error in finalize") signalChan := make(chan struct{}, 3) - resultChan := make(chan error, 1) + type result struct { + nonFatals []error + err error + } + resultChan := make(chan result, 1) var runs int go func() { defer close(resultChan) - resultChan <- uploads.Worker(t.Context(), signalChan, func() error { - runs++ - return nil - }, func() error { - return finalizerErr - }) + nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + runs++ + return nil, nil + }, + func() error { + return finalizerErr + }, + ) + resultChan <- result{nonFatals, err} }() // Send three signals; all should run @@ -114,32 +140,39 @@ func TestWorker(t *testing.T) { signalChan <- struct{}{} close(signalChan) - result, ok := (<-resultChan).(unwrappableError) - require.True(t, ok, "result should be a wrapped error") - require.ErrorContains(t, result, "worker finalize encountered an error: error in finalize") - require.Equal(t, finalizerErr, result.Unwrap(), "worker should send back the error it encountered, wrapped") + res := <-resultChan + require.ErrorContains(t, res.err, "worker finalize encountered an error: error in finalize") + require.ErrorIs(t, res.err, finalizerErr) require.Equal(t, 3, runs, "worker should have run all three times") }) t.Run("ignores a nil finalizer", func(t *testing.T) { signalChan := make(chan struct{}, 1) - resultChan := make(chan error, 1) + type result struct { + nonFatals []error + err error + } + resultChan := make(chan result, 1) var ran bool go func() { defer close(resultChan) - - resultChan <- uploads.Worker(t.Context(), signalChan, func() error { - ran = true - return nil - }, nil) + nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, + func(ctx context.Context) ([]func(context.Context) (error, error), error) { + ran = true + return nil, nil + }, + nil, + ) + resultChan <- result{nonFatals, err} }() require.False(t, ran, "worker should not run before signal") signalChan <- struct{}{} e(t, func(t *assert.CollectT) { require.True(t, ran, "worker should run after signal") }) close(signalChan) - result := <-resultChan - require.Nil(t, result, "result should be nil after successful runs and no finalizer") + res := <-resultChan + require.Nil(t, res.err, "error should be nil after successful runs and no finalizer") + require.Empty(t, res.nonFatals) }) } From 3082cdf76c4430e84b4173c6bf96ee1ca577a390 Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Tue, 31 Mar 2026 10:26:13 -0400 Subject: [PATCH 4/7] fix: Kick off `closedIndexesAvailable` too --- pkg/preparation/uploads/uploads.go | 1 + 1 file changed, 1 insertion(+) diff --git a/pkg/preparation/uploads/uploads.go b/pkg/preparation/uploads/uploads.go index b201a4d8..ccdf9343 100644 --- a/pkg/preparation/uploads/uploads.go +++ b/pkg/preparation/uploads/uploads.go @@ -222,6 +222,7 @@ func (a API) ExecuteUpload(ctx context.Context, uploadID id.UploadID, spaceDID d signal(dagScansAvailable) signal(nodeUploadsAvailable) signal(closedShardsAvailable) + signal(closedIndexesAvailable) signal(uploadedShardsAvailable) signal(uploadedIndexesAvailable) close(scansAvailable) From 95a51771d6ee7a10e7b9d2da7aa086a55a63d666 Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Tue, 31 Mar 2026 11:54:46 -0400 Subject: [PATCH 5/7] feat: Post-processing uses better parallelism, too --- pkg/preparation/preparation.go | 36 +++--- pkg/preparation/storacha/storacha.go | 142 +++++++++------------- pkg/preparation/storacha/storacha_test.go | 60 +++++++-- pkg/preparation/types/errors.go | 6 +- pkg/preparation/uploads/uploads.go | 89 ++++++++------ 5 files changed, 184 insertions(+), 149 deletions(-) diff --git a/pkg/preparation/preparation.go b/pkg/preparation/preparation.go index 2ade6b03..801f61b3 100644 --- a/pkg/preparation/preparation.go +++ b/pkg/preparation/preparation.go @@ -174,24 +174,24 @@ func NewAPI(repo Repo, client StorachaClient, options ...Option) API { } uploadsAPI = uploads.API{ - Repo: repo, - AssumeUnchangedSources: cfg.assumeUnchangedSources, - ExecuteScan: scansAPI.ExecuteScan, - ExecuteDagScansForUpload: dagsAPI.ExecuteDagScansForUpload, - AddNodesToUploadShards: blobsAPI.AddNodesToUploadShards, - AddShardsToUploadIndexes: blobsAPI.AddShardsToUploadIndexes, - CloseUploadShards: blobsAPI.CloseUploadShards, - CloseUploadIndexes: blobsAPI.CloseUploadIndexes, - FindShardAddTasksForUpload: storachaAPI.FindShardAddTasksForUpload, - FindIndexAddTasksForUpload: storachaAPI.FindIndexAddTasksForUpload, - BlobUploadParallelism: cfg.blobUploadParallelism, - PostProcessUploadedShards: storachaAPI.PostProcessUploadedShards, - PostProcessUploadedIndexes: storachaAPI.PostProcessUploadedIndexes, - AddStorachaUploadForUpload: storachaAPI.AddStorachaUploadForUpload, - RemoveBadFSEntry: scansAPI.RemoveBadFSEntry, - RemoveBadNodes: dagsAPI.RemoveBadNodes, - RemoveShard: blobsAPI.RemoveShard, - Publisher: cfg.bus, + Repo: repo, + AssumeUnchangedSources: cfg.assumeUnchangedSources, + ExecuteScan: scansAPI.ExecuteScan, + ExecuteDagScansForUpload: dagsAPI.ExecuteDagScansForUpload, + AddNodesToUploadShards: blobsAPI.AddNodesToUploadShards, + AddShardsToUploadIndexes: blobsAPI.AddShardsToUploadIndexes, + CloseUploadShards: blobsAPI.CloseUploadShards, + CloseUploadIndexes: blobsAPI.CloseUploadIndexes, + FindShardAddTasksForUpload: storachaAPI.FindShardAddTasksForUpload, + FindIndexAddTasksForUpload: storachaAPI.FindIndexAddTasksForUpload, + BlobUploadParallelism: cfg.blobUploadParallelism, + FindShardPostProcessTasksForUpload: storachaAPI.FindShardPostProcessTasksForUpload, + FindIndexPostProcessTasksForUpload: storachaAPI.FindIndexPostProcessTasksForUpload, + AddStorachaUploadForUpload: storachaAPI.AddStorachaUploadForUpload, + RemoveBadFSEntry: scansAPI.RemoveBadFSEntry, + RemoveBadNodes: dagsAPI.RemoveBadNodes, + RemoveShard: blobsAPI.RemoveShard, + Publisher: cfg.bus, } return API{ diff --git a/pkg/preparation/storacha/storacha.go b/pkg/preparation/storacha/storacha.go index 7c5a7153..7ec66230 100644 --- a/pkg/preparation/storacha/storacha.go +++ b/pkg/preparation/storacha/storacha.go @@ -20,11 +20,6 @@ import ( "github.com/storacha/go-ucanto/core/delegation" "github.com/storacha/go-ucanto/core/receipt/fx" "github.com/storacha/go-ucanto/did" - "go.opentelemetry.io/otel" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" - "github.com/storacha/guppy/pkg/bus" "github.com/storacha/guppy/pkg/bus/events" "github.com/storacha/guppy/pkg/client" @@ -34,6 +29,9 @@ import ( gtypes "github.com/storacha/guppy/pkg/preparation/types" "github.com/storacha/guppy/pkg/preparation/types/id" "github.com/storacha/guppy/pkg/preparation/uploads" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/trace" ) var ( @@ -72,9 +70,11 @@ type API struct { var _ uploads.FindShardAddTasksForUploadFunc = API{}.FindShardAddTasksForUpload var _ uploads.FindIndexAddTasksForUploadFunc = API{}.FindIndexAddTasksForUpload +var _ uploads.FindShardPostProcessTasksForUploadFunc = API{}.FindShardPostProcessTasksForUpload +var _ uploads.FindIndexPostProcessTasksForUploadFunc = API{}.FindIndexPostProcessTasksForUpload var _ uploads.AddStorachaUploadForUploadFunc = API{}.AddStorachaUploadForUpload -func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.BlobAddTask, error) { +func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { ctx, span := tracer.Start(ctx, "find-shard-add-tasks-for-upload") defer span.End() @@ -84,9 +84,9 @@ func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadI } span.AddEvent("found closed shards", trace.WithAttributes(attribute.Int("shards", len(closedShards)))) - tasks := make([]gtypes.BlobAddTask, 0, len(closedShards)) + tasks := make([]gtypes.IDTask, 0, len(closedShards)) for _, shard := range closedShards { - tasks = append(tasks, gtypes.BlobAddTask{ + tasks = append(tasks, gtypes.IDTask{ ID: shard.ID(), Run: func(ctx context.Context) (error, error) { if err := a.addBlob(ctx, shard, spaceDID); err != nil { @@ -106,7 +106,7 @@ func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadI return tasks, nil } -func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.BlobAddTask, error) { +func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { ctx, span := tracer.Start(ctx, "find-index-add-tasks-for-upload") defer span.End() @@ -116,9 +116,9 @@ func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadI } span.AddEvent("found closed indexes", trace.WithAttributes(attribute.Int("indexes", len(closedIndexes)))) - tasks := make([]gtypes.BlobAddTask, 0, len(closedIndexes)) + tasks := make([]gtypes.IDTask, 0, len(closedIndexes)) for _, index := range closedIndexes { - tasks = append(tasks, gtypes.BlobAddTask{ + tasks = append(tasks, gtypes.IDTask{ ID: index.ID(), Run: func(ctx context.Context) (error, error) { if err := a.addBlob(ctx, index, spaceDID); err != nil { @@ -138,74 +138,41 @@ func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadI return tasks, nil } -func (a API) PostProcessUploadedShards(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error { - ctx, span := tracer.Start(ctx, "post-process-uploaded-shards") +func (a API) FindShardPostProcessTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { + ctx, span := tracer.Start(ctx, "find-shard-post-process-tasks-for-upload") defer span.End() + uploadedShards, err := a.Repo.ShardsForUploadByState(ctx, uploadID, model.BlobStateUploaded) if err != nil { - return fmt.Errorf("failed to get uploaded shards for post processing %s: %w", uploadID, err) + return nil, fmt.Errorf("failed to get uploaded shards for upload %s: %w", uploadID, err) } span.AddEvent("found uploaded shards", trace.WithAttributes(attribute.Int("shards", len(uploadedShards)))) - blobs := make([]model.Blob, len(uploadedShards)) - for i, shard := range uploadedShards { - blobs[i] = shard - } - return a.postProcessBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { - var opts []client.FilecoinOfferOption - if blob.PDPAccept() != nil { - opts = append(opts, client.WithPDPAcceptInvocation(blob.PDPAccept())) - } - if err := a.filecoinOffer(ctx, blob, spaceDID, opts...); err != nil { - return gtypes.NewBlobUploadError(blob.ID(), err) - } - return nil - }) -} - -func (a API) postProcessBlobs(ctx context.Context, blobs []model.Blob, spaceDID did.DID, afterAdded func(blob model.Blob) error) error { - // Ensure at least 1 parallelism - if a.BlobUploadParallelism < 1 { - a.BlobUploadParallelism = 1 - } - - sem := make(chan struct{}, a.BlobUploadParallelism) - blobUploadErrorCh := make(chan error, len(blobs)) - eg, gctx := errgroup.WithContext(ctx) - for _, blob := range blobs { - sem <- struct{}{} - eg.Go(func() error { - defer func() { <-sem }() - if err := a.postProcessBlob(gctx, blob, spaceDID, afterAdded); err != nil { - err = fmt.Errorf("failed to add blob %s: %w", blob, err) - var errBlobUpload gtypes.BlobUploadError - if errors.As(err, &errBlobUpload) { - blobUploadErrorCh <- err + tasks := make([]gtypes.IDTask, 0, len(uploadedShards)) + for _, shard := range uploadedShards { + tasks = append(tasks, gtypes.IDTask{ + ID: shard.ID(), + Run: func(ctx context.Context) (error, error) { + err := a.postProcessBlob(ctx, shard, spaceDID, func(blob model.Blob) error { + var opts []client.FilecoinOfferOption + if blob.PDPAccept() != nil { + opts = append(opts, client.WithPDPAcceptInvocation(blob.PDPAccept())) + } + if err := a.filecoinOffer(ctx, blob, spaceDID, opts...); err != nil { + return fmt.Errorf("failed to `filecoin/offer` shard %s: %w", blob, err) + } return nil + }) + if err != nil { + log.Errorf("failed to post-process shard %s: %v", shard, err) + return nil, fmt.Errorf("failed to post-process shard %s: %w", shard, err) } - log.Errorf("%v", err) - return err - } - log.Infof("Successfully post-processed blob %s", blob.ID()) - return nil + log.Infof("Successfully post-processed shard %s", shard.ID()) + return nil, nil + }, }) } - - terminalErr := eg.Wait() - close(blobUploadErrorCh) - - if terminalErr != nil { - return terminalErr - } - - var blobUploadErrors []error - for err := range blobUploadErrorCh { - blobUploadErrors = append(blobUploadErrors, err) - } - if len(blobUploadErrors) > 0 { - return gtypes.NewBlobUploadErrors(blobUploadErrors) - } - return nil + return tasks, nil } func (a API) readerForBlob(ctx context.Context, blob model.Blob) (io.ReadCloser, error) { @@ -381,28 +348,37 @@ func (a API) filecoinOffer(ctx context.Context, blob model.Blob, spaceDID did.DI return nil } -// PostProcessUploadedIndexes runs post-processing for uploaded indexes, including -// adding them to the space via `space/index/add`. -func (a API) PostProcessUploadedIndexes(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error { - ctx, span := tracer.Start(ctx, "post-process-uploaded-indexes") +func (a API) FindIndexPostProcessTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { + ctx, span := tracer.Start(ctx, "find-index-post-process-tasks-for-upload") defer span.End() uploadedIndexes, err := a.Repo.IndexesForUploadByState(ctx, uploadID, model.BlobStateUploaded) if err != nil { - return fmt.Errorf("failed to get uploaded indexes for upload %s: %w", uploadID, err) + return nil, fmt.Errorf("failed to get uploaded indexes for upload %s: %w", uploadID, err) } span.AddEvent("found uploaded indexes", trace.WithAttributes(attribute.Int("indexes", len(uploadedIndexes)))) - blobs := make([]model.Blob, len(uploadedIndexes)) - for i, shard := range uploadedIndexes { - blobs[i] = shard + tasks := make([]gtypes.IDTask, 0, len(uploadedIndexes)) + for _, index := range uploadedIndexes { + tasks = append(tasks, gtypes.IDTask{ + ID: index.ID(), + Run: func(ctx context.Context) (error, error) { + err := a.postProcessBlob(ctx, index, spaceDID, func(blob model.Blob) error { + // Use a placeholder for the root because it doesn't matter what it is, + // and we don't want to wait for it to be known. It shouldn't really be + // something the index knows at all. + return a.Client.SpaceIndexAdd(ctx, blob.CID(), blob.Size(), util.PlaceholderCID, spaceDID) + }) + if err != nil { + log.Errorf("failed to post-process index %s: %v", index, err) + return nil, fmt.Errorf("failed to post-process index %s: %w", index, err) + } + log.Infof("Successfully post-processed index %s", index.ID()) + return nil, nil + }, + }) } - return a.postProcessBlobs(ctx, blobs, spaceDID, func(blob model.Blob) error { - // Use a placeholder for the root because it doesn't matter what it is, - // and we don't want to wait for it to be known. It shouldn't really be - // something the index knows at all. - return a.Client.SpaceIndexAdd(ctx, blob.CID(), blob.Size(), util.PlaceholderCID, spaceDID) - }) + return tasks, nil } func (a API) AddStorachaUploadForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error { diff --git a/pkg/preparation/storacha/storacha_test.go b/pkg/preparation/storacha/storacha_test.go index 72b63935..0cca97b3 100644 --- a/pkg/preparation/storacha/storacha_test.go +++ b/pkg/preparation/storacha/storacha_test.go @@ -105,8 +105,13 @@ func TestFindShardAddTasksForUpload(t *testing.T) { require.Equal(t, spaceDID, client.SpaceBlobAddInvocations[0].Space) // Now run post processing. - err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) + ppTasks, err := api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, err := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload shards firstShard, err = repo.GetShardByID(t.Context(), firstShard.ID()) require.NoError(t, err) @@ -154,8 +159,13 @@ func TestFindShardAddTasksForUpload(t *testing.T) { require.Equal(t, spaceDID, client.SpaceBlobAddInvocations[1].Space) // Now run post processing. - err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) + ppTasks, err = api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, err := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload second shard secondShard, err = repo.GetShardByID(t.Context(), secondShard.ID()) @@ -218,8 +228,13 @@ func TestFindShardAddTasksForUpload(t *testing.T) { require.NoError(t, err) require.ErrorContains(t, nonFatal, "simulated SpaceBlobAdd error") } - err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) + ppTasks, err := api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, err := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // It should have `space/blob/add`ed (and failed)... require.Len(t, client.SpaceBlobAddInvocations, 1) @@ -245,8 +260,13 @@ func TestFindShardAddTasksForUpload(t *testing.T) { require.NoError(t, nonFatal) require.NoError(t, err) } - err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) - require.ErrorContains(t, err, "simulated SpaceBlobReplicate error") + ppTasks, err = api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) + require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, fatalErr := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.ErrorContains(t, fatalErr, "simulated SpaceBlobReplicate error") + } // It should have `space/blob/add`ed again... require.Len(t, client.SpaceBlobAddInvocations, 2) @@ -268,8 +288,13 @@ func TestFindShardAddTasksForUpload(t *testing.T) { require.NoError(t, nonFatal) require.NoError(t, err) } - err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) - require.ErrorContains(t, err, "simulated FilecoinOffer error") + ppTasks, err = api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) + require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, fatalErr := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.ErrorContains(t, fatalErr, "simulated FilecoinOffer error") + } // It should NOT `space/blob/add` again... require.Len(t, client.SpaceBlobAddInvocations, 2) @@ -322,8 +347,13 @@ func TestFindShardAddTasksForUpload(t *testing.T) { require.NoError(t, nonFatal) require.NoError(t, err) } - err = api.PostProcessUploadedShards(t.Context(), upload.ID(), spaceDID) + ppTasks, err := api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, err := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // It should `space/blob/add`... require.Len(t, client.SpaceBlobAddInvocations, 1) @@ -414,8 +444,13 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { require.Equal(t, spaceDID, client.SpaceBlobAddInvocations[0].Space) // Now run post processing. - err = api.PostProcessUploadedIndexes(t.Context(), upload.ID(), spaceDID) + ppTasks, err := api.FindIndexPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, err := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload first shard firstIndex, err = repo.GetIndexByID(t.Context(), indexes[0].ID()) @@ -466,8 +501,13 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { require.Equal(t, spaceDID, client.SpaceBlobAddInvocations[1].Space) // Now run post processing. - err = api.PostProcessUploadedIndexes(t.Context(), upload.ID(), spaceDID) + ppTasks, err = api.FindIndexPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) + for _, task := range ppTasks { + nonFatal, err := task.Run(t.Context()) + require.NoError(t, nonFatal) + require.NoError(t, err) + } // Reload second shard secondIndex, err = repo.GetIndexByID(t.Context(), indexes[1].ID()) diff --git a/pkg/preparation/types/errors.go b/pkg/preparation/types/errors.go index 978f43e7..5a3573c0 100644 --- a/pkg/preparation/types/errors.go +++ b/pkg/preparation/types/errors.go @@ -193,7 +193,11 @@ func (e BlobUploadErrors) Unwrap() []error { return e.errs } -type BlobAddTask struct { +// IDTask represents a task that can be identified and deduplicated by an +// [id.ID]. The task's Run function returns two errors: a non-fatal error that +// can be collected and reported after all tasks have completed, and a fatal +// error that should cause immediate cancellation of all other tasks. +type IDTask struct { ID id.ID Run func(context.Context) (error, error) } diff --git a/pkg/preparation/uploads/uploads.go b/pkg/preparation/uploads/uploads.go index ccdf9343..2563167f 100644 --- a/pkg/preparation/uploads/uploads.go +++ b/pkg/preparation/uploads/uploads.go @@ -35,29 +35,29 @@ type AddNodeToUploadShardsFunc func(ctx context.Context, uploadID id.UploadID, s type AddShardsToUploadIndexesFunc func(ctx context.Context, uploadID id.UploadID, indexCB func(index *blobsmodel.Index) error) error type CloseUploadShardsFunc func(ctx context.Context, uploadID id.UploadID, shardCB func(shard *blobsmodel.Shard) error) error type CloseUploadIndexesFunc func(ctx context.Context, uploadID id.UploadID, indexCB func(index *blobsmodel.Index) error) error -type FindShardAddTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.BlobAddTask, error) -type FindIndexAddTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.BlobAddTask, error) +type FindShardAddTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) +type FindIndexAddTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) type AddNodesToUploadShardsFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, shardCB func(shard *blobsmodel.Shard) error) error -type PostProcessUploadedShardsFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error -type PostProcessUploadedIndexesFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error +type FindShardPostProcessTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) +type FindIndexPostProcessTasksForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) type AddStorachaUploadForUploadFunc func(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) error type RemoveBadFSEntryFunc func(ctx context.Context, spaceDID did.DID, fsEntryID id.FSEntryID) error type RemoveBadNodesFunc func(ctx context.Context, spaceDID did.DID, nodeCIDs []cid.Cid) error type RemoveShardFunc func(ctx context.Context, shardID id.ShardID) error type API struct { - Repo Repo - ExecuteScan ExecuteScanFunc - ExecuteDagScansForUpload ExecuteDagScansForUploadFunc - BlobUploadParallelism int - FindShardAddTasksForUpload FindShardAddTasksForUploadFunc - FindIndexAddTasksForUpload FindIndexAddTasksForUploadFunc - PostProcessUploadedShards PostProcessUploadedShardsFunc - PostProcessUploadedIndexes PostProcessUploadedIndexesFunc - AddStorachaUploadForUpload AddStorachaUploadForUploadFunc - RemoveBadFSEntry RemoveBadFSEntryFunc - RemoveBadNodes RemoveBadNodesFunc - RemoveShard RemoveShardFunc + Repo Repo + ExecuteScan ExecuteScanFunc + ExecuteDagScansForUpload ExecuteDagScansForUploadFunc + BlobUploadParallelism int + FindShardAddTasksForUpload FindShardAddTasksForUploadFunc + FindIndexAddTasksForUpload FindIndexAddTasksForUploadFunc + FindShardPostProcessTasksForUpload FindShardPostProcessTasksForUploadFunc + FindIndexPostProcessTasksForUpload FindIndexPostProcessTasksForUploadFunc + AddStorachaUploadForUpload AddStorachaUploadForUploadFunc + RemoveBadFSEntry RemoveBadFSEntryFunc + RemoveBadNodes RemoveBadNodesFunc + RemoveShard RemoveShardFunc // AddNodesToUploadShards assigns all unsharded nodes for an upload to shards. AddNodesToUploadShards AddNodesToUploadShardsFunc @@ -747,22 +747,30 @@ func runPostProcessShardWorker( span.End() }() + var inFlightShards sync.Map + _, err = Worker( ctx, uploadedShardsAvailable, - 1, + api.BlobUploadParallelism, // findWork func(ctx context.Context) ([]func(context.Context) (error, error), error) { - return []func(context.Context) (error, error){ - func(ctx context.Context) (error, error) { - err := api.PostProcessUploadedShards(ctx, uploadID, spaceDID) - if err != nil { - return nil, fmt.Errorf("`post-processing shards for upload %s: %w", uploadID, err) - } - return nil, nil - }, - }, nil + rawTasks, err := api.FindShardPostProcessTasksForUpload(ctx, uploadID, spaceDID) + if err != nil { + return nil, err + } + var tasks []func(context.Context) (error, error) + for _, raw := range rawTasks { + if _, loaded := inFlightShards.LoadOrStore(raw.ID, struct{}{}); loaded { + continue + } + tasks = append(tasks, func(ctx context.Context) (error, error) { + defer inFlightShards.Delete(raw.ID) + return raw.Run(ctx) + }) + } + return tasks, nil }, // finalize @@ -771,7 +779,6 @@ func runPostProcessShardWorker( if err != nil { return fmt.Errorf("`upload/add`ing upload %s: %w", uploadID, err) } - return nil }, ) @@ -881,22 +888,30 @@ func runPostProcessIndexWorker( span.End() }() + var inFlightIndexes sync.Map + _, err = Worker( ctx, uploadedIndexesAvailable, - 1, + api.BlobUploadParallelism, // findWork func(ctx context.Context) ([]func(context.Context) (error, error), error) { - return []func(context.Context) (error, error){ - func(ctx context.Context) (error, error) { - err := api.PostProcessUploadedIndexes(ctx, uploadID, spaceDID) - if err != nil { - return nil, fmt.Errorf("`post-processing indexes for upload %s: %w", uploadID, err) - } - return nil, nil - }, - }, nil + rawTasks, err := api.FindIndexPostProcessTasksForUpload(ctx, uploadID, spaceDID) + if err != nil { + return nil, err + } + var tasks []func(context.Context) (error, error) + for _, raw := range rawTasks { + if _, loaded := inFlightIndexes.LoadOrStore(raw.ID, struct{}{}); loaded { + continue + } + tasks = append(tasks, func(ctx context.Context) (error, error) { + defer inFlightIndexes.Delete(raw.ID) + return raw.Run(ctx) + }) + } + return tasks, nil }, // finalize From f00d08d39758bba6c82d9754fac197fa3b9e593f Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Thu, 23 Apr 2026 11:19:57 -0400 Subject: [PATCH 6/7] refactor: Rm unused param `spaceDID` --- pkg/preparation/uploads/uploads.go | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/pkg/preparation/uploads/uploads.go b/pkg/preparation/uploads/uploads.go index 2563167f..c1e76e7c 100644 --- a/pkg/preparation/uploads/uploads.go +++ b/pkg/preparation/uploads/uploads.go @@ -242,12 +242,12 @@ func (a API) ExecuteUpload(ctx context.Context, uploadID id.UploadID, spaceDID d return cid.Undef, fmt.Errorf("handling bad FS entries worker error [%w]: %w", workersErr, err) } case errors.As(workersErr, &blobUploadErrors): - err := a.handleBadBlobUploads(ctx, uploadID, spaceDID, blobUploadErrors) + err := a.handleBadBlobUploads(ctx, uploadID, blobUploadErrors) if err != nil { return cid.Undef, fmt.Errorf("handling bad shard uploads worker error [%w]: %w", workersErr, err) } case errors.As(workersErr, &badNodesErr): - err := a.handleBadNodes(ctx, uploadID, spaceDID, badNodesErr) + err := a.handleBadNodes(ctx, uploadID, badNodesErr) if err != nil { return cid.Undef, fmt.Errorf("handling bad nodes worker error [%w]: %w", workersErr, err) } @@ -289,13 +289,13 @@ func (a API) handleBadFSEntries(ctx context.Context, uploadID id.UploadID, badFS return nil } -func (a API) handleBadBlobUploads(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, blobUploadErrors types.BlobUploadErrors) error { +func (a API) handleBadBlobUploads(ctx context.Context, uploadID id.UploadID, blobUploadErrors types.BlobUploadErrors) error { // when there's a bad shard upload, it's not based on a problem locally usually, unless bad nodes were read during upload for _, e := range blobUploadErrors.Unwrap() { // bad nodes error can happen from reading car during upload var badNodesErr types.BadNodesError if errors.As(e, &badNodesErr) { - err := a.handleBadNodes(ctx, uploadID, spaceDID, badNodesErr) + err := a.handleBadNodes(ctx, uploadID, badNodesErr) if err != nil { return err } @@ -305,7 +305,7 @@ func (a API) handleBadBlobUploads(ctx context.Context, uploadID id.UploadID, spa return nil } -func (a API) handleBadNodes(ctx context.Context, uploadID id.UploadID, spaceDID did.DID, badNodesErr types.BadNodesError) error { +func (a API) handleBadNodes(ctx context.Context, uploadID id.UploadID, badNodesErr types.BadNodesError) error { upload, err := a.Repo.GetUploadByID(ctx, uploadID) if err != nil { return fmt.Errorf("getting upload %s after finding bad nodes: %w", uploadID, err) From 44c6c7ec8182d41689d2c3989513b71105ff6b73 Mon Sep 17 00:00:00 2001 From: Petra Jaros Date: Fri, 24 Apr 2026 19:14:34 -0400 Subject: [PATCH 7/7] refactor: Rework the worker code and its errors --- pkg/preparation/internal/worker/group.go | 128 ++++++++++++++++ pkg/preparation/internal/worker/group_test.go | 78 ++++++++++ pkg/preparation/internal/worker/types.go | 133 +++++++++++++++++ pkg/preparation/internal/worker/types_test.go | 94 ++++++++++++ pkg/preparation/internal/worker/worker.go | 81 ++++++++++ .../worker}/worker_test.go | 118 +++++++++------ pkg/preparation/storacha/storacha.go | 77 +++++----- pkg/preparation/storacha/storacha_test.go | 74 ++++------ pkg/preparation/types/errors.go | 10 -- pkg/preparation/types/idtask.go | 13 ++ pkg/preparation/uploads/uploads.go | 139 +++++++++++------- pkg/preparation/uploads/worker.go | 101 ------------- 12 files changed, 745 insertions(+), 301 deletions(-) create mode 100644 pkg/preparation/internal/worker/group.go create mode 100644 pkg/preparation/internal/worker/group_test.go create mode 100644 pkg/preparation/internal/worker/types.go create mode 100644 pkg/preparation/internal/worker/types_test.go create mode 100644 pkg/preparation/internal/worker/worker.go rename pkg/preparation/{uploads => internal/worker}/worker_test.go (53%) create mode 100644 pkg/preparation/types/idtask.go delete mode 100644 pkg/preparation/uploads/worker.go diff --git a/pkg/preparation/internal/worker/group.go b/pkg/preparation/internal/worker/group.go new file mode 100644 index 00000000..65b7561f --- /dev/null +++ b/pkg/preparation/internal/worker/group.go @@ -0,0 +1,128 @@ +package worker + +import ( + "context" + "sync" +) + +// Group runs a collection of tasks that return [TaskError]. Non-fatal errors +// are accumulated; a fatal error from any task cancels the context derived by +// [WithContext] so siblings watching that context can exit early. It's +// analogous to [golang.org/x/sync/errgroup.Group], with two deliberate +// differences: (1) it distinguishes fatal from non-fatal errors, and (2) it +// accumulates every task's errors rather than keeping only the first. +// +// If the group is cancelled, the cause will be the fatal error. +type Group struct { + ctx context.Context + cancel context.CancelCauseFunc + wg sync.WaitGroup + sem chan struct{} + + mu sync.Mutex + err taskError +} + +// WithContext returns a new [Group] and a derived context. The context is +// cancelled the first time a task reports a fatal error, or the first time +// [Group.Wait] returns, whichever occurs first. +func WithContext(ctx context.Context) (*Group, context.Context) { + ctx, cancel := context.WithCancelCause(ctx) + return &Group{ctx: ctx, cancel: cancel}, ctx +} + +// SetLimit limits the number of concurrently-running tasks to n. Must be +// called before any call to [Group.Go]. A value <= 0 removes any limit. +func (g *Group) SetLimit(n int) { + if n <= 0 { + g.sem = nil + return + } + g.sem = make(chan struct{}, n) +} + +// Go runs the given task in a new goroutine. If a concurrency limit is set, Go +// blocks until a slot is available. The task's [TaskError] is folded into the +// group's accumulated result; if the task returns a fatal error, the derived +// context is cancelled. +// +// If the group's context is cancelled when Go is called or while waiting for a +// slot, the task is not started and Go returns immediately. +func (g *Group) Go(task func() TaskError) { + if g.sem != nil { + select { + case g.sem <- struct{}{}: + case <-g.ctx.Done(): + return + } + } + g.launch(task) +} + +// TryGo runs the given task in a new goroutine if a slot is available. It +// returns true if the task was started, false otherwise. If no concurrency +// limit is set, TryGo always starts the task and returns true. +func (g *Group) TryGo(task func() TaskError) bool { + if g.sem != nil { + select { + case g.sem <- struct{}{}: + default: + return false + } + } + g.launch(task) + return true +} + +func (g *Group) launch(task func() TaskError) { + g.wg.Add(1) + go func() { + defer func() { + if g.sem != nil { + <-g.sem + } + g.wg.Done() + }() + select { + case <-g.ctx.Done(): + // Group has cancelled before this task started. Skip it silently: + // the real fatal is already in the accumulator, and an un-started + // task has nothing to report. Tasks already executing when cancel + // fires still run to completion and report whatever they want. + return + default: + } + g.collect(task()) + }() +} + +// Wait blocks until all goroutines started with [Group.Go] or [Group.TryGo] +// have returned, then returns the accumulated [TaskError]. It returns nil iff +// no task reported a fatal error and no task reported any non-fatal errors. +func (g *Group) Wait() TaskError { + g.wg.Wait() + + // Clean up resources: a `cancel` must always be called eventually. We've just + // `Wait()`ed for all tasks to finish, so this won't stop anything. + g.cancel(nil) + + g.mu.Lock() + defer g.mu.Unlock() + if g.err.isEmpty() { + return nil + } + result := g.err + return &result +} + +func (g *Group) collect(result TaskError) { + if result == nil { + return + } + g.mu.Lock() + g.err.add(result) + g.mu.Unlock() + if result.FatalError() != nil { + g.cancel(result.FatalError()) + } +} diff --git a/pkg/preparation/internal/worker/group_test.go b/pkg/preparation/internal/worker/group_test.go new file mode 100644 index 00000000..1530ccec --- /dev/null +++ b/pkg/preparation/internal/worker/group_test.go @@ -0,0 +1,78 @@ +package worker_test + +import ( + "errors" + "sync/atomic" + "testing" + + "github.com/storacha/guppy/pkg/preparation/internal/worker" + "github.com/stretchr/testify/require" +) + +func TestGroupLaunchSkipsIfCancelled(t *testing.T) { + t.Run("task is skipped if group ctx is already cancelled before launch", func(t *testing.T) { + fatalErr := errors.New("fatal") + + g, gctx := worker.WithContext(t.Context()) + + // First task cancels the group by returning a fatal. + g.Go(func() worker.TaskError { + return worker.NewFatalError(fatalErr) + }) + + // Wait for cancellation to propagate before enqueuing the second + // task, so the second task's goroutine sees a cancelled ctx on entry. + <-gctx.Done() + + var ranSecond atomic.Bool + g.Go(func() worker.TaskError { + ranSecond.Store(true) + return worker.NewFatalError(errors.New("should not be reported")) + }) + + res := g.Wait() + require.False(t, ranSecond.Load(), "task enqueued after cancellation should be skipped") + require.NotNil(t, res) + require.True(t, res.IsFatal()) + require.ErrorIs(t, res.FatalError(), fatalErr) + require.NotContains(t, res.FatalError().Error(), "should not be reported") + }) + + t.Run("in-flight task can still report errors after cancellation", func(t *testing.T) { + firstFatal := errors.New("first fatal") + secondFatal := errors.New("second fatal from in-flight task") + + g, gctx := worker.WithContext(t.Context()) + + started := make(chan struct{}) + release := make(chan struct{}) + + // In-flight task: signals it has started, waits to be released, then + // returns a fatal. It's already executing when the group cancels. + g.Go(func() worker.TaskError { + close(started) + <-release + return worker.NewFatalError(secondFatal) + }) + + // Wait for the in-flight task to be running. + <-started + + // Fire a fatal from another task to cancel the group. + g.Go(func() worker.TaskError { + return worker.NewFatalError(firstFatal) + }) + + // Wait for gctx to be cancelled. + <-gctx.Done() + + // Release the in-flight task so it can complete. + close(release) + + res := g.Wait() + require.NotNil(t, res) + require.True(t, res.IsFatal()) + require.ErrorIs(t, res.FatalError(), firstFatal, "first fatal should be recorded") + require.ErrorIs(t, res.FatalError(), secondFatal, "in-flight task's fatal should also be recorded") + }) +} diff --git a/pkg/preparation/internal/worker/types.go b/pkg/preparation/internal/worker/types.go new file mode 100644 index 00000000..1da80af4 --- /dev/null +++ b/pkg/preparation/internal/worker/types.go @@ -0,0 +1,133 @@ +package worker + +import ( + "context" + "errors" + "fmt" + "strings" +) + +// TaskError is the error type returned by a [Task]. A TaskError may carry a +// fatal error (which causes the worker to cancel dispatch of remaining tasks) +// and/or a set of non-fatal errors (which are collected and reported after +// all tasks complete). +type TaskError interface { + error + NonFatalErrors() []error + FatalError() error + IsFatal() bool +} + +// NewFatalError returns a [TaskError] carrying the given error as a fatal +// error. If err is nil, it returns nil. +func NewFatalError(err error) TaskError { + if err == nil { + return nil + } + return &taskError{fatalError: err} +} + +// NewNonFatalError returns a [TaskError] carrying the given errors as +// non-fatal errors. Any nil entries are dropped; if the result would be empty, +// it returns nil. +func NewNonFatalError(errs ...error) TaskError { + var filtered []error + for _, err := range errs { + if err != nil { + filtered = append(filtered, err) + } + } + if len(filtered) == 0 { + return nil + } + return &taskError{nonFatalErrors: filtered} +} + +type taskError struct { + nonFatalErrors []error + fatalError error +} + +// Task is a worker task. It returns a [TaskError] to signal the result: +// - return nil for success, +// - return [NewFatalError] (or any TaskError whose IsFatal is true) to abort +// further dispatch, +// - return [NewNonFatalError] to report errors that should be collected but +// not abort dispatch. +type Task func(context.Context) TaskError + +func (e *taskError) NonFatalErrors() []error { + return e.nonFatalErrors +} + +func (e *taskError) FatalError() error { + return e.fatalError +} + +func (e *taskError) IsFatal() bool { + return e.fatalError != nil +} + +func (e *taskError) Error() string { + nonFatalString := "" + if len(e.nonFatalErrors) > 0 { + nonFatalString = "non-fatal errors:" + } + for _, err := range e.nonFatalErrors { + nonFatalString += fmt.Sprintf("\n- %s", err) + } + + fatalString := "" + if e.fatalError != nil { + fatalString = fmt.Sprintf("fatal error: %s", e.fatalError) + } + + return fmt.Sprintf("worker encountered %s", strings.Join([]string{nonFatalString, fatalString}, "\n")) +} + +func (e *taskError) Unwrap() []error { + errs := make([]error, len(e.nonFatalErrors)) + copy(errs, e.nonFatalErrors) + if e.fatalError != nil { + errs = append(errs, e.fatalError) + } + return errs +} + +// add folds another [TaskError] into the receiver in place. +func (e *taskError) add(other TaskError) { + if other == nil { + return + } + e.nonFatalErrors = append(e.nonFatalErrors, other.NonFatalErrors()...) + if other.FatalError() != nil { + e.fatalError = errors.Join(e.fatalError, other.FatalError()) + } +} + +// isEmpty reports whether the taskError carries neither a fatal error nor any +// non-fatal errors. +func (e *taskError) isEmpty() bool { + return e.fatalError == nil && len(e.nonFatalErrors) == 0 +} + +// Join combines two [TaskError]s. Non-fatal errors are concatenated; fatal +// errors are joined with [errors.Join]. Either operand may be nil. +func Join(a TaskError, b TaskError) TaskError { + switch { + case a == nil && b == nil: + return nil + case a == nil: + return b + case b == nil: + return a + } + joined := &taskError{ + nonFatalErrors: append(a.NonFatalErrors(), b.NonFatalErrors()...), + fatalError: errors.Join(a.FatalError(), b.FatalError()), + } + if joined.isEmpty() { + return nil + } + return joined +} diff --git a/pkg/preparation/internal/worker/types_test.go b/pkg/preparation/internal/worker/types_test.go new file mode 100644 index 00000000..55aa0b8e --- /dev/null +++ b/pkg/preparation/internal/worker/types_test.go @@ -0,0 +1,94 @@ +package worker_test + +import ( + "errors" + "testing" + + "github.com/storacha/guppy/pkg/preparation/internal/worker" + "github.com/stretchr/testify/require" +) + +func TestNewFatalError(t *testing.T) { + t.Run("returns nil when err is nil", func(t *testing.T) { + require.Nil(t, worker.NewFatalError(nil)) + }) + + t.Run("wraps a non-nil error as fatal", func(t *testing.T) { + err := errors.New("boom") + te := worker.NewFatalError(err) + require.NotNil(t, te) + require.True(t, te.IsFatal()) + require.ErrorIs(t, te.FatalError(), err) + require.Empty(t, te.NonFatalErrors()) + }) +} + +func TestNewNonFatalError(t *testing.T) { + t.Run("returns nil when no errors are supplied", func(t *testing.T) { + require.Nil(t, worker.NewNonFatalError()) + }) + + t.Run("returns nil when all supplied errors are nil", func(t *testing.T) { + require.Nil(t, worker.NewNonFatalError(nil, nil)) + }) + + t.Run("drops nil entries and keeps non-nil ones", func(t *testing.T) { + a := errors.New("a") + b := errors.New("b") + te := worker.NewNonFatalError(nil, a, nil, b) + require.NotNil(t, te) + require.False(t, te.IsFatal()) + require.Nil(t, te.FatalError()) + require.Len(t, te.NonFatalErrors(), 2) + require.ErrorIs(t, te.NonFatalErrors()[0], a) + require.ErrorIs(t, te.NonFatalErrors()[1], b) + }) +} + +func TestTaskErrorUnwrap(t *testing.T) { + t.Run("exposes non-fatal and fatal errors via errors.Is", func(t *testing.T) { + nonFatal := errors.New("non-fatal") + fatal := errors.New("fatal") + te := worker.Join( + worker.NewNonFatalError(nonFatal), + worker.NewFatalError(fatal), + ) + require.NotNil(t, te) + require.ErrorIs(t, te, nonFatal) + require.ErrorIs(t, te, fatal) + }) +} + +func TestJoin(t *testing.T) { + t.Run("returns nil when both operands are nil", func(t *testing.T) { + require.Nil(t, worker.Join(nil, nil)) + }) + + t.Run("returns the non-nil operand when the other is nil", func(t *testing.T) { + err := errors.New("err") + te := worker.NewFatalError(err) + require.Equal(t, te, worker.Join(te, nil)) + require.Equal(t, te, worker.Join(nil, te)) + }) + + t.Run("concatenates non-fatal errors and joins fatals", func(t *testing.T) { + nonFatalA := errors.New("non-fatal A") + nonFatalB := errors.New("non-fatal B") + fatalA := errors.New("fatal A") + fatalB := errors.New("fatal B") + + a := worker.Join(worker.NewNonFatalError(nonFatalA), worker.NewFatalError(fatalA)) + b := worker.Join(worker.NewNonFatalError(nonFatalB), worker.NewFatalError(fatalB)) + + joined := worker.Join(a, b) + require.NotNil(t, joined) + require.True(t, joined.IsFatal()) + + require.Len(t, joined.NonFatalErrors(), 2) + require.ErrorIs(t, joined.NonFatalErrors()[0], nonFatalA) + require.ErrorIs(t, joined.NonFatalErrors()[1], nonFatalB) + + require.ErrorIs(t, joined.FatalError(), fatalA) + require.ErrorIs(t, joined.FatalError(), fatalB) + }) +} diff --git a/pkg/preparation/internal/worker/worker.go b/pkg/preparation/internal/worker/worker.go new file mode 100644 index 00000000..27a8866c --- /dev/null +++ b/pkg/preparation/internal/worker/worker.go @@ -0,0 +1,81 @@ +package worker + +import ( + "context" + "fmt" + + "github.com/storacha/guppy/internal/ctxutil" +) + +// Run executes a worker loop. The loop runs tasks in parallel. Run blocks until +// the `workAvailable` channel closes and all tasks are complete. +// +// The loop waits for a signal on the `workAvailable` channel, then calls +// `findWork` to get a batch of tasks to run and runs them with a maximum +// parallelism of `parallelism`, queuing any additional tasks. Any time the +// queue becomes empty, the loop waits for the next signal to find more work, +// until the `workAvailable` channel is closed. When the channel is closed, and +// all work is complete, the loop calls `finalize` and returns. +// +// If any task returns a [TaskError] whose [TaskError.IsFatal] is true, all +// running tasks are cancelled and no new tasks are started. +// +// Returns nil iff no task reported any errors, fatal or non-fatal. Otherwise, +// returns a [TaskError] describing what accumulated. +func Run( + ctx context.Context, + workAvailable <-chan struct{}, + parallelism int, + findWork func(ctx context.Context) ([]Task, error), + finalize func() error, +) TaskError { + g, gctx := WithContext(ctx) + g.SetLimit(parallelism) + + for { + select { + case <-ctx.Done(): + // External cancellation. The group has already been cancelled, so wait + // and return, but include the cancellation cause as a fatal error. + return Join(g.Wait(), NewFatalError(ctxutil.Cause(ctx))) + + case <-gctx.Done(): + // Internal cancellation. Wait, and return the result. + return g.Wait() + + case _, ok := <-workAvailable: + if !ok { + result := g.Wait() + if result != nil && result.IsFatal() { + return result + } + if finalize != nil { + if ferr := finalize(); ferr != nil { + return Join(result, NewFatalError(fmt.Errorf("worker finalize encountered an error: %w", ferr))) + } + } + return result + } + + tasks, err := findWork(ctx) + if err != nil { + // Cancel in-flight siblings and drain before surfacing the + // fatal. + result := g.Wait() + return Join(result, NewFatalError(fmt.Errorf("worker findWork encountered an error: %w", err))) + } + + for _, task := range tasks { + g.Go(func() TaskError { + select { + case <-gctx.Done(): + // If the context is already cancelled, skip starting the task. + return nil + default: + return task(gctx) + } + }) + } + } + } +} diff --git a/pkg/preparation/uploads/worker_test.go b/pkg/preparation/internal/worker/worker_test.go similarity index 53% rename from pkg/preparation/uploads/worker_test.go rename to pkg/preparation/internal/worker/worker_test.go index 84d8fcb1..4a2c8f19 100644 --- a/pkg/preparation/uploads/worker_test.go +++ b/pkg/preparation/internal/worker/worker_test.go @@ -1,4 +1,4 @@ -package uploads_test +package worker_test import ( "context" @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/storacha/guppy/pkg/preparation/uploads" + "github.com/storacha/guppy/pkg/preparation/internal/worker" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -18,27 +18,25 @@ func e(t *testing.T, condition func(collect *assert.CollectT)) { require.EventuallyWithT(t, condition, time.Second, 10*time.Millisecond) } -func task(fn func() error) func(context.Context) (error, error) { - return func(ctx context.Context) (error, error) { - return nil, fn() +// fatalTask wraps a fn returning error into a Task that reports any error as +// fatal. +func fatalTask(fn func() error) worker.Task { + return func(ctx context.Context) worker.TaskError { + return worker.NewFatalError(fn()) } } func TestWorker(t *testing.T) { t.Run("runs the work function for every signal received, then the finalize function when the channel closes", func(t *testing.T) { signalChan := make(chan struct{}, 1) - type result struct { - nonFatals []error - err error - } - resultChan := make(chan result, 1) + resultChan := make(chan worker.TaskError, 1) var runs int var finalizes int go func() { defer close(resultChan) - nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + resultChan <- worker.Run(t.Context(), signalChan, 1, + func(ctx context.Context) ([]worker.Task, error) { runs++ return nil, nil }, @@ -47,7 +45,6 @@ func TestWorker(t *testing.T) { return nil }, ) - resultChan <- result{nonFatals, err} }() require.Equal(t, 0, runs, "worker should not run before signal") @@ -63,30 +60,23 @@ func TestWorker(t *testing.T) { }) res := <-resultChan - require.Nil(t, res.err, "error should be nil after successful runs") - require.Empty(t, res.nonFatals, "non-fatal errors should be empty after successful runs") + require.Nil(t, res, "result should be nil after a clean run") }) t.Run("immediately responds with any work error, skipping the finalizer", func(t *testing.T) { workerErr := errors.New("error in doWork") signalChan := make(chan struct{}, 3) - type result struct { - nonFatals []error - err error - } - resultChan := make(chan result, 1) + resultChan := make(chan worker.TaskError, 1) var runs int var finalizes int go func() { defer close(resultChan) - nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + resultChan <- worker.Run(t.Context(), signalChan, 1, + func(ctx context.Context) ([]worker.Task, error) { runs++ if runs == 2 { - return []func(context.Context) (error, error){ - task(func() error { return workerErr }), - }, nil + return []worker.Task{fatalTask(func() error { return workerErr })}, nil } return nil, nil }, @@ -95,7 +85,6 @@ func TestWorker(t *testing.T) { return nil }, ) - resultChan <- result{nonFatals, err} }() // Send three signals; the second should cause an error, the third should not run @@ -104,26 +93,64 @@ func TestWorker(t *testing.T) { signalChan <- struct{}{} res := <-resultChan - require.ErrorContains(t, res.err, "worker task encountered a fatal error: error in doWork") - require.ErrorIs(t, res.err, workerErr) + require.NotNil(t, res) + require.True(t, res.IsFatal(), "a fatal task error should make the result fatal") + require.ErrorIs(t, res.FatalError(), workerErr) require.LessOrEqual(t, runs, 3, "worker should have stopped after encountering an error") require.Equal(t, 0, finalizes, "finalize function should not have been called") }) + t.Run("collects non-fatal errors and continues dispatching", func(t *testing.T) { + nonFatalErr := errors.New("non-fatal failure") + signalChan := make(chan struct{}, 3) + resultChan := make(chan worker.TaskError, 1) + var runs int + var finalizes int + + go func() { + defer close(resultChan) + resultChan <- worker.Run(t.Context(), signalChan, 1, + func(ctx context.Context) ([]worker.Task, error) { + runs++ + return []worker.Task{ + func(ctx context.Context) worker.TaskError { + return worker.NewNonFatalError(nonFatalErr) + }, + }, nil + }, + func() error { + finalizes++ + return nil + }, + ) + }() + + signalChan <- struct{}{} + signalChan <- struct{}{} + signalChan <- struct{}{} + close(signalChan) + + res := <-resultChan + require.NotNil(t, res) + require.False(t, res.IsFatal(), "non-fatal errors should not make the result fatal") + require.Len(t, res.NonFatalErrors(), 3, "every non-fatal error should be collected") + for _, err := range res.NonFatalErrors() { + require.ErrorIs(t, err, nonFatalErr) + } + require.Equal(t, 3, runs, "worker should have kept dispatching after non-fatal errors") + require.Equal(t, 1, finalizes, "finalize should still run when only non-fatal errors occurred") + }) + t.Run("responds with any finalize error", func(t *testing.T) { finalizerErr := errors.New("error in finalize") signalChan := make(chan struct{}, 3) - type result struct { - nonFatals []error - err error - } - resultChan := make(chan result, 1) + resultChan := make(chan worker.TaskError, 1) var runs int go func() { defer close(resultChan) - nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + resultChan <- worker.Run(t.Context(), signalChan, 1, + func(ctx context.Context) ([]worker.Task, error) { runs++ return nil, nil }, @@ -131,7 +158,6 @@ func TestWorker(t *testing.T) { return finalizerErr }, ) - resultChan <- result{nonFatals, err} }() // Send three signals; all should run @@ -141,30 +167,27 @@ func TestWorker(t *testing.T) { close(signalChan) res := <-resultChan - require.ErrorContains(t, res.err, "worker finalize encountered an error: error in finalize") - require.ErrorIs(t, res.err, finalizerErr) + require.NotNil(t, res) + require.True(t, res.IsFatal(), "finalize error should make the result fatal") + require.ErrorContains(t, res.FatalError(), "worker finalize encountered an error: error in finalize") + require.ErrorIs(t, res.FatalError(), finalizerErr) require.Equal(t, 3, runs, "worker should have run all three times") }) t.Run("ignores a nil finalizer", func(t *testing.T) { signalChan := make(chan struct{}, 1) - type result struct { - nonFatals []error - err error - } - resultChan := make(chan result, 1) + resultChan := make(chan worker.TaskError, 1) var ran bool go func() { defer close(resultChan) - nonFatals, err := uploads.Worker(t.Context(), signalChan, 1, - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + resultChan <- worker.Run(t.Context(), signalChan, 1, + func(ctx context.Context) ([]worker.Task, error) { ran = true return nil, nil }, nil, ) - resultChan <- result{nonFatals, err} }() require.False(t, ran, "worker should not run before signal") @@ -172,7 +195,6 @@ func TestWorker(t *testing.T) { e(t, func(t *assert.CollectT) { require.True(t, ran, "worker should run after signal") }) close(signalChan) res := <-resultChan - require.Nil(t, res.err, "error should be nil after successful runs and no finalizer") - require.Empty(t, res.nonFatals) + require.Nil(t, res, "result should be nil after a clean run with no finalizer") }) } diff --git a/pkg/preparation/storacha/storacha.go b/pkg/preparation/storacha/storacha.go index 7ec66230..e1aa4d9b 100644 --- a/pkg/preparation/storacha/storacha.go +++ b/pkg/preparation/storacha/storacha.go @@ -15,7 +15,7 @@ import ( "github.com/multiformats/go-multicodec" filecoincap "github.com/storacha/go-libstoracha/capabilities/filecoin" spaceblobcap "github.com/storacha/go-libstoracha/capabilities/space/blob" - "github.com/storacha/go-libstoracha/capabilities/types" + captypes "github.com/storacha/go-libstoracha/capabilities/types" "github.com/storacha/go-libstoracha/capabilities/upload" "github.com/storacha/go-ucanto/core/delegation" "github.com/storacha/go-ucanto/core/receipt/fx" @@ -26,7 +26,8 @@ import ( "github.com/storacha/guppy/pkg/internal/util" "github.com/storacha/guppy/pkg/preparation/blobs/model" "github.com/storacha/guppy/pkg/preparation/internal/meteredwriter" - gtypes "github.com/storacha/guppy/pkg/preparation/types" + "github.com/storacha/guppy/pkg/preparation/internal/worker" + "github.com/storacha/guppy/pkg/preparation/types" "github.com/storacha/guppy/pkg/preparation/types/id" "github.com/storacha/guppy/pkg/preparation/uploads" "go.opentelemetry.io/otel" @@ -46,7 +47,7 @@ type Client interface { SpaceIndexAdd(ctx context.Context, indexCID cid.Cid, indexSize uint64, rootCID cid.Cid, space did.DID) error FilecoinOffer(ctx context.Context, space did.DID, content ipld.Link, piece ipld.Link, opts ...client.FilecoinOfferOption) (filecoincap.OfferOk, error) UploadAdd(ctx context.Context, space did.DID, root ipld.Link, shards []ipld.Link) (upload.AddOk, error) - SpaceBlobReplicate(ctx context.Context, space did.DID, blob types.Blob, replicaCount uint, locationCommitment delegation.Delegation) (spaceblobcap.ReplicateOk, fx.Effects, error) + SpaceBlobReplicate(ctx context.Context, space did.DID, blob captypes.Blob, replicaCount uint, locationCommitment delegation.Delegation) (spaceblobcap.ReplicateOk, fx.Effects, error) } var _ Client = (*client.Client)(nil) @@ -74,7 +75,7 @@ var _ uploads.FindShardPostProcessTasksForUploadFunc = API{}.FindShardPostProces var _ uploads.FindIndexPostProcessTasksForUploadFunc = API{}.FindIndexPostProcessTasksForUpload var _ uploads.AddStorachaUploadForUploadFunc = API{}.AddStorachaUploadForUpload -func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { +func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) { ctx, span := tracer.Start(ctx, "find-shard-add-tasks-for-upload") defer span.End() @@ -84,29 +85,29 @@ func (a API) FindShardAddTasksForUpload(ctx context.Context, uploadID id.UploadI } span.AddEvent("found closed shards", trace.WithAttributes(attribute.Int("shards", len(closedShards)))) - tasks := make([]gtypes.IDTask, 0, len(closedShards)) + tasks := make([]types.IDTask, 0, len(closedShards)) for _, shard := range closedShards { - tasks = append(tasks, gtypes.IDTask{ + tasks = append(tasks, types.IDTask{ ID: shard.ID(), - Run: func(ctx context.Context) (error, error) { + Run: func(ctx context.Context) worker.TaskError { if err := a.addBlob(ctx, shard, spaceDID); err != nil { err = fmt.Errorf("failed to add shard %s: %w", shard, err) - // [gtypes.BlobUploadError]s are non-fatal. - var errBlobUpload gtypes.BlobUploadError + // [types.BlobUploadError]s are non-fatal. + var errBlobUpload types.BlobUploadError if errors.As(err, &errBlobUpload) { - return err, nil + return worker.NewNonFatalError(err) } log.Errorf("%v", err) - return nil, err + return worker.NewFatalError(err) } - return nil, nil + return nil }, }) } return tasks, nil } -func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { +func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) { ctx, span := tracer.Start(ctx, "find-index-add-tasks-for-upload") defer span.End() @@ -116,29 +117,29 @@ func (a API) FindIndexAddTasksForUpload(ctx context.Context, uploadID id.UploadI } span.AddEvent("found closed indexes", trace.WithAttributes(attribute.Int("indexes", len(closedIndexes)))) - tasks := make([]gtypes.IDTask, 0, len(closedIndexes)) + tasks := make([]types.IDTask, 0, len(closedIndexes)) for _, index := range closedIndexes { - tasks = append(tasks, gtypes.IDTask{ + tasks = append(tasks, types.IDTask{ ID: index.ID(), - Run: func(ctx context.Context) (error, error) { + Run: func(ctx context.Context) worker.TaskError { if err := a.addBlob(ctx, index, spaceDID); err != nil { err = fmt.Errorf("failed to add index %s: %w", index, err) - // [gtypes.BlobUploadError]s are non-fatal. - var errBlobUpload gtypes.BlobUploadError + // [types.BlobUploadError]s are non-fatal. + var errBlobUpload types.BlobUploadError if errors.As(err, &errBlobUpload) { - return err, nil + return worker.NewNonFatalError(err) } log.Errorf("%v", err) - return nil, err + return worker.NewFatalError(err) } - return nil, nil + return nil }, }) } return tasks, nil } -func (a API) FindShardPostProcessTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { +func (a API) FindShardPostProcessTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) { ctx, span := tracer.Start(ctx, "find-shard-post-process-tasks-for-upload") defer span.End() @@ -148,11 +149,11 @@ func (a API) FindShardPostProcessTasksForUpload(ctx context.Context, uploadID id } span.AddEvent("found uploaded shards", trace.WithAttributes(attribute.Int("shards", len(uploadedShards)))) - tasks := make([]gtypes.IDTask, 0, len(uploadedShards)) + tasks := make([]types.IDTask, 0, len(uploadedShards)) for _, shard := range uploadedShards { - tasks = append(tasks, gtypes.IDTask{ + tasks = append(tasks, types.IDTask{ ID: shard.ID(), - Run: func(ctx context.Context) (error, error) { + Run: func(ctx context.Context) worker.TaskError { err := a.postProcessBlob(ctx, shard, spaceDID, func(blob model.Blob) error { var opts []client.FilecoinOfferOption if blob.PDPAccept() != nil { @@ -165,10 +166,10 @@ func (a API) FindShardPostProcessTasksForUpload(ctx context.Context, uploadID id }) if err != nil { log.Errorf("failed to post-process shard %s: %v", shard, err) - return nil, fmt.Errorf("failed to post-process shard %s: %w", shard, err) + return worker.NewFatalError(fmt.Errorf("failed to post-process shard %s: %w", shard, err)) } log.Infof("Successfully post-processed shard %s", shard.ID()) - return nil, nil + return nil }, }) } @@ -246,7 +247,7 @@ func (a API) addBlob(ctx context.Context, blob model.Blob, spaceDID did.DID) err })) addedBlob, err := a.spaceBlobAdd(ctx, addReader, spaceDID, opts...) if err != nil { - return gtypes.NewBlobUploadError(blob.ID(), fmt.Errorf("failed to add blob %s to space %s: %w", blob, spaceDID, err)) + return types.NewBlobUploadError(blob.ID(), fmt.Errorf("failed to add blob %s to space %s: %w", blob, spaceDID, err)) } if err := blob.SpaceBlobAdded(addedBlob); err != nil { @@ -271,7 +272,7 @@ func (a API) addBlob(ctx context.Context, blob model.Blob, spaceDID did.DID) err func (a API) postProcessBlob(ctx context.Context, blob model.Blob, spaceDID did.DID, afterAdded func(blob model.Blob) error) error { if err := a.spaceBlobReplicate(ctx, blob, spaceDID, blob.Location()); err != nil { - return gtypes.NewBlobUploadError(blob.ID(), fmt.Errorf("failed to replicate blob %s: %w", blob, err)) + return types.NewBlobUploadError(blob.ID(), fmt.Errorf("failed to replicate blob %s: %w", blob, err)) } if afterAdded != nil { @@ -310,7 +311,7 @@ func (a API) spaceBlobReplicate(ctx context.Context, blob model.Blob, spaceDID d _, _, err := a.Client.SpaceBlobReplicate( ctx, spaceDID, - types.Blob{ + captypes.Blob{ Digest: blob.Digest(), Size: blob.Size(), }, @@ -328,8 +329,8 @@ func (a API) filecoinOffer(ctx context.Context, blob model.Blob, spaceDID did.DI switch { case blob.Size() == 0: return fmt.Errorf("blob %s has no set size yet", blob) - case blob.Size() < gtypes.MinPiecePayload: - log.Warnf("skipping `filecoin/offer` for blob %s: size %d is below minimum %d", blob, blob.Size(), gtypes.MinPiecePayload) + case blob.Size() < types.MinPiecePayload: + log.Warnf("skipping `filecoin/offer` for blob %s: size %d is below minimum %d", blob, blob.Size(), types.MinPiecePayload) return nil case blob.Size() > commp.MaxPiecePayload: log.Warnf("skipping `filecoin/offer` for blob %s: size %d is above maximum %d", blob, blob.Size(), commp.MaxPiecePayload) @@ -348,7 +349,7 @@ func (a API) filecoinOffer(ctx context.Context, blob model.Blob, spaceDID did.DI return nil } -func (a API) FindIndexPostProcessTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]gtypes.IDTask, error) { +func (a API) FindIndexPostProcessTasksForUpload(ctx context.Context, uploadID id.UploadID, spaceDID did.DID) ([]types.IDTask, error) { ctx, span := tracer.Start(ctx, "find-index-post-process-tasks-for-upload") defer span.End() @@ -358,11 +359,11 @@ func (a API) FindIndexPostProcessTasksForUpload(ctx context.Context, uploadID id } span.AddEvent("found uploaded indexes", trace.WithAttributes(attribute.Int("indexes", len(uploadedIndexes)))) - tasks := make([]gtypes.IDTask, 0, len(uploadedIndexes)) + tasks := make([]types.IDTask, 0, len(uploadedIndexes)) for _, index := range uploadedIndexes { - tasks = append(tasks, gtypes.IDTask{ + tasks = append(tasks, types.IDTask{ ID: index.ID(), - Run: func(ctx context.Context) (error, error) { + Run: func(ctx context.Context) worker.TaskError { err := a.postProcessBlob(ctx, index, spaceDID, func(blob model.Blob) error { // Use a placeholder for the root because it doesn't matter what it is, // and we don't want to wait for it to be known. It shouldn't really be @@ -371,10 +372,10 @@ func (a API) FindIndexPostProcessTasksForUpload(ctx context.Context, uploadID id }) if err != nil { log.Errorf("failed to post-process index %s: %v", index, err) - return nil, fmt.Errorf("failed to post-process index %s: %w", index, err) + return worker.NewFatalError(fmt.Errorf("failed to post-process index %s: %w", index, err)) } log.Infof("Successfully post-processed index %s", index.ID()) - return nil, nil + return nil }, }) } diff --git a/pkg/preparation/storacha/storacha_test.go b/pkg/preparation/storacha/storacha_test.go index 0cca97b3..9e8c7ed6 100644 --- a/pkg/preparation/storacha/storacha_test.go +++ b/pkg/preparation/storacha/storacha_test.go @@ -85,9 +85,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload shards @@ -108,9 +106,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { ppTasks, err := api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload shards firstShard, err = repo.GetShardByID(t.Context(), firstShard.ID()) @@ -143,9 +139,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload second shard @@ -162,9 +156,7 @@ func TestFindShardAddTasksForUpload(t *testing.T) { ppTasks, err = api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload second shard @@ -224,16 +216,16 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, err) - require.ErrorContains(t, nonFatal, "simulated SpaceBlobAdd error") + te := task.Run(t.Context()) + require.NotNil(t, te) + require.False(t, te.IsFatal(), "BlobUploadError should be non-fatal") + require.Len(t, te.NonFatalErrors(), 1) + require.ErrorContains(t, te.NonFatalErrors()[0], "simulated SpaceBlobAdd error") } ppTasks, err := api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // It should have `space/blob/add`ed (and failed)... @@ -256,16 +248,15 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } ppTasks, err = api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, fatalErr := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.ErrorContains(t, fatalErr, "simulated SpaceBlobReplicate error") + te := task.Run(t.Context()) + require.NotNil(t, te) + require.True(t, te.IsFatal(), "post-process error should be fatal") + require.ErrorContains(t, te.FatalError(), "simulated SpaceBlobReplicate error") } // It should have `space/blob/add`ed again... @@ -284,16 +275,15 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err = api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } ppTasks, err = api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, fatalErr := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.ErrorContains(t, fatalErr, "simulated FilecoinOffer error") + te := task.Run(t.Context()) + require.NotNil(t, te) + require.True(t, te.IsFatal(), "post-process error should be fatal") + require.ErrorContains(t, te.FatalError(), "simulated FilecoinOffer error") } // It should NOT `space/blob/add` again... @@ -343,16 +333,12 @@ func TestFindShardAddTasksForUpload(t *testing.T) { tasks, err := api.FindShardAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } ppTasks, err := api.FindShardPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // It should `space/blob/add`... @@ -427,9 +413,7 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { tasks, err := api.FindIndexAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload first shard @@ -447,9 +431,7 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { ppTasks, err := api.FindIndexPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload first shard @@ -485,9 +467,7 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { tasks, err = api.FindIndexAddTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range tasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload second shard @@ -504,9 +484,7 @@ func TestFindIndexAddTasksForUpload(t *testing.T) { ppTasks, err = api.FindIndexPostProcessTasksForUpload(t.Context(), upload.ID(), spaceDID) require.NoError(t, err) for _, task := range ppTasks { - nonFatal, err := task.Run(t.Context()) - require.NoError(t, nonFatal) - require.NoError(t, err) + require.Nil(t, task.Run(t.Context())) } // Reload second shard diff --git a/pkg/preparation/types/errors.go b/pkg/preparation/types/errors.go index 5a3573c0..4e10dab9 100644 --- a/pkg/preparation/types/errors.go +++ b/pkg/preparation/types/errors.go @@ -1,7 +1,6 @@ package types import ( - "context" "errors" "fmt" "strings" @@ -192,12 +191,3 @@ func (e BlobUploadErrors) Error() string { func (e BlobUploadErrors) Unwrap() []error { return e.errs } - -// IDTask represents a task that can be identified and deduplicated by an -// [id.ID]. The task's Run function returns two errors: a non-fatal error that -// can be collected and reported after all tasks have completed, and a fatal -// error that should cause immediate cancellation of all other tasks. -type IDTask struct { - ID id.ID - Run func(context.Context) (error, error) -} diff --git a/pkg/preparation/types/idtask.go b/pkg/preparation/types/idtask.go new file mode 100644 index 00000000..687bbdc3 --- /dev/null +++ b/pkg/preparation/types/idtask.go @@ -0,0 +1,13 @@ +package types + +import ( + "github.com/storacha/guppy/pkg/preparation/internal/worker" + "github.com/storacha/guppy/pkg/preparation/types/id" +) + +// IDTask represents a worker task that can be identified and deduplicated by an +// [id.ID]. +type IDTask struct { + ID id.ID + Run worker.Task +} diff --git a/pkg/preparation/uploads/uploads.go b/pkg/preparation/uploads/uploads.go index c1e76e7c..9bb3dad9 100644 --- a/pkg/preparation/uploads/uploads.go +++ b/pkg/preparation/uploads/uploads.go @@ -18,6 +18,7 @@ import ( "github.com/storacha/guppy/pkg/preparation/bettererrgroup" blobsmodel "github.com/storacha/guppy/pkg/preparation/blobs/model" dagmodel "github.com/storacha/guppy/pkg/preparation/dags/model" + "github.com/storacha/guppy/pkg/preparation/internal/worker" scanmodel "github.com/storacha/guppy/pkg/preparation/scans/model" "github.com/storacha/guppy/pkg/preparation/types" "github.com/storacha/guppy/pkg/preparation/types/id" @@ -374,23 +375,23 @@ func runScanWorker( span.End() }() - _, err = Worker( + if te := worker.Run( ctx, scansAvailable, 1, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { - return []func(context.Context) (error, error){ - func(ctx context.Context) (error, error) { + func(ctx context.Context) ([]worker.Task, error) { + return []worker.Task{ + func(ctx context.Context) worker.TaskError { if api.AssumeUnchangedSources { upload, err := api.Repo.GetUploadByID(ctx, uploadID) if err != nil { - return nil, fmt.Errorf("checking upload for existing scan: %w", err) + return worker.NewFatalError(fmt.Errorf("checking upload for existing scan: %w", err)) } if upload.HasRootFSEntryID() { log.Infow("Skipping FS rescan (--assume-unchanged-sources): scan already exists", "upload", uploadID) - return nil, nil + return nil } log.Infow("No existing scan found, performing FS scan despite --assume-unchanged-sources", "upload", uploadID) } @@ -406,10 +407,10 @@ func runScanWorker( }) if err != nil { - return nil, fmt.Errorf("running scans: %w", err) + return worker.NewFatalError(fmt.Errorf("running scans: %w", err)) } - return nil, nil + return nil }, }, nil }, @@ -419,7 +420,9 @@ func runScanWorker( close(dagScansAvailable) return nil }, - ) + ); te != nil { + err = te + } return err } @@ -456,25 +459,25 @@ func runDAGScanWorker( span.End() }() - _, err = Worker( + if te := worker.Run( ctx, dagScansAvailable, 1, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { - return []func(context.Context) (error, error){ - func(ctx context.Context) (error, error) { + func(ctx context.Context) ([]worker.Task, error) { + return []worker.Task{ + func(ctx context.Context) worker.TaskError { err := api.ExecuteDagScansForUpload(ctx, uploadID, func(node dagmodel.Node, data []byte) error { signal(nodeUploadsAvailable) return nil }) if err != nil { - return nil, fmt.Errorf("running dag scans for upload %s: %w", uploadID, err) + return worker.NewFatalError(fmt.Errorf("running dag scans for upload %s: %w", uploadID, err)) } - return nil, nil + return nil }, }, nil }, @@ -501,7 +504,9 @@ func runDAGScanWorker( close(nodeUploadsAvailable) return nil }, - ) + ); te != nil { + err = te + } return err } @@ -542,20 +547,20 @@ func runShardingWorker( return nil } - _, err = Worker( + if te := worker.Run( ctx, nodeUploadsAvailable, 1, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { - return []func(context.Context) (error, error){ - func(ctx context.Context) (error, error) { + func(ctx context.Context) ([]worker.Task, error) { + return []worker.Task{ + func(ctx context.Context) worker.TaskError { err := api.AddNodesToUploadShards(ctx, uploadID, spaceDID, handleClosedShard) if err != nil { - return nil, fmt.Errorf("adding nodes to shards for upload %s: %w", uploadID, err) + return worker.NewFatalError(fmt.Errorf("adding nodes to shards for upload %s: %w", uploadID, err)) } - return nil, nil + return nil }, }, nil }, @@ -571,7 +576,9 @@ func runShardingWorker( return nil }, - ) + ); te != nil { + err = te + } return err } @@ -611,20 +618,20 @@ func runIndexingWorker( return nil } - _, err = Worker( + if te := worker.Run( ctx, shardsNeedIndexing, 1, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { - return []func(context.Context) (error, error){ - func(ctx context.Context) (error, error) { + func(ctx context.Context) ([]worker.Task, error) { + return []worker.Task{ + func(ctx context.Context) worker.TaskError { err := api.AddShardsToUploadIndexes(ctx, uploadID, handleClosedIndex) if err != nil { - return nil, fmt.Errorf("adding shards to indexes for upload %s: %w", uploadID, err) + return worker.NewFatalError(fmt.Errorf("adding shards to indexes for upload %s: %w", uploadID, err)) } - return nil, nil + return nil }, }, nil }, @@ -640,7 +647,9 @@ func runIndexingWorker( return nil }, - ) + ); te != nil { + err = te + } return err } @@ -678,30 +687,30 @@ func runShardUploadWorker( var inFlightShards sync.Map - nonFatals, err := Worker( + te := worker.Run( ctx, closedShardsAvailable, api.BlobUploadParallelism, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + func(ctx context.Context) ([]worker.Task, error) { rawTasks, err := api.FindShardAddTasksForUpload(ctx, uploadID, spaceDID) if err != nil { return nil, err } - var tasks []func(context.Context) (error, error) + var tasks []worker.Task for _, raw := range rawTasks { // Ignore tasks that are already in flight. if _, loaded := inFlightShards.LoadOrStore(raw.ID, struct{}{}); loaded { continue } - tasks = append(tasks, func(ctx context.Context) (error, error) { + tasks = append(tasks, func(ctx context.Context) worker.TaskError { defer inFlightShards.Delete(raw.ID) - nonFatal, fatal := raw.Run(ctx) - if nonFatal == nil && fatal == nil { + taskErr := raw.Run(ctx) + if taskErr == nil { signal(uploadedShardsAvailable) } - return nonFatal, fatal + return taskErr }) } return tasks, nil @@ -714,7 +723,14 @@ func runShardUploadWorker( }, ) - return errors.Join(err, types.NewBlobUploadErrors(nonFatals)) + var nonFatals []error + var fatal error + if te != nil { + nonFatals = te.NonFatalErrors() + fatal = te.FatalError() + } + err = errors.Join(fatal, types.NewBlobUploadErrors(nonFatals)) + return err } func runPostProcessShardWorker( @@ -749,23 +765,23 @@ func runPostProcessShardWorker( var inFlightShards sync.Map - _, err = Worker( + if te := worker.Run( ctx, uploadedShardsAvailable, api.BlobUploadParallelism, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + func(ctx context.Context) ([]worker.Task, error) { rawTasks, err := api.FindShardPostProcessTasksForUpload(ctx, uploadID, spaceDID) if err != nil { return nil, err } - var tasks []func(context.Context) (error, error) + var tasks []worker.Task for _, raw := range rawTasks { if _, loaded := inFlightShards.LoadOrStore(raw.ID, struct{}{}); loaded { continue } - tasks = append(tasks, func(ctx context.Context) (error, error) { + tasks = append(tasks, func(ctx context.Context) worker.TaskError { defer inFlightShards.Delete(raw.ID) return raw.Run(ctx) }) @@ -781,7 +797,9 @@ func runPostProcessShardWorker( } return nil }, - ) + ); te != nil { + err = te + } return err } @@ -819,30 +837,30 @@ func runIndexUploadWorker( var inFlightIndexes sync.Map - nonFatals, err := Worker( + te := worker.Run( ctx, closedIndexesAvailable, api.BlobUploadParallelism, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + func(ctx context.Context) ([]worker.Task, error) { rawTasks, err := api.FindIndexAddTasksForUpload(ctx, uploadID, spaceDID) if err != nil { return nil, err } - var tasks []func(context.Context) (error, error) + var tasks []worker.Task for _, raw := range rawTasks { // Ignore tasks that are already in flight. if _, loaded := inFlightIndexes.LoadOrStore(raw.ID, struct{}{}); loaded { continue } - tasks = append(tasks, func(ctx context.Context) (error, error) { + tasks = append(tasks, func(ctx context.Context) worker.TaskError { defer inFlightIndexes.Delete(raw.ID) - nonFatal, fatal := raw.Run(ctx) - if nonFatal == nil && fatal == nil { + taskErr := raw.Run(ctx) + if taskErr == nil { signal(uploadedIndexesAvailable) } - return nonFatal, fatal + return taskErr }) } return tasks, nil @@ -855,7 +873,14 @@ func runIndexUploadWorker( }, ) - return errors.Join(err, types.NewBlobUploadErrors(nonFatals)) + var nonFatals []error + var fatal error + if te != nil { + nonFatals = te.NonFatalErrors() + fatal = te.FatalError() + } + err = errors.Join(fatal, types.NewBlobUploadErrors(nonFatals)) + return err } func runPostProcessIndexWorker( @@ -890,23 +915,23 @@ func runPostProcessIndexWorker( var inFlightIndexes sync.Map - _, err = Worker( + if te := worker.Run( ctx, uploadedIndexesAvailable, api.BlobUploadParallelism, // findWork - func(ctx context.Context) ([]func(context.Context) (error, error), error) { + func(ctx context.Context) ([]worker.Task, error) { rawTasks, err := api.FindIndexPostProcessTasksForUpload(ctx, uploadID, spaceDID) if err != nil { return nil, err } - var tasks []func(context.Context) (error, error) + var tasks []worker.Task for _, raw := range rawTasks { if _, loaded := inFlightIndexes.LoadOrStore(raw.ID, struct{}{}); loaded { continue } - tasks = append(tasks, func(ctx context.Context) (error, error) { + tasks = append(tasks, func(ctx context.Context) worker.TaskError { defer inFlightIndexes.Delete(raw.ID) return raw.Run(ctx) }) @@ -916,6 +941,8 @@ func runPostProcessIndexWorker( // finalize nil, - ) + ); te != nil { + err = te + } return err } diff --git a/pkg/preparation/uploads/worker.go b/pkg/preparation/uploads/worker.go deleted file mode 100644 index d94a6bd0..00000000 --- a/pkg/preparation/uploads/worker.go +++ /dev/null @@ -1,101 +0,0 @@ -package uploads - -import ( - "context" - "fmt" - "sync" - - "github.com/storacha/guppy/internal/ctxutil" - "golang.org/x/sync/errgroup" -) - -func Worker( - ctx context.Context, - workAvailable <-chan struct{}, - parallelism int, - findWork func(ctx context.Context) ([]func(context.Context) (error, error), error), - finalize func() error, -) ([]error, error) { - var ( - queue []func(context.Context) (error, error) - nonFatalErrorsMu sync.Mutex - nonFatalErrors []error - ) - - // gctx is cancelled when any task returns a fatal error, allowing the outer - // loop to detect failure and stop dispatching. Tasks receive the outer ctx - // (not gctx) so sibling tasks are not cancelled when one fails. - sem := make(chan struct{}, parallelism) - eg, gctx := errgroup.WithContext(ctx) - - dispatchNext := func() bool { - if len(queue) == 0 { - return false - } - task := queue[0] - queue = queue[1:] - select { - case sem <- struct{}{}: - case <-gctx.Done(): - return false - } - eg.Go(func() error { - defer func() { <-sem }() - nonFatal, fatal := task(ctx) - if fatal != nil { - return fmt.Errorf("worker task encountered a fatal error: %w", fatal) - } - if nonFatal != nil { - nonFatalErrorsMu.Lock() - nonFatalErrors = append(nonFatalErrors, nonFatal) - nonFatalErrorsMu.Unlock() - } - return nil - }) - return true - } - - for { - select { - case <-gctx.Done(): - // gctx is cancelled either because ctx was cancelled (external stop) - // or because a task returned a fatal error (internal failure). Wait - // for all in-flight tasks to finish before determining which it was. - fatalErr := eg.Wait() - if fatalErr != nil { - return nonFatalErrors, fatalErr - } - // No task error — must be an external cancellation. - return nonFatalErrors, ctxutil.Cause(ctx) - case _, ok := <-workAvailable: - if !ok { - // Drain the queue before finalizing. - for dispatchNext() { - } - if fatalErr := eg.Wait(); fatalErr != nil { - return nonFatalErrors, fatalErr - } - if finalize != nil { - if err := finalize(); err != nil { - return nonFatalErrors, fmt.Errorf("worker finalize encountered an error: %w", err) - } - } - if len(nonFatalErrors) > 0 { - return nonFatalErrors, nil - } - return nil, nil - } - - tasks, err := findWork(ctx) - if err != nil { - _ = eg.Wait() - return nonFatalErrors, fmt.Errorf("worker findWork encountered an error: %w", err) - } - queue = append(queue, tasks...) - - // Fill available parallelism slots. - for dispatchNext() { - } - } - } -}