diff --git a/block.go b/block.go index 74e89a6..744272f 100644 --- a/block.go +++ b/block.go @@ -1,8 +1,11 @@ package proton import ( + "bytes" "context" "io" + "mime/multipart" + "net/http" "github.com/go-resty/resty/v2" ) @@ -33,10 +36,33 @@ func (c *Client) RequestBlockUpload(ctx context.Context, req BlockUploadReq) ([] } func (c *Client) UploadBlock(ctx context.Context, bareURL, token string, block io.Reader) error { + var body bytes.Buffer + writer := multipart.NewWriter(&body) + part, err := writer.CreateFormFile("Block", "blob") + if err != nil { + return err + } + if _, err := io.Copy(part, block); err != nil { + return err + } + if err := writer.Close(); err != nil { + return err + } + contentType := writer.FormDataContentType() + payload := body.Bytes() + return c.do(ctx, func(r *resty.Request) (*resty.Response, error) { return r. SetHeader("pm-storage-token", token). - SetMultipartField("Block", "blob", "application/octet-stream", block). + SetHeader("Content-Type", contentType). + SetBody(payload). + AddRetryCondition(func(res *resty.Response, _ error) bool { + if res == nil { + return false + } + return res.StatusCode() >= http.StatusInternalServerError && + res.StatusCode() < 600 + }). Post(bareURL) }) } diff --git a/block_test.go b/block_test.go new file mode 100644 index 0000000..11941d9 --- /dev/null +++ b/block_test.go @@ -0,0 +1,130 @@ +package proton_test + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "github.com/rclone/go-proton-api" + "github.com/stretchr/testify/require" +) + +func TestUploadBlockReplaysMultipartBodyAfterConnectionDrop(t *testing.T) { + payload := []byte("encrypted Proton block payload") + + var ( + mu sync.Mutex + attempts [][]byte + ) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var received []byte + file, _, err := r.FormFile("Block") + if err == nil { + received, err = io.ReadAll(file) + require.NoError(t, err) + require.NoError(t, file.Close()) + } + + mu.Lock() + attempts = append(attempts, received) + attempt := len(attempts) + mu.Unlock() + + if attempt == 1 { + hijacker, ok := w.(http.Hijacker) + require.True(t, ok) + conn, _, err := hijacker.Hijack() + require.NoError(t, err) + require.NoError(t, conn.Close()) + return + } + + w.Header().Set("Date", time.Now().UTC().Format(http.TimeFormat)) + w.Header().Set("Content-Type", "application/json") + _, err = io.WriteString(w, `{"Code":1000}`) + require.NoError(t, err) + })) + defer server.Close() + + manager := proton.New( + proton.WithHostURL(server.URL), + proton.WithRetryCount(1), + ) + defer manager.Close() + client := manager.NewClient("", "", "") + defer client.Close() + + err := client.UploadBlock( + context.Background(), + server.URL+"/storage/blocks", + "test-token", + bytes.NewReader(payload), + ) + require.NoError(t, err) + + mu.Lock() + defer mu.Unlock() + require.Len(t, attempts, 2) + require.Equal(t, payload, attempts[0]) + require.Equal(t, payload, attempts[1]) +} + +func TestUploadBlockRetriesBadGatewayWithSameBody(t *testing.T) { + payload := []byte("encrypted block survives 502 retry") + + var ( + mu sync.Mutex + attempts [][]byte + ) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + file, _, err := r.FormFile("Block") + require.NoError(t, err) + received, err := io.ReadAll(file) + require.NoError(t, err) + require.NoError(t, file.Close()) + + mu.Lock() + attempts = append(attempts, received) + attempt := len(attempts) + mu.Unlock() + + w.Header().Set("Date", time.Now().UTC().Format(http.TimeFormat)) + w.Header().Set("Content-Type", "application/json") + if attempt == 1 { + w.WriteHeader(http.StatusBadGateway) + _, err = io.WriteString(w, `{"Code":0,"Error":"simulated bad gateway"}`) + require.NoError(t, err) + return + } + _, err = io.WriteString(w, `{"Code":1000}`) + require.NoError(t, err) + })) + defer server.Close() + + manager := proton.New( + proton.WithHostURL(server.URL), + proton.WithRetryCount(1), + ) + defer manager.Close() + client := manager.NewClient("", "", "") + defer client.Close() + + err := client.UploadBlock( + context.Background(), + server.URL+"/storage/blocks", + "test-token", + bytes.NewReader(payload), + ) + require.NoError(t, err) + + mu.Lock() + defer mu.Unlock() + require.Len(t, attempts, 2) + require.Equal(t, payload, attempts[0]) + require.Equal(t, payload, attempts[1]) +}