Initial release of bdrclone

This commit is contained in:
2026-08-13 23:02:25 +08:00
commit d6d1956050
23 changed files with 3042 additions and 0 deletions

327
internal/baidu/client.go Normal file
View File

@@ -0,0 +1,327 @@
package baidu
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strconv"
"strings"
"sync"
"time"
"gitea.dddbg.com/youbin/bdrclone/internal/config"
)
const (
defaultAPIBase = "https://pan.baidu.com"
defaultOAuthBase = "https://openapi.baidu.com"
defaultUploadBase = "https://d.pcs.baidu.com"
)
type Client struct {
httpClient *http.Client
apiBase string
oauthBase string
uploadBase string
userAgent string
root string
uploadParts int
partSize int64
downloadMu sync.Mutex
downloadURL map[int64]cachedDownloadURL
mu sync.Mutex
clientID string
secret string
accessToken string
refresh string
expiresAt time.Time
onToken func(Token, time.Time) error
}
type Option func(*Client)
func WithHTTPClient(client *http.Client) Option {
return func(c *Client) { c.httpClient = client }
}
func WithEndpoints(api, oauth, upload string) Option {
return func(c *Client) {
if api != "" {
c.apiBase = strings.TrimSuffix(api, "/")
}
if oauth != "" {
c.oauthBase = strings.TrimSuffix(oauth, "/")
}
if upload != "" {
c.uploadBase = strings.TrimSuffix(upload, "/")
}
}
}
func WithTokenSaver(fn func(Token, time.Time) error) Option {
return func(c *Client) { c.onToken = fn }
}
func New(cfg *config.Config, opts ...Option) *Client {
c := &Client{
httpClient: &http.Client{Timeout: 0},
apiBase: defaultAPIBase, oauthBase: defaultOAuthBase, uploadBase: defaultUploadBase,
userAgent: cfg.UserAgent, root: cfg.Root,
uploadParts: cfg.UploadParts, partSize: cfg.PartSize,
downloadURL: make(map[int64]cachedDownloadURL),
clientID: cfg.ClientID, secret: cfg.ClientSecret,
accessToken: cfg.AccessToken, refresh: cfg.RefreshToken, expiresAt: cfg.ExpiresAt,
}
for _, opt := range opts {
opt(c)
}
return c
}
type cachedDownloadURL struct {
url string
expiresAt time.Time
}
func (c *Client) AuthorizationURL(redirectURI, state string) string {
q := url.Values{
"response_type": {"code"},
"client_id": {c.clientID},
"redirect_uri": {redirectURI},
"scope": {"basic,netdisk"},
}
if state != "" {
q.Set("state", state)
}
return c.oauthBase + "/oauth/2.0/authorize?" + q.Encode()
}
func (c *Client) ExchangeCode(ctx context.Context, code, redirectURI string) (Token, error) {
q := url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"client_id": {c.clientID},
"client_secret": {c.secret},
"redirect_uri": {redirectURI},
}
return c.fetchToken(ctx, q)
}
func (c *Client) RefreshToken(ctx context.Context) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.refreshLocked(ctx)
}
func (c *Client) refreshLocked(ctx context.Context) error {
if c.refresh == "" {
return errors.New("refresh token is empty; run `bdrclone auth`")
}
q := url.Values{
"grant_type": {"refresh_token"},
"refresh_token": {c.refresh},
"client_id": {c.clientID},
"client_secret": {c.secret},
}
tok, err := c.fetchTokenUnlocked(ctx, q)
if err != nil {
return err
}
return c.applyToken(tok)
}
func (c *Client) fetchToken(ctx context.Context, q url.Values) (Token, error) {
c.mu.Lock()
defer c.mu.Unlock()
tok, err := c.fetchTokenUnlocked(ctx, q)
if err != nil {
return Token{}, err
}
if err := c.applyToken(tok); err != nil {
return Token{}, err
}
return tok, nil
}
func (c *Client) fetchTokenUnlocked(ctx context.Context, q url.Values) (Token, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.oauthBase+"/oauth/2.0/token?"+q.Encode(), nil)
if err != nil {
return Token{}, err
}
resp, err := c.httpClient.Do(req)
if err != nil {
return Token{}, fmt.Errorf("OAuth request: %w", err)
}
defer resp.Body.Close()
b, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return Token{}, fmt.Errorf("read OAuth response: %w", err)
}
var oauthErr OAuthError
if json.Unmarshal(b, &oauthErr) == nil && oauthErr.ErrorCode != "" {
return Token{}, &oauthErr
}
if resp.StatusCode/100 != 2 {
return Token{}, fmt.Errorf("OAuth HTTP %s: %s", resp.Status, strings.TrimSpace(string(b)))
}
var tok Token
if err := json.Unmarshal(b, &tok); err != nil {
return Token{}, fmt.Errorf("decode OAuth response: %w", err)
}
if tok.AccessToken == "" {
return Token{}, errors.New("OAuth response contains no access_token")
}
return tok, nil
}
func (c *Client) applyToken(tok Token) error {
if tok.RefreshToken == "" {
tok.RefreshToken = c.refresh
}
c.accessToken, c.refresh = tok.AccessToken, tok.RefreshToken
// Refresh one minute early so a request never starts with an expiring token.
c.expiresAt = time.Now().Add(time.Duration(tok.ExpiresIn)*time.Second - time.Minute)
if c.onToken != nil {
return c.onToken(tok, c.expiresAt)
}
return nil
}
func (c *Client) ensureToken(ctx context.Context) error {
c.mu.Lock()
defer c.mu.Unlock()
if c.accessToken != "" && (c.expiresAt.IsZero() || time.Now().Before(c.expiresAt)) {
return nil
}
return c.refreshLocked(ctx)
}
func (c *Client) token() string {
c.mu.Lock()
defer c.mu.Unlock()
return c.accessToken
}
func (c *Client) RemotePath(name string) string {
name = "/" + strings.TrimPrefix(name, "/")
clean := path.Clean(name)
if c.root == "/" {
return clean
}
return path.Join(c.root, clean)
}
func (c *Client) APIPath(remote string) string {
if c.root == "/" {
return path.Clean(remote)
}
remote = path.Clean(remote)
if remote != c.root && !strings.HasPrefix(remote, c.root+"/") {
return "/"
}
trimmed := strings.TrimPrefix(remote, c.root)
if trimmed == "" {
return "/"
}
return "/" + strings.TrimPrefix(trimmed, "/")
}
func (c *Client) request(ctx context.Context, method, endpoint string, query url.Values, form url.Values, out any) error {
if err := c.ensureToken(ctx); err != nil {
return err
}
for attempt := 0; attempt < 3; attempt++ {
q := cloneValues(query)
q.Set("access_token", c.token())
var body io.Reader
if form != nil {
body = strings.NewReader(form.Encode())
}
req, err := http.NewRequestWithContext(ctx, method, endpoint+"?"+q.Encode(), body)
if err != nil {
return err
}
req.Header.Set("User-Agent", c.userAgent)
if form != nil {
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
}
resp, err := c.httpClient.Do(req)
if err != nil {
if attempt < 2 {
if err := sleepContext(ctx, time.Duration(1<<attempt)*time.Second); err != nil {
return err
}
continue
}
return fmt.Errorf("Baidu API request: %w", err)
}
b, readErr := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
resp.Body.Close()
if readErr != nil {
return fmt.Errorf("read Baidu API response: %w", readErr)
}
var apiErr APIError
_ = json.Unmarshal(b, &apiErr)
if apiErr.Code() == 111 || apiErr.Code() == -6 {
c.mu.Lock()
refreshErr := c.refreshLocked(ctx)
c.mu.Unlock()
if refreshErr != nil {
return refreshErr
}
continue
}
if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 {
if attempt < 2 {
if err := sleepContext(ctx, time.Duration(1<<attempt)*time.Second); err != nil {
return err
}
continue
}
}
if resp.StatusCode/100 != 2 {
return fmt.Errorf("Baidu API HTTP %s: %s", resp.Status, strings.TrimSpace(string(b)))
}
if apiErr.Code() != 0 {
return &apiErr
}
if out != nil && len(bytes.TrimSpace(b)) > 0 {
if err := json.Unmarshal(b, out); err != nil {
return fmt.Errorf("decode Baidu API response: %w", err)
}
}
return nil
}
return errors.New("Baidu API request failed after token refresh")
}
func (c *Client) xpanFile() string { return c.apiBase + "/rest/2.0/xpan/file" }
func (c *Client) xpanMultimedia() string { return c.apiBase + "/rest/2.0/xpan/multimedia" }
func cloneValues(src url.Values) url.Values {
dst := make(url.Values, len(src))
for k, values := range src {
dst[k] = append([]string(nil), values...)
}
return dst
}
func sleepContext(ctx context.Context, delay time.Duration) error {
t := time.NewTimer(delay)
defer t.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-t.C:
return nil
}
}
func intString(v int64) string { return strconv.FormatInt(v, 10) }

View File

@@ -0,0 +1,211 @@
package baidu
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"os"
"strconv"
"sync"
"testing"
"time"
"gitea.dddbg.com/youbin/bdrclone/internal/config"
)
func TestRemotePathWithRoot(t *testing.T) {
client := New(testConfig("token"))
for input, want := range map[string]string{
"/": "/apps/bdrclone", "docs/a.txt": "/apps/bdrclone/docs/a.txt", "/../a": "/apps/bdrclone/a",
} {
if got := client.RemotePath(input); got != want {
t.Errorf("RemotePath(%q) = %q, want %q", input, got, want)
}
}
if got := client.APIPath("/apps/bdrclone/docs/a.txt"); got != "/docs/a.txt" {
t.Fatalf("APIPath = %q", got)
}
}
func TestExpiredTokenRefreshesAndPersists(t *testing.T) {
var saved Token
var listCalls int
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/oauth/2.0/token":
if r.URL.Query().Get("refresh_token") != "refresh-old" {
t.Errorf("unexpected refresh token: %q", r.URL.Query().Get("refresh_token"))
}
fmt.Fprint(w, `{"access_token":"access-new","refresh_token":"refresh-new","expires_in":3600}`)
case "/rest/2.0/xpan/file":
listCalls++
if r.URL.Query().Get("access_token") != "access-new" {
t.Errorf("unexpected access token: %q", r.URL.Query().Get("access_token"))
}
if got := r.Header.Get("User-Agent"); got != "pan.baidu.com" {
t.Errorf("User-Agent = %q", got)
}
fmt.Fprint(w, `{"errno":0,"list":[]}`)
default:
http.NotFound(w, r)
}
}))
defer server.Close()
cfg := testConfig("expired")
cfg.RefreshToken = "refresh-old"
cfg.ExpiresAt = time.Now().Add(-time.Hour)
client := New(cfg, WithEndpoints(server.URL, server.URL, server.URL), WithTokenSaver(func(token Token, _ time.Time) error {
saved = token
return nil
}))
if _, err := client.List(context.Background(), "/"); err != nil {
t.Fatal(err)
}
if saved.AccessToken != "access-new" || saved.RefreshToken != "refresh-new" {
t.Fatalf("saved token = %+v", saved)
}
if listCalls != 1 {
t.Fatalf("list calls = %d", listCalls)
}
}
func TestOpenUsesOfficialUserAgentAndRange(t *testing.T) {
var server *httptest.Server
metadataCalls := 0
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/rest/2.0/xpan/multimedia":
metadataCalls++
fmt.Fprintf(w, `{"errno":0,"list":[{"fs_id":42,"dlink":%q}]}`, server.URL+"/download?x=1")
case "/download":
if got := r.Header.Get("User-Agent"); got != "pan.baidu.com" {
t.Errorf("User-Agent = %q", got)
}
if got := r.Header.Get("Range"); got != "bytes=2-5" {
t.Errorf("Range = %q", got)
}
if got := r.URL.Query().Get("access_token"); got != "token" {
t.Errorf("access_token = %q", got)
}
w.WriteHeader(http.StatusPartialContent)
fmt.Fprint(w, "2345")
default:
http.NotFound(w, r)
}
}))
defer server.Close()
client := New(testConfig("token"), WithEndpoints(server.URL, server.URL, server.URL))
body, err := client.Open(context.Background(), File{FSID: 42, Path: "/a"}, 2, 4)
if err != nil {
t.Fatal(err)
}
defer body.Close()
b, err := io.ReadAll(body)
if err != nil {
t.Fatal(err)
}
if string(b) != "2345" {
t.Fatalf("body = %q", b)
}
body, err = client.Open(context.Background(), File{FSID: 42, Path: "/a"}, 2, 4)
if err != nil {
t.Fatal(err)
}
body.Close()
if metadataCalls != 1 {
t.Fatalf("download metadata calls = %d, want cached URL to be reused", metadataCalls)
}
}
func TestUploadMultipartFlow(t *testing.T) {
content := append(bytes.Repeat([]byte("a"), int(defaultPartSize)), []byte("tail")...)
file, err := os.CreateTemp(t.TempDir(), "upload-*")
if err != nil {
t.Fatal(err)
}
if _, err := file.Write(content); err != nil {
t.Fatal(err)
}
var mu sync.Mutex
parts := map[int][]byte{}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/rest/2.0/xpan/file":
if err := r.ParseForm(); err != nil {
t.Error(err)
}
switch r.URL.Query().Get("method") {
case "precreate":
if r.Form.Get("path") != "/apps/bdrclone/remote.bin" {
t.Errorf("precreate path = %q", r.Form.Get("path"))
}
var blocks []string
if err := json.Unmarshal([]byte(r.Form.Get("block_list")), &blocks); err != nil || len(blocks) != 2 {
t.Errorf("block list = %q, err=%v", r.Form.Get("block_list"), err)
}
fmt.Fprint(w, `{"errno":0,"return_type":1,"uploadid":"upload-1","block_list":[0,1]}`)
case "create":
if r.Form.Get("uploadid") != "upload-1" {
t.Errorf("create uploadid = %q", r.Form.Get("uploadid"))
}
fmt.Fprint(w, `{"errno":0,"fs_id":99,"path":"/apps/bdrclone/remote.bin","server_filename":"remote.bin","size":4194308}`)
default:
http.Error(w, "unexpected method", http.StatusBadRequest)
}
case "/rest/2.0/pcs/superfile2":
part, _ := strconv.Atoi(r.URL.Query().Get("partseq"))
if r.URL.Query().Get("uploadid") != "upload-1" {
t.Errorf("part uploadid = %q", r.URL.Query().Get("uploadid"))
}
if err := r.ParseMultipartForm(defaultPartSize + 1024); err != nil {
t.Error(err)
return
}
partFile, _, err := r.FormFile("file")
if err != nil {
t.Error(err)
return
}
b, err := io.ReadAll(partFile)
partFile.Close()
if err != nil {
t.Error(err)
}
mu.Lock()
parts[part] = b
mu.Unlock()
fmt.Fprint(w, `{"md5":"ok"}`)
default:
http.NotFound(w, r)
}
}))
defer server.Close()
cfg := testConfig("token")
cfg.PartSize = defaultPartSize
cfg.UploadParts = 2
client := New(cfg, WithEndpoints(server.URL, server.URL, server.URL))
entry, err := client.Upload(context.Background(), file, int64(len(content)), "/remote.bin", time.Unix(1_700_000_000, 0), nil)
if err != nil {
t.Fatal(err)
}
if entry.FSID != 99 {
t.Fatalf("entry = %+v", entry)
}
mu.Lock()
defer mu.Unlock()
if !bytes.Equal(parts[0], content[:defaultPartSize]) || !bytes.Equal(parts[1], content[defaultPartSize:]) {
t.Fatalf("uploaded parts do not match input: sizes %d, %d", len(parts[0]), len(parts[1]))
}
}
func testConfig(token string) *config.Config {
return &config.Config{
ClientID: "client", ClientSecret: "secret", AccessToken: token,
Root: "/apps/bdrclone", UserAgent: "pan.baidu.com", UploadParts: 1,
}
}

186
internal/baidu/files.go Normal file
View File

@@ -0,0 +1,186 @@
package baidu
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strconv"
"strings"
"time"
)
var ErrNotFound = errors.New("remote path not found")
func (c *Client) List(ctx context.Context, name string) ([]File, error) {
remote := c.RemotePath(name)
var result []File
for start := 0; ; start += 1000 {
var response struct {
List []File `json:"list"`
}
q := url.Values{
"method": {"list"}, "dir": {remote}, "start": {strconv.Itoa(start)},
"limit": {"1000"}, "order": {"name"},
}
if err := c.request(ctx, http.MethodGet, c.xpanFile(), q, nil, &response); err != nil {
return nil, err
}
result = append(result, response.List...)
if len(response.List) < 1000 {
return result, nil
}
}
}
func (c *Client) Stat(ctx context.Context, name string) (File, error) {
remote := c.RemotePath(name)
if remote == c.root || (c.root == "/" && remote == "/") {
return File{Path: remote, ServerFilename: path.Base(remote), IsDir: 1}, nil
}
parent := path.Dir(c.APIPath(remote))
entries, err := c.List(ctx, parent)
if err != nil {
return File{}, err
}
for _, entry := range entries {
if entry.Path == remote || entry.Name() == path.Base(remote) {
return entry, nil
}
}
return File{}, fmt.Errorf("%w: %s", ErrNotFound, name)
}
func (c *Client) Mkdir(ctx context.Context, name string) (File, error) {
form := url.Values{
"path": {c.RemotePath(name)}, "size": {"0"}, "isdir": {"1"}, "rtype": {"3"},
}
var result File
if err := c.request(ctx, http.MethodPost, c.xpanFile(), url.Values{"method": {"create"}}, form, &result); err != nil {
return File{}, err
}
return result, nil
}
func (c *Client) Delete(ctx context.Context, name string) error {
remote := c.RemotePath(name)
if remote == c.root {
return errors.New("refusing to delete the configured remote root")
}
paths, _ := json.Marshal([]string{remote})
return c.manage(ctx, "delete", string(paths))
}
func (c *Client) Move(ctx context.Context, source, destination string) error {
src := c.RemotePath(source)
if src == c.root {
return errors.New("refusing to move the configured remote root")
}
dst := c.RemotePath(destination)
items := []map[string]string{{"path": src, "dest": path.Dir(dst), "newname": path.Base(dst)}}
b, _ := json.Marshal(items)
return c.manage(ctx, "move", string(b))
}
func (c *Client) Copy(ctx context.Context, source, destination string) error {
src := c.RemotePath(source)
dst := c.RemotePath(destination)
items := []map[string]string{{"path": src, "dest": path.Dir(dst), "newname": path.Base(dst)}}
b, _ := json.Marshal(items)
return c.manage(ctx, "copy", string(b))
}
func (c *Client) manage(ctx context.Context, operation, fileList string) error {
q := url.Values{"method": {"filemanager"}, "opera": {operation}}
form := url.Values{"async": {"0"}, "ondup": {"fail"}, "filelist": {fileList}}
return c.request(ctx, http.MethodPost, c.xpanFile(), q, form, nil)
}
func (c *Client) Quota(ctx context.Context) (Quota, error) {
var result Quota
err := c.request(ctx, http.MethodGet, c.apiBase+"/api/quota", nil, nil, &result)
return result, err
}
func (c *Client) DownloadURL(ctx context.Context, file File) (string, error) {
if file.FSID == 0 {
return "", errors.New("file has no fs_id")
}
c.downloadMu.Lock()
cached, ok := c.downloadURL[file.FSID]
c.downloadMu.Unlock()
if ok && time.Now().Before(cached.expiresAt) {
return cached.url, nil
}
ids, _ := json.Marshal([]int64{file.FSID})
var response struct {
List []File `json:"list"`
}
q := url.Values{"method": {"filemetas"}, "fsids": {string(ids)}, "dlink": {"1"}}
if err := c.request(ctx, http.MethodGet, c.xpanMultimedia(), q, nil, &response); err != nil {
return "", err
}
if len(response.List) == 0 || response.List[0].DLink == "" {
return "", errors.New("Baidu returned no download URL")
}
u, err := url.Parse(response.List[0].DLink)
if err != nil {
return "", fmt.Errorf("parse download URL: %w", err)
}
q2 := u.Query()
q2.Set("access_token", c.token())
u.RawQuery = q2.Encode()
result := u.String()
c.downloadMu.Lock()
c.downloadURL[file.FSID] = cachedDownloadURL{url: result, expiresAt: time.Now().Add(50 * time.Minute)}
c.downloadMu.Unlock()
return result, nil
}
func (c *Client) Open(ctx context.Context, file File, offset, length int64) (io.ReadCloser, error) {
for attempt := 0; attempt < 2; attempt++ {
downloadURL, err := c.DownloadURL(ctx, file)
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
if err != nil {
return nil, err
}
req.Header.Set("User-Agent", c.userAgent)
if offset > 0 || length > 0 {
end := ""
if length > 0 {
end = strconv.FormatInt(offset+length-1, 10)
}
req.Header.Set("Range", "bytes="+strconv.FormatInt(offset, 10)+"-"+end)
}
resp, err := c.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("download %s: %w", file.Path, err)
}
if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusPartialContent {
if resp.StatusCode == http.StatusOK && offset > 0 {
if _, err := io.CopyN(io.Discard, resp.Body, offset); err != nil {
resp.Body.Close()
return nil, fmt.Errorf("seek download %s to %d: %w", file.Path, offset, err)
}
}
return resp.Body, nil
}
b, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10))
resp.Body.Close()
if attempt == 0 && (resp.StatusCode == http.StatusForbidden || resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusNotFound) {
c.downloadMu.Lock()
delete(c.downloadURL, file.FSID)
c.downloadMu.Unlock()
continue
}
return nil, fmt.Errorf("download %s: HTTP %s: %s", file.Path, resp.Status, strings.TrimSpace(string(b)))
}
return nil, fmt.Errorf("download %s failed after refreshing its temporary URL", file.Path)
}

91
internal/baidu/types.go Normal file
View File

@@ -0,0 +1,91 @@
package baidu
import (
"fmt"
"path"
"time"
)
type APIError struct {
Errno int `json:"errno"`
ErrorCode int `json:"error_code"`
ErrorMsg string `json:"error_msg"`
Errmsg string `json:"errmsg"`
RequestID any `json:"request_id"`
}
func (e *APIError) Code() int {
if e.Errno != 0 {
return e.Errno
}
return e.ErrorCode
}
func (e *APIError) Error() string {
msg := e.ErrorMsg
if msg == "" {
msg = e.Errmsg
}
if msg == "" {
msg = "Baidu API request failed"
}
return fmt.Sprintf("%s (code=%d, request_id=%v)", msg, e.Code(), e.RequestID)
}
type File struct {
FSID int64 `json:"fs_id"`
Category int `json:"category"`
Size int64 `json:"size"`
Path string `json:"path"`
ServerFilename string `json:"server_filename"`
MD5 string `json:"md5"`
IsDir int `json:"isdir"`
ServerCTime int64 `json:"server_ctime"`
ServerMTime int64 `json:"server_mtime"`
LocalCTime int64 `json:"local_ctime"`
LocalMTime int64 `json:"local_mtime"`
CTime int64 `json:"ctime"`
MTime int64 `json:"mtime"`
DLink string `json:"dlink,omitempty"`
}
func (f File) Name() string {
if f.ServerFilename != "" {
return f.ServerFilename
}
return path.Base(f.Path)
}
func (f File) IsDirectory() bool { return f.IsDir == 1 }
func (f File) ModTime() time.Time {
stamp := f.LocalMTime
if stamp == 0 {
stamp = f.ServerMTime
}
if stamp == 0 {
stamp = f.MTime
}
return time.Unix(stamp, 0)
}
type Quota struct {
Total int64 `json:"total"`
Used int64 `json:"used"`
}
type Token struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
ExpiresIn int64 `json:"expires_in"`
Scope string `json:"scope"`
}
type OAuthError struct {
ErrorCode string `json:"error"`
ErrorDescription string `json:"error_description"`
}
func (e *OAuthError) Error() string {
return fmt.Sprintf("OAuth error %s: %s", e.ErrorCode, e.ErrorDescription)
}

321
internal/baidu/upload.go Normal file
View 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
}