diff --git a/file_upload.go b/file_upload.go index be3d4d3..2c4cf9b 100644 --- a/file_upload.go +++ b/file_upload.go @@ -19,6 +19,16 @@ import ( "github.com/rclone/go-proton-api" ) +func collectUploadErrors(errChan <-chan error, count int) error { + var firstErr error + for range count { + if err := <-errChan; err != nil && firstErr == nil { + firstErr = err + } + } + return firstErr +} + func (protonDrive *ProtonDrive) handleRevisionConflict(ctx context.Context, link *proton.Link, createFileResp *proton.CreateFileRes) (string, bool, error) { if link != nil { linkID := link.LinkID @@ -292,15 +302,13 @@ func (protonDrive *ProtonDrive) uploadAndCollectBlockData(ctx context.Context, n return err } - errChan := make(chan error) + errChan := make(chan error, len(blockUploadResp)) uploadBlockWrapper := func(ctx context.Context, errChan chan error, bareURL, token string, block io.Reader) { - // log.Println("Before semaphore") if err := protonDrive.blockUploadSemaphore.Acquire(ctx, 1); err != nil { errChan <- err + return } defer protonDrive.blockUploadSemaphore.Release(1) - // log.Println("After semaphore") - // defer log.Println("Release semaphore") errChan <- protonDrive.c.UploadBlock(ctx, bareURL, token, block) } @@ -308,11 +316,8 @@ func (protonDrive *ProtonDrive) uploadAndCollectBlockData(ctx context.Context, n go uploadBlockWrapper(ctx, errChan, blockUploadResp[i].BareURL, blockUploadResp[i].Token, bytes.NewReader(pendingUploadBlocks[i].encData)) } - for i := 0; i < len(blockUploadResp); i++ { - err := <-errChan - if err != nil { - return err - } + if err := collectUploadErrors(errChan, len(blockUploadResp)); err != nil { + return err } pendingUploadBlocks = pendingUploadBlocks[:0] diff --git a/file_upload_batch_permits_test.go b/file_upload_batch_permits_test.go new file mode 100644 index 0000000..df740b5 --- /dev/null +++ b/file_upload_batch_permits_test.go @@ -0,0 +1,123 @@ +package proton_api_bridge + +import ( + "bytes" + "context" + "fmt" + "io" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/ProtonMail/gopenpgp/v3/crypto" + "github.com/rclone/go-proton-api" + "golang.org/x/sync/semaphore" +) + +func TestFailedBlockBatchReleasesEverySemaphorePermit(t *testing.T) { + originalBlockSize := UPLOAD_BLOCK_SIZE + originalBatchSize := UPLOAD_BATCH_BLOCK_SIZE + UPLOAD_BLOCK_SIZE = 16 + UPLOAD_BATCH_BLOCK_SIZE = 2 + t.Cleanup(func() { + UPLOAD_BLOCK_SIZE = originalBlockSize + UPLOAD_BATCH_BLOCK_SIZE = originalBatchSize + }) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Date", time.Now().UTC().Format(http.TimeFormat)) + w.Header().Set("Content-Type", "application/json") + + switch r.URL.Path { + case "/drive/blocks": + _, err := io.WriteString(w, fmt.Sprintf( + `{"Code":1000,"UploadLinks":[{"Token":"error-token","BareURL":%q},{"Token":"success-token","BareURL":%q}]}`, + serverURL(r, "/storage/blocks/error"), + serverURL(r, "/storage/blocks/success"), + )) + if err != nil { + t.Error(err) + } + case "/storage/blocks/error": + w.WriteHeader(http.StatusBadGateway) + _, err := io.WriteString(w, `{"Code":0,"Error":"simulated bad gateway"}`) + if err != nil { + t.Error(err) + } + case "/storage/blocks/success": + time.Sleep(100 * time.Millisecond) + _, err := io.WriteString(w, `{"Code":1000}`) + if err != nil { + t.Error(err) + } + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + manager := proton.New( + proton.WithHostURL(server.URL), + proton.WithRetryCount(0), + ) + defer manager.Close() + client := manager.NewClient("", "", "") + defer client.Close() + + pgp := crypto.PGP() + signingKey, err := pgp.KeyGeneration().AddUserId("test", "test@example.com").New().GenerateKey() + if err != nil { + t.Fatal(err) + } + signingKeyRing, err := crypto.NewKeyRing(signingKey) + if err != nil { + t.Fatal(err) + } + nodeKey, err := pgp.KeyGeneration().AddUserId("node", "node@example.com").New().GenerateKey() + if err != nil { + t.Fatal(err) + } + nodeKeyRing, err := crypto.NewKeyRing(nodeKey) + if err != nil { + t.Fatal(err) + } + sessionKey, err := pgp.GenerateSessionKey() + if err != nil { + t.Fatal(err) + } + + blockSemaphore := semaphore.NewWeighted(2) + drive := &ProtonDrive{ + MainShare: &proton.Share{ + ShareMetadata: proton.ShareMetadata{ShareID: "share-id"}, + AddressID: "address-id", + }, + DefaultAddrKR: signingKeyRing, + c: client, + blockUploadSemaphore: blockSemaphore, + } + + _, _, _, _, err = drive.uploadAndCollectBlockData( + context.Background(), + sessionKey, + nodeKeyRing, + bytes.NewReader([]byte("0123456789abcdef0123456789abcdef")), + "link-id", + "revision-id", + ) + if err == nil { + t.Fatal("expected failed block upload") + } + + acquireCtx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := blockSemaphore.Acquire(acquireCtx, 2); err != nil { + t.Fatalf("failed batch leaked a block-upload permit: %v", err) + } + blockSemaphore.Release(2) +} + +func serverURL(r *http.Request, path string) string { + return "http://" + r.Host + path +} diff --git a/file_upload_concurrency_test.go b/file_upload_concurrency_test.go new file mode 100644 index 0000000..598c1ea --- /dev/null +++ b/file_upload_concurrency_test.go @@ -0,0 +1,48 @@ +package proton_api_bridge + +import ( + "context" + "errors" + "testing" + "time" + + "golang.org/x/sync/semaphore" +) + +func TestCollectUploadErrorsReleasesAllWorkersAfterFailure(t *testing.T) { + const ( + batchSize = int64(8) + slotCount = int64(20) + ) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + slots := semaphore.NewWeighted(slotCount) + + for batch := 0; batch < 4; batch++ { + results := make(chan error) + for block := int64(0); block < batchSize; block++ { + go func(fail bool) { + if err := slots.Acquire(ctx, 1); err != nil { + results <- err + return + } + defer slots.Release(1) + if fail { + results <- errors.New("synthetic upload failure") + return + } + results <- nil + }(block == 0) + } + + if err := collectUploadErrors(results, int(batchSize)); err == nil { + t.Fatal("expected the first upload failure to be returned") + } + } + + if err := slots.Acquire(ctx, slotCount); err != nil { + t.Fatalf("upload workers leaked semaphore slots: %v", err) + } + slots.Release(slotCount) +}