diff --git a/file_upload.go b/file_upload.go index be3d4d3..938ca17 100644 --- a/file_upload.go +++ b/file_upload.go @@ -8,6 +8,8 @@ import ( "crypto/sha256" "encoding/base64" "encoding/hex" + "errors" + "fmt" "io" "mime" "os" @@ -19,6 +21,156 @@ import ( "github.com/rclone/go-proton-api" ) +const ( + blockUploadMaxAttempts = 5 + blockUploadRetryBaseDelay = time.Second + blockUploadRetryMaxDelay = 15 * time.Second +) + +type pendingUploadBlock struct { + blockUploadInfo proton.BlockUploadInfo + encData []byte +} + +type blockUploadResult struct { + index int + err error +} + +type blockUploadRetryLogger interface { + Warnf(format string, v ...interface{}) +} + +func retryableBlockUploadError(err error) bool { + if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return false + } + + var apiErr *proton.APIError + if errors.As(err, &apiErr) { + return apiErr.Status >= 500 && apiErr.Status <= 599 + } + + var protonNetErr *proton.NetError + return errors.As(err, &protonNetErr) +} + +func blockUploadRetryDelay(failedAttempt int) time.Duration { + delay := blockUploadRetryBaseDelay + for i := 1; i < failedAttempt && delay < blockUploadRetryMaxDelay; i++ { + delay *= 2 + } + if delay > blockUploadRetryMaxDelay { + return blockUploadRetryMaxDelay + } + return delay +} + +func waitForBlockUploadRetry(ctx context.Context, delay time.Duration) error { + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + +func uploadBlockBatchWithRetry( + ctx context.Context, + blocks []pendingUploadBlock, + maxAttempts int, + requestLinks func(context.Context, []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error), + uploadBlock func(context.Context, proton.BlockUploadLink, []byte) error, + wait func(context.Context, time.Duration) error, + logger blockUploadRetryLogger, +) error { + remaining := append([]pendingUploadBlock(nil), blocks...) + var lastErr error + + for attempt := 1; attempt <= maxAttempts; attempt++ { + blockList := make([]proton.BlockUploadInfo, len(remaining)) + for i := range remaining { + blockList[i] = remaining[i].blockUploadInfo + } + + links, err := requestLinks(ctx, blockList) + if err != nil { + lastErr = err + if !retryableBlockUploadError(err) || attempt == maxAttempts { + return err + } + } else { + if len(links) != len(remaining) { + return fmt.Errorf( + "requested %d Proton block upload links, received %d", + len(remaining), + len(links), + ) + } + + results := make(chan blockUploadResult, len(remaining)) + for i := range remaining { + go func(index int) { + results <- blockUploadResult{ + index: index, + err: uploadBlock(ctx, links[index], remaining[index].encData), + } + }(i) + } + + errorsByIndex := make([]error, len(remaining)) + for range remaining { + result := <-results + errorsByIndex[result.index] = result.err + } + + failed := make([]pendingUploadBlock, 0, len(remaining)) + var terminalErr error + lastErr = nil + for i, uploadErr := range errorsByIndex { + if uploadErr == nil { + continue + } + if !retryableBlockUploadError(uploadErr) && terminalErr == nil { + terminalErr = uploadErr + } + if lastErr == nil { + lastErr = uploadErr + } + failed = append(failed, remaining[i]) + } + if terminalErr != nil { + return terminalErr + } + if len(failed) == 0 { + return nil + } + if attempt == maxAttempts { + return lastErr + } + remaining = failed + } + + delay := blockUploadRetryDelay(attempt) + if logger != nil { + logger.Warnf( + "Retrying %d transient Proton block upload(s) after %s (attempt %d/%d)", + len(remaining), + delay, + attempt+1, + maxAttempts, + ) + } + if err := wait(ctx, delay); err != nil { + return err + } + } + + return lastErr +} + func (protonDrive *ProtonDrive) handleRevisionConflict(ctx context.Context, link *proton.Link, createFileResp *proton.CreateFileRes) (string, bool, error) { if link != nil { linkID := link.LinkID @@ -257,62 +409,46 @@ func (protonDrive *ProtonDrive) createFileUploadDraft(ctx context.Context, paren } func (protonDrive *ProtonDrive) uploadAndCollectBlockData(ctx context.Context, newSessionKey *crypto.SessionKey, newNodeKR *crypto.KeyRing, file io.Reader, linkID, revisionID string) ([]byte, int64, []int64, string, error) { - type PendingUploadBlocks struct { - blockUploadInfo proton.BlockUploadInfo - encData []byte - } - if newSessionKey == nil || newNodeKR == nil { return nil, 0, nil, "", ErrMissingInputUploadAndCollectBlockData } totalFileSize := int64(0) - pendingUploadBlocks := make([]PendingUploadBlocks, 0) + pendingUploadBlocks := make([]pendingUploadBlock, 0) manifestSignatureData := make([]byte, 0) uploadPendingBlocks := func() error { if len(pendingUploadBlocks) == 0 { return nil } - blockList := make([]proton.BlockUploadInfo, 0) - for i := range pendingUploadBlocks { - blockList = append(blockList, pendingUploadBlocks[i].blockUploadInfo) - } - blockUploadReq := proton.BlockUploadReq{ - AddressID: protonDrive.MainShare.AddressID, - ShareID: protonDrive.MainShare.ShareID, - LinkID: linkID, - RevisionID: revisionID, - - BlockList: blockList, - } - blockUploadResp, err := protonDrive.c.RequestBlockUpload(ctx, blockUploadReq) - if err != nil { - return err + requestLinks := func(ctx context.Context, blockList []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + return protonDrive.c.RequestBlockUpload(ctx, proton.BlockUploadReq{ + AddressID: protonDrive.MainShare.AddressID, + ShareID: protonDrive.MainShare.ShareID, + LinkID: linkID, + RevisionID: revisionID, + BlockList: blockList, + }) } - - errChan := make(chan error) - uploadBlockWrapper := func(ctx context.Context, errChan chan error, bareURL, token string, block io.Reader) { - // log.Println("Before semaphore") + uploadBlock := func(ctx context.Context, link proton.BlockUploadLink, block []byte) error { if err := protonDrive.blockUploadSemaphore.Acquire(ctx, 1); err != nil { - errChan <- err + return err } defer protonDrive.blockUploadSemaphore.Release(1) - // log.Println("After semaphore") - // defer log.Println("Release semaphore") - errChan <- protonDrive.c.UploadBlock(ctx, bareURL, token, block) + return protonDrive.c.UploadBlock(ctx, link.BareURL, link.Token, bytes.NewReader(block)) } - for i := range blockUploadResp { - 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 := uploadBlockBatchWithRetry( + ctx, + pendingUploadBlocks, + blockUploadMaxAttempts, + requestLinks, + uploadBlock, + waitForBlockUploadRetry, + protonDrive.Config.GetLogger(), + ); err != nil { + return err } pendingUploadBlocks = pendingUploadBlocks[:0] @@ -405,7 +541,7 @@ func (protonDrive *ProtonDrive) uploadAndCollectBlockData(ctx context.Context, n } manifestSignatureData = append(manifestSignatureData, hash...) - pendingUploadBlocks = append(pendingUploadBlocks, PendingUploadBlocks{ + pendingUploadBlocks = append(pendingUploadBlocks, pendingUploadBlock{ blockUploadInfo: proton.BlockUploadInfo{ Index: i, // iOS drive: BE starts with 1 Size: int64(len(encData)), diff --git a/file_upload_concurrency_test.go b/file_upload_concurrency_test.go new file mode 100644 index 0000000..06d3677 --- /dev/null +++ b/file_upload_concurrency_test.go @@ -0,0 +1,235 @@ +package proton_api_bridge + +import ( + "context" + "errors" + "reflect" + "sync" + "testing" + "time" + + "github.com/rclone/go-proton-api" + "golang.org/x/sync/semaphore" +) + +func testPendingBlocks(indexes ...int) []pendingUploadBlock { + blocks := make([]pendingUploadBlock, len(indexes)) + for i, index := range indexes { + blocks[i] = pendingUploadBlock{ + blockUploadInfo: proton.BlockUploadInfo{Index: index}, + encData: []byte{byte(index)}, + } + } + return blocks +} + +func testUploadLinks(blocks []proton.BlockUploadInfo) []proton.BlockUploadLink { + links := make([]proton.BlockUploadLink, len(blocks)) + for i, block := range blocks { + links[i] = proton.BlockUploadLink{Token: string(rune(block.Index))} + } + return links +} + +func noRetryWait(_ context.Context, _ time.Duration) error { return nil } + +func TestUploadBlockBatchRetriesOnlyTransientFailures(t *testing.T) { + var requestIndexes [][]int + requestLinks := func(_ context.Context, blocks []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + indexes := make([]int, len(blocks)) + for i, block := range blocks { + indexes[i] = block.Index + } + requestIndexes = append(requestIndexes, indexes) + return testUploadLinks(blocks), nil + } + + var mu sync.Mutex + uploads := map[int]int{} + uploadBlock := func(_ context.Context, _ proton.BlockUploadLink, block []byte) error { + index := int(block[0]) + mu.Lock() + uploads[index]++ + attempt := uploads[index] + mu.Unlock() + if attempt == 1 && index != 2 { + return &proton.APIError{Status: 502, Message: "temporary storage failure"} + } + return nil + } + + err := uploadBlockBatchWithRetry( + context.Background(), + testPendingBlocks(1, 2, 3), + blockUploadMaxAttempts, + requestLinks, + uploadBlock, + noRetryWait, + nil, + ) + if err != nil { + t.Fatalf("retrying transient block uploads failed: %v", err) + } + if want := [][]int{{1, 2, 3}, {1, 3}}; !reflect.DeepEqual(requestIndexes, want) { + t.Fatalf("requested block indexes %v, want %v", requestIndexes, want) + } + if want := map[int]int{1: 2, 2: 1, 3: 2}; !reflect.DeepEqual(uploads, want) { + t.Fatalf("block upload counts %v, want %v", uploads, want) + } +} + +func TestUploadBlockBatchReturnsNonRetryableError(t *testing.T) { + terminalErr := &proton.APIError{Status: 422, Message: "draft conflict"} + requests := 0 + err := uploadBlockBatchWithRetry( + context.Background(), + testPendingBlocks(1), + blockUploadMaxAttempts, + func(_ context.Context, blocks []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + requests++ + return testUploadLinks(blocks), nil + }, + func(_ context.Context, _ proton.BlockUploadLink, _ []byte) error { return terminalErr }, + noRetryWait, + nil, + ) + if !errors.Is(err, terminalErr) { + t.Fatalf("returned error %v, want %v", err, terminalErr) + } + if requests != 1 { + t.Fatalf("requested upload links %d times after a terminal error, want 1", requests) + } +} + +func TestUploadBlockBatchRetriesTransientLinkRequest(t *testing.T) { + transientErr := &proton.APIError{Status: 502, Message: "temporary API failure"} + requests := 0 + uploads := 0 + err := uploadBlockBatchWithRetry( + context.Background(), + testPendingBlocks(1), + blockUploadMaxAttempts, + func(_ context.Context, blocks []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + requests++ + if requests == 1 { + return nil, transientErr + } + return testUploadLinks(blocks), nil + }, + func(_ context.Context, _ proton.BlockUploadLink, _ []byte) error { + uploads++ + return nil + }, + noRetryWait, + nil, + ) + if err != nil { + t.Fatalf("retrying a transient link request failed: %v", err) + } + if requests != 2 || uploads != 1 { + t.Fatalf("observed %d link requests and %d uploads, want 2 and 1", requests, uploads) + } +} + +func TestUploadBlockBatchRejectsMismatchedLinkCount(t *testing.T) { + err := uploadBlockBatchWithRetry( + context.Background(), + testPendingBlocks(1, 2), + blockUploadMaxAttempts, + func(_ context.Context, _ []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + return []proton.BlockUploadLink{{}}, nil + }, + func(_ context.Context, _ proton.BlockUploadLink, _ []byte) error { + t.Fatal("upload must not start with a mismatched link response") + return nil + }, + noRetryWait, + nil, + ) + if err == nil { + t.Fatal("expected a mismatched link count to fail") + } +} + +func TestUploadBlockBatchReturnsLastErrorAfterLimit(t *testing.T) { + transientErr := &proton.APIError{Status: 502, Message: "temporary storage failure"} + requests := 0 + err := uploadBlockBatchWithRetry( + context.Background(), + testPendingBlocks(1), + 3, + func(_ context.Context, blocks []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + requests++ + return testUploadLinks(blocks), nil + }, + func(_ context.Context, _ proton.BlockUploadLink, _ []byte) error { return transientErr }, + noRetryWait, + nil, + ) + if !errors.Is(err, transientErr) { + t.Fatalf("returned error %v, want %v", err, transientErr) + } + if requests != 3 { + t.Fatalf("requested upload links %d times, want 3", requests) + } +} + +func TestUploadBlockBatchHonorsCancellationDuringBackoff(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + err := uploadBlockBatchWithRetry( + ctx, + testPendingBlocks(1), + blockUploadMaxAttempts, + func(_ context.Context, blocks []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + return testUploadLinks(blocks), nil + }, + func(_ context.Context, _ proton.BlockUploadLink, _ []byte) error { + return &proton.APIError{Status: 502, Message: "temporary storage failure"} + }, + waitForBlockUploadRetry, + nil, + ) + if !errors.Is(err, context.Canceled) { + t.Fatalf("returned error %v, want context cancellation", err) + } +} + +func TestUploadBlockBatchReleasesAllWorkersAfterFailure(t *testing.T) { + const slotCount = int64(20) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + slots := semaphore.NewWeighted(slotCount) + + for batch := 0; batch < 4; batch++ { + err := uploadBlockBatchWithRetry( + ctx, + testPendingBlocks(1, 2, 3, 4, 5, 6, 7, 8), + blockUploadMaxAttempts, + func(_ context.Context, blocks []proton.BlockUploadInfo) ([]proton.BlockUploadLink, error) { + return testUploadLinks(blocks), nil + }, + func(_ context.Context, _ proton.BlockUploadLink, block []byte) error { + if err := slots.Acquire(ctx, 1); err != nil { + return err + } + defer slots.Release(1) + if block[0] == 1 { + return errors.New("synthetic upload failure") + } + return nil + }, + noRetryWait, + nil, + ) + if 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) +}