Initial release of bdrclone
This commit is contained in:
327
internal/baidu/client.go
Normal file
327
internal/baidu/client.go
Normal 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) }
|
||||
Reference in New Issue
Block a user