Files
bdrclone/internal/baidu/upload.go
youbin c602698ce9
All checks were successful
Build / Test and build (push) Successful in 5m53s
Retry transient uploads and continue backups
2026-08-15 15:32:43 +08:00

328 lines
9.3 KiB
Go

package baidu
import (
"bytes"
"context"
"crypto/md5"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"hash"
"io"
"mime/multipart"
"net/http"
"net/url"
"os"
"path/filepath"
"strconv"
"sync"
"time"
)
const (
defaultPartSize = int64(4 << 20)
vipPartSize = int64(16 << 20)
svipPartSize = int64(32 << 20)
maxPartCount = 2048
maxUploadAttempts = 6
)
type UploadProgress func(uploaded, total int64)
type uploadPlan struct {
Size int64
PartSize int64
PartMD5 []string
ContentMD5 string
SliceMD5 string
}
type precreateResponse struct {
ReturnType int `json:"return_type"`
UploadID string `json:"uploadid"`
BlockList []int `json:"block_list"`
Info File `json:"info"`
}
func (c *Client) UploadFile(ctx context.Context, localPath, remotePath string, progress UploadProgress) (File, error) {
f, err := os.Open(localPath)
if err != nil {
return File{}, fmt.Errorf("open local file: %w", err)
}
defer f.Close()
stat, err := f.Stat()
if err != nil {
return File{}, fmt.Errorf("stat local file: %w", err)
}
return c.Upload(ctx, f, stat.Size(), remotePath, stat.ModTime(), progress)
}
func (c *Client) Upload(ctx context.Context, file *os.File, size int64, remotePath string, modTime time.Time, progress UploadProgress) (File, error) {
if size == 0 {
return File{}, errors.New("Baidu Netdisk API does not allow empty files")
}
partSize, err := c.choosePartSize(ctx)
if err != nil {
return File{}, err
}
plan, err := buildUploadPlan(ctx, file, size, partSize)
if err != nil {
return File{}, err
}
if len(plan.PartMD5) > maxPartCount {
return File{}, fmt.Errorf("file needs %d parts; Baidu allows at most %d with the selected %d MiB part size", len(plan.PartMD5), maxPartCount, partSize>>20)
}
remote := c.RemotePath(remotePath)
pre, err := c.precreate(ctx, remote, plan, modTime, true)
if err != nil {
return File{}, err
}
if pre.ReturnType == 2 {
return pre.Info, nil
}
if pre.UploadID == "" {
return File{}, errors.New("Baidu precreate response contains no uploadid")
}
parts := pre.BlockList
if len(parts) == 0 {
parts = make([]int, len(plan.PartMD5))
for i := range parts {
parts[i] = i
}
}
if err := c.uploadPartsParallel(ctx, file, remote, pre.UploadID, plan, parts, progress); err != nil {
return File{}, err
}
return c.createUploadedFile(ctx, remote, pre.UploadID, plan, modTime)
}
func (c *Client) choosePartSize(ctx context.Context) (int64, error) {
if c.partSize > 0 {
return c.partSize, nil
}
var info struct {
VIPType int `json:"vip_type"`
}
err := c.request(ctx, http.MethodGet, c.apiBase+"/rest/2.0/xpan/nas", url.Values{"method": {"uinfo"}}, nil, &info)
if err != nil {
return 0, fmt.Errorf("query Baidu membership for upload part size: %w", err)
}
switch info.VIPType {
case 1:
return vipPartSize, nil
case 2:
return svipPartSize, nil
default:
return defaultPartSize, nil
}
}
func buildUploadPlan(ctx context.Context, file *os.File, size, partSize int64) (uploadPlan, error) {
if partSize < defaultPartSize {
return uploadPlan{}, errors.New("part size must be at least 4 MiB")
}
if _, err := file.Seek(0, io.SeekStart); err != nil {
return uploadPlan{}, fmt.Errorf("seek upload source: %w", err)
}
fullHash := md5.New()
firstHash := md5.New()
var firstRemaining int64 = 256 << 10
partHashes := make([]string, 0, (size+partSize-1)/partSize)
buf := make([]byte, 1<<20)
for offset := int64(0); offset < size; {
if err := ctx.Err(); err != nil {
return uploadPlan{}, err
}
partBytes := min(partSize, size-offset)
partHash := md5.New()
remaining := partBytes
for remaining > 0 {
n, readErr := file.Read(buf[:min(int64(len(buf)), remaining)])
if n > 0 {
chunk := buf[:n]
_, _ = fullHash.Write(chunk)
_, _ = partHash.Write(chunk)
if firstRemaining > 0 {
firstN := min(int64(n), firstRemaining)
_, _ = firstHash.Write(chunk[:firstN])
firstRemaining -= firstN
}
remaining -= int64(n)
}
if readErr != nil {
if errors.Is(readErr, io.EOF) && remaining == 0 {
break
}
return uploadPlan{}, fmt.Errorf("hash upload source: %w", readErr)
}
}
partHashes = append(partHashes, hashString(partHash))
offset += partBytes
}
return uploadPlan{Size: size, PartSize: partSize, PartMD5: partHashes, ContentMD5: hashString(fullHash), SliceMD5: hashString(firstHash)}, nil
}
func hashString(h hash.Hash) string { return hex.EncodeToString(h.Sum(nil)) }
func (c *Client) precreate(ctx context.Context, remote string, plan uploadPlan, modTime time.Time, rapid bool) (precreateResponse, error) {
blocks, _ := json.Marshal(plan.PartMD5)
form := url.Values{
"path": {remote}, "size": {intString(plan.Size)}, "isdir": {"0"}, "autoinit": {"1"},
"rtype": {"3"}, "block_list": {string(blocks)},
"local_ctime": {intString(modTime.Unix())}, "local_mtime": {intString(modTime.Unix())},
}
if rapid {
form.Set("content-md5", plan.ContentMD5)
form.Set("slice-md5", plan.SliceMD5)
}
var result precreateResponse
err := c.request(ctx, http.MethodPost, c.xpanFile(), url.Values{"method": {"precreate"}}, form, &result)
return result, err
}
func (c *Client) uploadPartsParallel(ctx context.Context, file *os.File, remote, uploadID string, plan uploadPlan, parts []int, progress UploadProgress) error {
workers := c.uploadParts
if workers < 1 {
workers = 1
}
ctx, cancel := context.WithCancel(ctx)
defer cancel()
jobs := make(chan int)
errCh := make(chan error, 1)
var wg sync.WaitGroup
var mu sync.Mutex
var uploaded int64
worker := func() {
defer wg.Done()
for part := range jobs {
offset := int64(part) * plan.PartSize
if offset >= plan.Size {
select {
case errCh <- fmt.Errorf("Baidu requested invalid part %d", part):
default:
}
cancel()
return
}
size := min(plan.PartSize, plan.Size-offset)
if err := c.uploadPartWithRetry(ctx, file, remote, uploadID, part, offset, size); err != nil {
select {
case errCh <- err:
default:
}
cancel()
return
}
mu.Lock()
uploaded += size
if progress != nil {
progress(uploaded, plan.Size)
}
mu.Unlock()
}
}
for range min(workers, len(parts)) {
wg.Add(1)
go worker()
}
for _, part := range parts {
select {
case jobs <- part:
case <-ctx.Done():
break
}
if ctx.Err() != nil {
break
}
}
close(jobs)
wg.Wait()
select {
case err := <-errCh:
return err
default:
return ctx.Err()
}
}
func (c *Client) uploadPartWithRetry(ctx context.Context, file *os.File, remote, uploadID string, part int, offset, size int64) error {
var lastErr error
for attempt := 0; attempt < maxUploadAttempts; attempt++ {
if err := ctx.Err(); err != nil {
return err
}
section := io.NewSectionReader(file, offset, size)
lastErr = c.uploadPart(ctx, section, filepath.Base(remote), remote, uploadID, part)
if lastErr == nil {
return nil
}
if attempt < maxUploadAttempts-1 {
if err := sleepContext(ctx, c.uploadRetryDelay(attempt)); err != nil {
return err
}
}
}
return fmt.Errorf("upload failed after %d attempts: %w", maxUploadAttempts, lastErr)
}
func defaultUploadRetryDelay(attempt int) time.Duration {
delay := time.Second * time.Duration(1<<attempt)
return min(delay, 30*time.Second)
}
func (c *Client) uploadPart(ctx context.Context, section *io.SectionReader, filename, remote, uploadID string, part int) error {
var envelope bytes.Buffer
mw := multipart.NewWriter(&envelope)
if _, err := mw.CreateFormFile("file", filename); err != nil {
return err
}
headerLen := envelope.Len()
if err := mw.Close(); err != nil {
return err
}
header := append([]byte(nil), envelope.Bytes()[:headerLen]...)
tail := append([]byte(nil), envelope.Bytes()[headerLen:]...)
q := url.Values{
"method": {"upload"}, "access_token": {c.token()}, "type": {"tmpfile"},
"path": {remote}, "uploadid": {uploadID}, "partseq": {strconv.Itoa(part)},
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.uploadBase+"/rest/2.0/pcs/superfile2?"+q.Encode(), io.MultiReader(bytes.NewReader(header), section, bytes.NewReader(tail)))
if err != nil {
return err
}
req.ContentLength = int64(len(header)+len(tail)) + section.Size()
req.Header.Set("Content-Type", mw.FormDataContentType())
req.Header.Set("User-Agent", c.userAgent)
resp, err := c.httpClient.Do(req)
if err != nil {
return fmt.Errorf("upload part %d: %w", part, err)
}
defer resp.Body.Close()
b, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return fmt.Errorf("read upload part %d response: %w", part, err)
}
var apiErr APIError
_ = json.Unmarshal(b, &apiErr)
if resp.StatusCode/100 != 2 || apiErr.Code() != 0 {
if apiErr.Code() != 0 {
return fmt.Errorf("upload part %d: %w", part, &apiErr)
}
return fmt.Errorf("upload part %d: HTTP %s: %s", part, resp.Status, string(b))
}
return nil
}
func (c *Client) createUploadedFile(ctx context.Context, remote, uploadID string, plan uploadPlan, modTime time.Time) (File, error) {
blocks, _ := json.Marshal(plan.PartMD5)
form := url.Values{
"path": {remote}, "size": {intString(plan.Size)}, "isdir": {"0"}, "rtype": {"3"},
"uploadid": {uploadID}, "block_list": {string(blocks)},
"local_ctime": {intString(modTime.Unix())}, "local_mtime": {intString(modTime.Unix())},
}
var result File
err := c.request(ctx, http.MethodPost, c.xpanFile(), url.Values{"method": {"create"}}, form, &result)
return result, err
}