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<= 500 { if attempt < 2 { if err := sleepContext(ctx, time.Duration(1< 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) }