330 lines
8.2 KiB
Go
330 lines
8.2 KiB
Go
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
|
|
uploadRetryDelay func(int) time.Duration
|
|
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,
|
|
uploadRetryDelay: defaultUploadRetryDelay,
|
|
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) }
|