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<