328 lines
9.3 KiB
Go
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
|
|
}
|