Initial release of bdrclone
This commit is contained in:
321
internal/baidu/upload.go
Normal file
321
internal/baidu/upload.go
Normal file
@@ -0,0 +1,321 @@
|
||||
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
|
||||
)
|
||||
|
||||
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 < 3; 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 < 2 {
|
||||
if err := sleepContext(ctx, time.Duration(1<<attempt)*time.Second); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return lastErr
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user