Initial release of bdrclone
This commit is contained in:
174
internal/auth/auth.go
Normal file
174
internal/auth/auth.go
Normal file
@@ -0,0 +1,174 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"html"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
||||
)
|
||||
|
||||
const OOBRedirectURI = "oob"
|
||||
|
||||
func Authorize(ctx context.Context, client *baidu.Client, redirectURI string, openBrowser bool) error {
|
||||
u, err := url.Parse(redirectURI)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse redirect_uri: %w", err)
|
||||
}
|
||||
if u.Scheme != "http" || u.Hostname() != "127.0.0.1" {
|
||||
return errors.New("automatic auth requires an http://127.0.0.1 redirect_uri")
|
||||
}
|
||||
state, err := randomState()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
listener, err := net.Listen("tcp", u.Host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("listen for OAuth callback on %s: %w", u.Host, err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
result := make(chan error, 1)
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc(u.Path, func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Query().Get("state") != state {
|
||||
http.Error(w, "OAuth state mismatch", http.StatusBadRequest)
|
||||
select {
|
||||
case result <- errors.New("OAuth state mismatch"):
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
if code := r.URL.Query().Get("error"); code != "" {
|
||||
message := r.URL.Query().Get("error_description")
|
||||
if message == "" {
|
||||
message = code
|
||||
}
|
||||
http.Error(w, message, http.StatusBadRequest)
|
||||
select {
|
||||
case result <- errors.New(message):
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
code := r.URL.Query().Get("code")
|
||||
if code == "" {
|
||||
http.Error(w, "Authorization code is missing", http.StatusBadRequest)
|
||||
select {
|
||||
case result <- errors.New("authorization code is missing"):
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
_, exchangeErr := client.ExchangeCode(r.Context(), code, redirectURI)
|
||||
if exchangeErr != nil {
|
||||
http.Error(w, html.EscapeString(exchangeErr.Error()), http.StatusBadGateway)
|
||||
} else {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = w.Write([]byte("<!doctype html><meta charset=utf-8><title>bdrclone</title><p>授权成功,可以关闭此页面。</p>"))
|
||||
}
|
||||
select {
|
||||
case result <- exchangeErr:
|
||||
default:
|
||||
}
|
||||
})
|
||||
server := &http.Server{Handler: mux, ReadHeaderTimeout: 10 * time.Second}
|
||||
serverErr := make(chan error, 1)
|
||||
go func() {
|
||||
if err := server.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
serverErr <- err
|
||||
}
|
||||
}()
|
||||
authorizeURL := client.AuthorizationURL(redirectURI, state)
|
||||
fmt.Printf("请在浏览器中授权:\n%s\n", authorizeURL)
|
||||
if openBrowser {
|
||||
_ = openURL(authorizeURL)
|
||||
}
|
||||
select {
|
||||
case err := <-result:
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
_ = server.Shutdown(shutdownCtx)
|
||||
return err
|
||||
case err := <-serverErr:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// AuthorizeOOB uses Baidu's documented out-of-band flow. Baidu displays the
|
||||
// authorization code in its own page instead of redirecting to a local server.
|
||||
func AuthorizeOOB(ctx context.Context, client *baidu.Client, openBrowser bool, input io.Reader, output io.Writer) error {
|
||||
state, err := randomState()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
authorizeURL := client.AuthorizationURL(OOBRedirectURI, state)
|
||||
fmt.Fprintf(output, "请在浏览器中授权:\n%s\n\n授权后,将页面显示的授权码粘贴到这里:", authorizeURL)
|
||||
if openBrowser {
|
||||
_ = openURL(authorizeURL)
|
||||
}
|
||||
|
||||
line, err := bufio.NewReader(input).ReadString('\n')
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return fmt.Errorf("read authorization code: %w", err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
code, err := parseAuthorizationCode(line)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = client.ExchangeCode(ctx, code, OOBRedirectURI)
|
||||
return err
|
||||
}
|
||||
|
||||
func parseAuthorizationCode(input string) (string, error) {
|
||||
value := strings.TrimSpace(input)
|
||||
if value == "" {
|
||||
return "", errors.New("authorization code is empty")
|
||||
}
|
||||
if parsed, err := url.Parse(value); err == nil && parsed.Query().Get("code") != "" {
|
||||
value = parsed.Query().Get("code")
|
||||
}
|
||||
if strings.ContainsAny(value, " \t\r\n") {
|
||||
return "", errors.New("authorization code contains whitespace")
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func randomState() (string, error) {
|
||||
b := make([]byte, 24)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", fmt.Errorf("generate OAuth state: %w", err)
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
func openURL(target string) error {
|
||||
var command string
|
||||
var args []string
|
||||
switch runtime.GOOS {
|
||||
case "darwin":
|
||||
command, args = "open", []string{target}
|
||||
case "windows":
|
||||
command, args = "rundll32", []string{"url.dll,FileProtocolHandler", target}
|
||||
default:
|
||||
command, args = "xdg-open", []string{target}
|
||||
}
|
||||
return exec.Command(command, args...).Start()
|
||||
}
|
||||
25
internal/auth/auth_test.go
Normal file
25
internal/auth/auth_test.go
Normal file
@@ -0,0 +1,25 @@
|
||||
package auth
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseAuthorizationCode(t *testing.T) {
|
||||
tests := map[string]string{
|
||||
"plain-code\n": "plain-code",
|
||||
"http://openapi.baidu.com/success?code=a%2Bb": "a+b",
|
||||
"https://example.test/?state=x&code=xyz": "xyz",
|
||||
}
|
||||
for input, want := range tests {
|
||||
got, err := parseAuthorizationCode(input)
|
||||
if err != nil {
|
||||
t.Fatalf("parseAuthorizationCode(%q): %v", input, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Errorf("parseAuthorizationCode(%q) = %q, want %q", input, got, want)
|
||||
}
|
||||
}
|
||||
for _, input := range []string{"", " \n", "two words"} {
|
||||
if _, err := parseAuthorizationCode(input); err == nil {
|
||||
t.Errorf("parseAuthorizationCode(%q) unexpectedly succeeded", input)
|
||||
}
|
||||
}
|
||||
}
|
||||
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) }
|
||||
211
internal/baidu/client_test.go
Normal file
211
internal/baidu/client_test.go
Normal 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
186
internal/baidu/files.go
Normal 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
91
internal/baidu/types.go
Normal 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
321
internal/baidu/upload.go
Normal 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
|
||||
}
|
||||
127
internal/config/config.go
Normal file
127
internal/config/config.go
Normal file
@@ -0,0 +1,127 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const DefaultUserAgent = "pan.baidu.com"
|
||||
|
||||
type Config struct {
|
||||
ClientID string `json:"client_id"`
|
||||
ClientSecret string `json:"client_secret"`
|
||||
RedirectURI string `json:"redirect_uri"`
|
||||
AccessToken string `json:"access_token,omitempty"`
|
||||
RefreshToken string `json:"refresh_token,omitempty"`
|
||||
ExpiresAt time.Time `json:"expires_at,omitempty"`
|
||||
Root string `json:"root"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
UploadParts int `json:"upload_parts"`
|
||||
PartSize int64 `json:"part_size,omitempty"`
|
||||
}
|
||||
|
||||
func DefaultPath() (string, error) {
|
||||
dir, err := os.UserConfigDir()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("find config directory: %w", err)
|
||||
}
|
||||
return filepath.Join(dir, "bdrclone", "config.json"), nil
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil, fmt.Errorf("config %q does not exist; run `bdrclone config` first", path)
|
||||
}
|
||||
return nil, fmt.Errorf("read config: %w", err)
|
||||
}
|
||||
var cfg Config
|
||||
if err := json.Unmarshal(b, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("parse config: %w", err)
|
||||
}
|
||||
cfg.applyDefaults()
|
||||
return &cfg, cfg.Validate(false)
|
||||
}
|
||||
|
||||
func Save(path string, cfg *Config) error {
|
||||
cfg.applyDefaults()
|
||||
if err := cfg.Validate(false); err != nil {
|
||||
return err
|
||||
}
|
||||
b, err := json.MarshalIndent(cfg, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode config: %w", err)
|
||||
}
|
||||
b = append(b, '\n')
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
|
||||
return fmt.Errorf("create config directory: %w", err)
|
||||
}
|
||||
tmp, err := os.CreateTemp(filepath.Dir(path), ".config-*")
|
||||
if err != nil {
|
||||
return fmt.Errorf("create temporary config: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
defer os.Remove(tmpName)
|
||||
if err := tmp.Chmod(0o600); err != nil {
|
||||
tmp.Close()
|
||||
return fmt.Errorf("secure config: %w", err)
|
||||
}
|
||||
if _, err := tmp.Write(b); err != nil {
|
||||
tmp.Close()
|
||||
return fmt.Errorf("write config: %w", err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return fmt.Errorf("close config: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
return fmt.Errorf("replace config: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Config) Validate(requireToken bool) error {
|
||||
if strings.TrimSpace(c.ClientID) == "" || strings.TrimSpace(c.ClientSecret) == "" {
|
||||
return errors.New("client_id and client_secret are required")
|
||||
}
|
||||
if requireToken && c.RefreshToken == "" && c.AccessToken == "" {
|
||||
return errors.New("no OAuth token; run `bdrclone auth` first")
|
||||
}
|
||||
if c.UploadParts < 1 || c.UploadParts > 32 {
|
||||
return errors.New("upload_parts must be between 1 and 32")
|
||||
}
|
||||
if c.PartSize != 0 && c.PartSize < 4<<20 {
|
||||
return errors.New("part_size must be 0 or at least 4 MiB")
|
||||
}
|
||||
if c.PartSize > 32<<20 {
|
||||
return errors.New("part_size cannot exceed Baidu's 32 MiB maximum")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Config) applyDefaults() {
|
||||
if c.RedirectURI == "" {
|
||||
c.RedirectURI = "http://127.0.0.1:53682/callback"
|
||||
}
|
||||
if c.Root == "" {
|
||||
c.Root = "/"
|
||||
}
|
||||
if !strings.HasPrefix(c.Root, "/") {
|
||||
c.Root = "/" + c.Root
|
||||
}
|
||||
c.Root = strings.TrimSuffix(c.Root, "/")
|
||||
if c.Root == "" {
|
||||
c.Root = "/"
|
||||
}
|
||||
if c.UserAgent == "" {
|
||||
c.UserAgent = DefaultUserAgent
|
||||
}
|
||||
if c.UploadParts == 0 {
|
||||
c.UploadParts = 3
|
||||
}
|
||||
}
|
||||
32
internal/config/config_test.go
Normal file
32
internal/config/config_test.go
Normal file
@@ -0,0 +1,32 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSaveLoadSecureConfig(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "nested", "config.json")
|
||||
cfg := &Config{ClientID: "id", ClientSecret: "secret"}
|
||||
if err := Save(path, cfg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loaded, err := Load(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if loaded.Root != "/" || loaded.UserAgent != DefaultUserAgent || loaded.UploadParts != 3 {
|
||||
t.Fatalf("defaults not applied: %+v", loaded)
|
||||
}
|
||||
if runtime.GOOS != "windows" {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info.Mode().Perm() != 0o600 {
|
||||
t.Fatalf("config permissions = %o", info.Mode().Perm())
|
||||
}
|
||||
}
|
||||
}
|
||||
341
internal/mount/cmount.go
Normal file
341
internal/mount/cmount.go
Normal file
@@ -0,0 +1,341 @@
|
||||
//go:build darwin && cgo && cmount
|
||||
|
||||
package mount
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
||||
"github.com/winfsp/cgofuse/fuse"
|
||||
)
|
||||
|
||||
type cMountFS struct {
|
||||
fuse.FileSystemBase
|
||||
client *baidu.Client
|
||||
readOnly bool
|
||||
|
||||
mu sync.Mutex
|
||||
next uint64
|
||||
handles map[uint64]*mountHandle
|
||||
}
|
||||
|
||||
type mountHandle struct {
|
||||
mu sync.Mutex
|
||||
read baidu.File
|
||||
write *os.File
|
||||
name string
|
||||
dirty bool
|
||||
}
|
||||
|
||||
func newCMountFS(client *baidu.Client, readOnly bool) *cMountFS {
|
||||
return &cMountFS{client: client, readOnly: readOnly, next: 1, handles: make(map[uint64]*mountHandle)}
|
||||
}
|
||||
|
||||
func (f *cMountFS) Getattr(name string, stat *fuse.Stat_t, fh uint64) int {
|
||||
if fh != ^uint64(0) {
|
||||
f.mu.Lock()
|
||||
handle, ok := f.handles[fh]
|
||||
f.mu.Unlock()
|
||||
if ok {
|
||||
handle.mu.Lock()
|
||||
defer handle.mu.Unlock()
|
||||
if handle.write == nil {
|
||||
return -fuse.EBADF
|
||||
}
|
||||
info, err := handle.write.Stat()
|
||||
if err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
stat.Mode, stat.Size = fuse.S_IFREG|0o644, info.Size()
|
||||
stat.Mtim = fuse.NewTimespec(info.ModTime())
|
||||
return 0
|
||||
}
|
||||
}
|
||||
if name == "/" {
|
||||
stat.Mode, stat.Nlink = fuse.S_IFDIR|0o755, 2
|
||||
return 0
|
||||
}
|
||||
entry, err := f.client.Stat(context.Background(), name)
|
||||
if err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
stat.Size = entry.Size
|
||||
stat.Mtim = fuse.NewTimespec(entry.ModTime())
|
||||
if entry.IsDirectory() {
|
||||
stat.Mode, stat.Nlink = fuse.S_IFDIR|0o755, 2
|
||||
} else {
|
||||
stat.Mode, stat.Nlink = fuse.S_IFREG|0o644, 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (f *cMountFS) Readdir(name string, fill func(string, *fuse.Stat_t, int64) bool, _ int64, _ uint64) int {
|
||||
entries, err := f.client.List(context.Background(), name)
|
||||
if err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
fill(".", nil, 0)
|
||||
fill("..", nil, 0)
|
||||
for _, entry := range entries {
|
||||
stat := &fuse.Stat_t{Size: entry.Size, Mtim: fuse.NewTimespec(entry.ModTime())}
|
||||
if entry.IsDirectory() {
|
||||
stat.Mode = fuse.S_IFDIR | 0o755
|
||||
} else {
|
||||
stat.Mode = fuse.S_IFREG | 0o644
|
||||
}
|
||||
if !fill(entry.Name(), stat, 0) {
|
||||
break
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (f *cMountFS) Open(name string, flags int) (int, uint64) {
|
||||
write := flags&(os.O_WRONLY|os.O_RDWR) != 0
|
||||
entry, err := f.client.Stat(context.Background(), name)
|
||||
if err != nil {
|
||||
return errno(err), 0
|
||||
}
|
||||
if !write {
|
||||
return 0, f.addHandle(&mountHandle{read: entry, name: name})
|
||||
}
|
||||
if f.readOnly {
|
||||
return -fuse.EROFS, 0
|
||||
}
|
||||
tmp, err := os.CreateTemp("", "bdrclone-write-*")
|
||||
if err != nil {
|
||||
return errno(err), 0
|
||||
}
|
||||
if flags&os.O_TRUNC == 0 && entry.Size > 0 {
|
||||
body, openErr := f.client.Open(context.Background(), entry, 0, 0)
|
||||
if openErr == nil {
|
||||
_, openErr = io.Copy(tmp, body)
|
||||
openErr = errors.Join(openErr, body.Close())
|
||||
}
|
||||
if openErr != nil {
|
||||
tmp.Close()
|
||||
os.Remove(tmp.Name())
|
||||
return errno(openErr), 0
|
||||
}
|
||||
}
|
||||
return 0, f.addHandle(&mountHandle{write: tmp, read: entry, name: name, dirty: flags&os.O_TRUNC != 0})
|
||||
}
|
||||
|
||||
func (f *cMountFS) Create(name string, flags int, mode uint32) (int, uint64) {
|
||||
if f.readOnly {
|
||||
return -fuse.EROFS, 0
|
||||
}
|
||||
tmp, err := os.CreateTemp("", "bdrclone-write-*")
|
||||
if err != nil {
|
||||
return errno(err), 0
|
||||
}
|
||||
return 0, f.addHandle(&mountHandle{write: tmp, name: name, dirty: true})
|
||||
}
|
||||
|
||||
func (f *cMountFS) Read(_ string, dest []byte, offset int64, fh uint64) int {
|
||||
handle, ok := f.handle(fh)
|
||||
if !ok {
|
||||
return -fuse.EBADF
|
||||
}
|
||||
handle.mu.Lock()
|
||||
defer handle.mu.Unlock()
|
||||
if handle.write != nil {
|
||||
n, err := handle.write.ReadAt(dest, offset)
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return errno(err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
body, err := f.client.Open(context.Background(), handle.read, offset, int64(len(dest)))
|
||||
if err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
n, readErr := io.ReadFull(body, dest)
|
||||
readErr = errors.Join(readErr, body.Close())
|
||||
if readErr != nil && !errors.Is(readErr, io.EOF) && !errors.Is(readErr, io.ErrUnexpectedEOF) {
|
||||
return errno(readErr)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func (f *cMountFS) Write(_ string, data []byte, offset int64, fh uint64) int {
|
||||
handle, ok := f.handle(fh)
|
||||
if !ok {
|
||||
return -fuse.EBADF
|
||||
}
|
||||
handle.mu.Lock()
|
||||
defer handle.mu.Unlock()
|
||||
if handle.write == nil {
|
||||
return -fuse.EBADF
|
||||
}
|
||||
n, err := handle.write.WriteAt(data, offset)
|
||||
if n > 0 {
|
||||
handle.dirty = true
|
||||
}
|
||||
if err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func (f *cMountFS) Flush(_ string, fh uint64) int { return f.flush(fh) }
|
||||
|
||||
func (f *cMountFS) Fsync(_ string, _ bool, fh uint64) int { return f.flush(fh) }
|
||||
|
||||
func (f *cMountFS) Truncate(_ string, size int64, fh uint64) int {
|
||||
handle, ok := f.handle(fh)
|
||||
if !ok {
|
||||
return -fuse.EBADF
|
||||
}
|
||||
handle.mu.Lock()
|
||||
defer handle.mu.Unlock()
|
||||
if handle.write == nil {
|
||||
return -fuse.EBADF
|
||||
}
|
||||
if err := handle.write.Truncate(size); err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
handle.dirty = true
|
||||
return 0
|
||||
}
|
||||
|
||||
func (f *cMountFS) Release(_ string, fh uint64) int {
|
||||
status := f.flush(fh)
|
||||
f.mu.Lock()
|
||||
handle, ok := f.handles[fh]
|
||||
delete(f.handles, fh)
|
||||
f.mu.Unlock()
|
||||
if ok {
|
||||
handle.mu.Lock()
|
||||
defer handle.mu.Unlock()
|
||||
if handle.write == nil {
|
||||
return status
|
||||
}
|
||||
closeErr := handle.write.Close()
|
||||
if status != 0 {
|
||||
recoveryPath, recoveryErr := preserveFailedWrite(handle.write.Name(), handle.name)
|
||||
fmt.Fprintf(os.Stderr, "bdrclone: upload failed for %s; local recovery file: %s\n", handle.name, recoveryPath)
|
||||
if closeErr != nil || recoveryErr != nil {
|
||||
status = -fuse.EIO
|
||||
}
|
||||
} else if err := errors.Join(closeErr, os.Remove(handle.write.Name())); err != nil {
|
||||
status = errno(err)
|
||||
}
|
||||
}
|
||||
return status
|
||||
}
|
||||
|
||||
func (f *cMountFS) flush(fh uint64) int {
|
||||
handle, ok := f.handle(fh)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
handle.mu.Lock()
|
||||
defer handle.mu.Unlock()
|
||||
if handle.write == nil || !handle.dirty {
|
||||
return 0
|
||||
}
|
||||
info, err := handle.write.Stat()
|
||||
if err == nil && info.Size() == 0 {
|
||||
err = errors.New("Baidu Netdisk does not allow empty files")
|
||||
}
|
||||
if err == nil {
|
||||
_, err = f.client.Upload(context.Background(), handle.write, info.Size(), handle.name, info.ModTime(), nil)
|
||||
}
|
||||
if err == nil {
|
||||
handle.dirty = false
|
||||
}
|
||||
return errno(err)
|
||||
}
|
||||
|
||||
func (f *cMountFS) Mkdir(name string, _ uint32) int {
|
||||
if f.readOnly {
|
||||
return -fuse.EROFS
|
||||
}
|
||||
_, err := f.client.Mkdir(context.Background(), name)
|
||||
return errno(err)
|
||||
}
|
||||
|
||||
func (f *cMountFS) Unlink(name string) int {
|
||||
if f.readOnly {
|
||||
return -fuse.EROFS
|
||||
}
|
||||
return errno(f.client.Delete(context.Background(), name))
|
||||
}
|
||||
func (f *cMountFS) Rmdir(name string) int {
|
||||
if f.readOnly {
|
||||
return -fuse.EROFS
|
||||
}
|
||||
entries, err := f.client.List(context.Background(), name)
|
||||
if err != nil {
|
||||
return errno(err)
|
||||
}
|
||||
if len(entries) > 0 {
|
||||
return -fuse.ENOTEMPTY
|
||||
}
|
||||
return errno(f.client.Delete(context.Background(), name))
|
||||
}
|
||||
func (f *cMountFS) Rename(oldName, newName string) int {
|
||||
if f.readOnly {
|
||||
return -fuse.EROFS
|
||||
}
|
||||
return errno(f.client.Move(context.Background(), oldName, newName))
|
||||
}
|
||||
|
||||
func (f *cMountFS) addHandle(handle *mountHandle) uint64 {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
id := f.next
|
||||
f.next++
|
||||
f.handles[id] = handle
|
||||
return id
|
||||
}
|
||||
|
||||
func (f *cMountFS) handle(id uint64) (*mountHandle, bool) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
h, ok := f.handles[id]
|
||||
return h, ok
|
||||
}
|
||||
|
||||
func errno(err error) int {
|
||||
if err == nil {
|
||||
return 0
|
||||
}
|
||||
if errors.Is(err, baidu.ErrNotFound) || errors.Is(err, os.ErrNotExist) {
|
||||
return -fuse.ENOENT
|
||||
}
|
||||
if errors.Is(err, syscall.EACCES) {
|
||||
return -fuse.EACCES
|
||||
}
|
||||
return -fuse.EIO
|
||||
}
|
||||
|
||||
func Mount(ctx context.Context, client *baidu.Client, mountpoint string, options Options) error {
|
||||
host := fuse.NewFileSystemHost(newCMountFS(client, options.ReadOnly))
|
||||
args := []string{"-o", "fsname=bdrclone", "-o", "subtype=bdrclone", "-o", "volname=Baidu Netdisk", "-o", "noappledouble", "-o", "noapplexattr"}
|
||||
if options.ReadOnly {
|
||||
args = append(args, "-o", "ro")
|
||||
}
|
||||
done := make(chan bool, 1)
|
||||
go func() { done <- host.Mount(mountpoint, args) }()
|
||||
select {
|
||||
case ok := <-done:
|
||||
if !ok {
|
||||
return errors.New("macFUSE mount failed")
|
||||
}
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
host.Unmount()
|
||||
<-done
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
var _ fuse.FileSystemInterface = (*cMountFS)(nil)
|
||||
368
internal/mount/fs.go
Normal file
368
internal/mount/fs.go
Normal file
@@ -0,0 +1,368 @@
|
||||
//go:build linux
|
||||
|
||||
package mount
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"io"
|
||||
"os"
|
||||
"path"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"bazil.org/fuse"
|
||||
"bazil.org/fuse/fs"
|
||||
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
||||
)
|
||||
|
||||
type FileSystem struct {
|
||||
client *baidu.Client
|
||||
readOnly bool
|
||||
}
|
||||
|
||||
func New(client *baidu.Client, readOnly bool) *FileSystem {
|
||||
return &FileSystem{client: client, readOnly: readOnly}
|
||||
}
|
||||
|
||||
func (f *FileSystem) Root() (fs.Node, error) {
|
||||
return &Dir{fs: f, name: "/"}, nil
|
||||
}
|
||||
|
||||
type Dir struct {
|
||||
fs *FileSystem
|
||||
name string
|
||||
}
|
||||
|
||||
func (d *Dir) Attr(_ context.Context, attr *fuse.Attr) error {
|
||||
attr.Inode = inode(d.name)
|
||||
attr.Mode = os.ModeDir | 0o755
|
||||
attr.Mtime = time.Now()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *Dir) Lookup(ctx context.Context, req *fuse.LookupRequest, _ *fuse.LookupResponse) (fs.Node, error) {
|
||||
remote := path.Join(d.name, req.Name)
|
||||
entry, err := d.fs.client.Stat(ctx, remote)
|
||||
if err != nil {
|
||||
if errors.Is(err, baidu.ErrNotFound) {
|
||||
return nil, syscall.ENOENT
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return d.fs.node(d.fs.client.APIPath(entry.Path), entry), nil
|
||||
}
|
||||
|
||||
func (d *Dir) ReadDirAll(ctx context.Context) ([]fuse.Dirent, error) {
|
||||
entries, err := d.fs.client.List(ctx, d.name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]fuse.Dirent, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
typeID := fuse.DT_File
|
||||
if entry.IsDirectory() {
|
||||
typeID = fuse.DT_Dir
|
||||
}
|
||||
result = append(result, fuse.Dirent{Inode: inode(d.fs.client.APIPath(entry.Path)), Name: entry.Name(), Type: typeID})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (d *Dir) Mkdir(ctx context.Context, req *fuse.MkdirRequest) (fs.Node, error) {
|
||||
if d.fs.readOnly {
|
||||
return nil, syscall.EROFS
|
||||
}
|
||||
name := path.Join(d.name, req.Name)
|
||||
entry, err := d.fs.client.Mkdir(ctx, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return d.fs.node(name, entry), nil
|
||||
}
|
||||
|
||||
func (d *Dir) Create(ctx context.Context, req *fuse.CreateRequest, resp *fuse.CreateResponse) (fs.Node, fs.Handle, error) {
|
||||
if d.fs.readOnly {
|
||||
return nil, nil, syscall.EROFS
|
||||
}
|
||||
name := path.Join(d.name, req.Name)
|
||||
tmp, err := os.CreateTemp("", "bdrclone-write-*")
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
node := &File{fs: d.fs, name: name, info: baidu.File{Path: d.fs.client.RemotePath(name), ServerFilename: req.Name}}
|
||||
handle := &writeHandle{node: node, file: tmp, dirty: true}
|
||||
node.active = handle
|
||||
resp.Flags |= fuse.OpenDirectIO
|
||||
return node, handle, nil
|
||||
}
|
||||
|
||||
func (d *Dir) Remove(ctx context.Context, req *fuse.RemoveRequest) error {
|
||||
if d.fs.readOnly {
|
||||
return syscall.EROFS
|
||||
}
|
||||
name := path.Join(d.name, req.Name)
|
||||
if req.Dir {
|
||||
entries, err := d.fs.client.List(ctx, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(entries) > 0 {
|
||||
return syscall.ENOTEMPTY
|
||||
}
|
||||
}
|
||||
return d.fs.client.Delete(ctx, name)
|
||||
}
|
||||
|
||||
func (d *Dir) Rename(ctx context.Context, req *fuse.RenameRequest, newDir fs.Node) error {
|
||||
if d.fs.readOnly {
|
||||
return syscall.EROFS
|
||||
}
|
||||
destinationDir, ok := newDir.(*Dir)
|
||||
if !ok {
|
||||
return syscall.ENOTDIR
|
||||
}
|
||||
return d.fs.client.Move(ctx, path.Join(d.name, req.OldName), path.Join(destinationDir.name, req.NewName))
|
||||
}
|
||||
|
||||
func (f *FileSystem) node(name string, entry baidu.File) fs.Node {
|
||||
if entry.IsDirectory() {
|
||||
return &Dir{fs: f, name: name}
|
||||
}
|
||||
return &File{fs: f, name: name, info: entry}
|
||||
}
|
||||
|
||||
type File struct {
|
||||
fs *FileSystem
|
||||
name string
|
||||
mu sync.Mutex
|
||||
info baidu.File
|
||||
active *writeHandle
|
||||
}
|
||||
|
||||
func (f *File) Attr(_ context.Context, attr *fuse.Attr) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
attr.Inode = inode(f.name)
|
||||
attr.Mode = 0o644
|
||||
attr.Size = uint64(f.info.Size)
|
||||
attr.Mtime = f.info.ModTime()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *File) Open(ctx context.Context, req *fuse.OpenRequest, resp *fuse.OpenResponse) (fs.Handle, error) {
|
||||
if !req.Flags.IsWriteOnly() && !req.Flags.IsReadWrite() {
|
||||
resp.Flags |= fuse.OpenDirectIO
|
||||
return &readHandle{file: f}, nil
|
||||
}
|
||||
if f.fs.readOnly {
|
||||
return nil, syscall.EROFS
|
||||
}
|
||||
f.mu.Lock()
|
||||
busy := f.active != nil
|
||||
f.mu.Unlock()
|
||||
if busy {
|
||||
return nil, syscall.EBUSY
|
||||
}
|
||||
tmp, err := os.CreateTemp("", "bdrclone-write-*")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if req.Flags&fuse.OpenTruncate == 0 && f.info.Size > 0 {
|
||||
body, err := f.fs.client.Open(ctx, f.info, 0, 0)
|
||||
if err != nil {
|
||||
tmp.Close()
|
||||
os.Remove(tmp.Name())
|
||||
return nil, err
|
||||
}
|
||||
_, copyErr := io.Copy(tmp, body)
|
||||
closeErr := body.Close()
|
||||
if copyErr != nil || closeErr != nil {
|
||||
tmp.Close()
|
||||
os.Remove(tmp.Name())
|
||||
return nil, errors.Join(copyErr, closeErr)
|
||||
}
|
||||
}
|
||||
resp.Flags |= fuse.OpenDirectIO
|
||||
handle := &writeHandle{node: f, file: tmp, dirty: req.Flags&fuse.OpenTruncate != 0}
|
||||
f.mu.Lock()
|
||||
f.active = handle
|
||||
f.mu.Unlock()
|
||||
return handle, nil
|
||||
}
|
||||
|
||||
func (f *File) Setattr(_ context.Context, req *fuse.SetattrRequest, _ *fuse.SetattrResponse) error {
|
||||
if f.fs.readOnly && req.Valid.Size() {
|
||||
return syscall.EROFS
|
||||
}
|
||||
if req.Valid.Size() {
|
||||
f.mu.Lock()
|
||||
handle := f.active
|
||||
f.mu.Unlock()
|
||||
if handle == nil {
|
||||
return syscall.EBADF
|
||||
}
|
||||
return handle.truncate(int64(req.Size))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *File) Fsync(ctx context.Context, _ *fuse.FsyncRequest) error {
|
||||
f.mu.Lock()
|
||||
handle := f.active
|
||||
f.mu.Unlock()
|
||||
if handle == nil {
|
||||
return nil
|
||||
}
|
||||
return handle.sync(ctx)
|
||||
}
|
||||
|
||||
type readHandle struct{ file *File }
|
||||
|
||||
func (h *readHandle) Read(ctx context.Context, req *fuse.ReadRequest, resp *fuse.ReadResponse) error {
|
||||
h.file.mu.Lock()
|
||||
info := h.file.info
|
||||
h.file.mu.Unlock()
|
||||
if req.Offset >= info.Size {
|
||||
resp.Data = nil
|
||||
return nil
|
||||
}
|
||||
size := min(int64(req.Size), info.Size-req.Offset)
|
||||
body, err := h.file.fs.client.Open(ctx, info, req.Offset, size)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer body.Close()
|
||||
resp.Data, err = io.ReadAll(io.LimitReader(body, size))
|
||||
return err
|
||||
}
|
||||
|
||||
type writeHandle struct {
|
||||
node *File
|
||||
file *os.File
|
||||
mu sync.Mutex
|
||||
dirty bool
|
||||
done bool
|
||||
}
|
||||
|
||||
func (h *writeHandle) Read(_ context.Context, req *fuse.ReadRequest, resp *fuse.ReadResponse) error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
buf := make([]byte, req.Size)
|
||||
n, err := h.file.ReadAt(buf, req.Offset)
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return err
|
||||
}
|
||||
resp.Data = buf[:n]
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *writeHandle) Write(_ context.Context, req *fuse.WriteRequest, resp *fuse.WriteResponse) error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
n, err := h.file.WriteAt(req.Data, req.Offset)
|
||||
resp.Size = n
|
||||
if n > 0 {
|
||||
h.dirty = true
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *writeHandle) Setattr(_ context.Context, req *fuse.SetattrRequest, _ *fuse.SetattrResponse) error {
|
||||
if !req.Valid.Size() {
|
||||
return nil
|
||||
}
|
||||
return h.truncate(int64(req.Size))
|
||||
}
|
||||
|
||||
func (h *writeHandle) truncate(size int64) error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if err := h.file.Truncate(size); err != nil {
|
||||
return err
|
||||
}
|
||||
h.dirty = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *writeHandle) Fsync(ctx context.Context, _ *fuse.FsyncRequest) error { return h.sync(ctx) }
|
||||
|
||||
func (h *writeHandle) Flush(ctx context.Context, _ *fuse.FlushRequest) error {
|
||||
return h.sync(ctx)
|
||||
}
|
||||
|
||||
func (h *writeHandle) Release(ctx context.Context, _ *fuse.ReleaseRequest) error {
|
||||
err := h.sync(ctx)
|
||||
h.mu.Lock()
|
||||
if !h.done {
|
||||
h.done = true
|
||||
closeErr := h.file.Close()
|
||||
if err != nil {
|
||||
recoveryPath, recoveryErr := preserveFailedWrite(h.file.Name(), h.node.name)
|
||||
fmt.Fprintf(os.Stderr, "bdrclone: upload failed for %s; local recovery file: %s\n", h.node.name, recoveryPath)
|
||||
err = errors.Join(err, closeErr, recoveryErr)
|
||||
} else {
|
||||
err = errors.Join(closeErr, os.Remove(h.file.Name()))
|
||||
}
|
||||
}
|
||||
h.mu.Unlock()
|
||||
h.node.mu.Lock()
|
||||
if h.node.active == h {
|
||||
h.node.active = nil
|
||||
}
|
||||
h.node.mu.Unlock()
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *writeHandle) sync(ctx context.Context) error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if !h.dirty || h.done {
|
||||
return nil
|
||||
}
|
||||
stat, err := h.file.Stat()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if stat.Size() == 0 {
|
||||
return errors.New("Baidu Netdisk does not support empty files")
|
||||
}
|
||||
entry, err := h.node.fs.client.Upload(ctx, h.file, stat.Size(), h.node.name, stat.ModTime(), nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("upload %s: %w", h.node.name, err)
|
||||
}
|
||||
h.node.mu.Lock()
|
||||
h.node.info = entry
|
||||
h.node.mu.Unlock()
|
||||
h.dirty = false
|
||||
return nil
|
||||
}
|
||||
|
||||
func inode(name string) uint64 {
|
||||
h := fnv.New64a()
|
||||
_, _ = h.Write([]byte(name))
|
||||
return h.Sum64()
|
||||
}
|
||||
|
||||
var (
|
||||
_ fs.FS = (*FileSystem)(nil)
|
||||
_ fs.Node = (*Dir)(nil)
|
||||
_ fs.NodeRequestLookuper = (*Dir)(nil)
|
||||
_ fs.HandleReadDirAller = (*Dir)(nil)
|
||||
_ fs.NodeMkdirer = (*Dir)(nil)
|
||||
_ fs.NodeCreater = (*Dir)(nil)
|
||||
_ fs.NodeRemover = (*Dir)(nil)
|
||||
_ fs.NodeRenamer = (*Dir)(nil)
|
||||
_ fs.Node = (*File)(nil)
|
||||
_ fs.NodeOpener = (*File)(nil)
|
||||
_ fs.NodeSetattrer = (*File)(nil)
|
||||
_ fs.NodeFsyncer = (*File)(nil)
|
||||
_ fs.HandleReader = (*readHandle)(nil)
|
||||
_ fs.HandleReader = (*writeHandle)(nil)
|
||||
_ fs.HandleWriter = (*writeHandle)(nil)
|
||||
_ fs.HandleFlusher = (*writeHandle)(nil)
|
||||
_ fs.HandleReleaser = (*writeHandle)(nil)
|
||||
)
|
||||
37
internal/mount/mount.go
Normal file
37
internal/mount/mount.go
Normal file
@@ -0,0 +1,37 @@
|
||||
//go:build linux
|
||||
|
||||
package mount
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"bazil.org/fuse"
|
||||
"bazil.org/fuse/fs"
|
||||
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
||||
)
|
||||
|
||||
func Mount(ctx context.Context, client *baidu.Client, mountpoint string, options Options) error {
|
||||
if options.Name == "" {
|
||||
options.Name = "bdrclone"
|
||||
}
|
||||
mountOptions := []fuse.MountOption{fuse.FSName(options.Name), fuse.Subtype("bdrclone")}
|
||||
if options.ReadOnly {
|
||||
mountOptions = append(mountOptions, fuse.ReadOnly())
|
||||
}
|
||||
conn, err := fuse.Mount(mountpoint, mountOptions...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("mount %s: %w", mountpoint, err)
|
||||
}
|
||||
defer conn.Close()
|
||||
server := fs.New(conn, &fs.Config{})
|
||||
serveErr := make(chan error, 1)
|
||||
go func() { serveErr <- server.Serve(New(client, options.ReadOnly)) }()
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
_ = fuse.Unmount(mountpoint)
|
||||
return <-serveErr
|
||||
}
|
||||
}
|
||||
14
internal/mount/mount_unsupported.go
Normal file
14
internal/mount/mount_unsupported.go
Normal file
@@ -0,0 +1,14 @@
|
||||
//go:build !linux && !(darwin && cgo && cmount)
|
||||
|
||||
package mount
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
||||
)
|
||||
|
||||
func Mount(context.Context, *baidu.Client, string, Options) error {
|
||||
return errors.New("this build has no mount backend; on macOS install macFUSE and rebuild with `go build -tags cmount ./cmd/bdrclone`")
|
||||
}
|
||||
6
internal/mount/options.go
Normal file
6
internal/mount/options.go
Normal file
@@ -0,0 +1,6 @@
|
||||
package mount
|
||||
|
||||
type Options struct {
|
||||
ReadOnly bool
|
||||
Name string
|
||||
}
|
||||
29
internal/mount/recovery.go
Normal file
29
internal/mount/recovery.go
Normal file
@@ -0,0 +1,29 @@
|
||||
package mount
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func preserveFailedWrite(tempPath, remotePath string) (string, error) {
|
||||
cacheDir, err := os.UserCacheDir()
|
||||
if err != nil {
|
||||
return tempPath, err
|
||||
}
|
||||
recoveryDir := filepath.Join(cacheDir, "bdrclone", "failed-writes")
|
||||
if err := os.MkdirAll(recoveryDir, 0o700); err != nil {
|
||||
return tempPath, err
|
||||
}
|
||||
name := strings.ReplaceAll(filepath.Base(remotePath), string(filepath.Separator), "_")
|
||||
if name == "" || name == "." {
|
||||
name = "remote-file"
|
||||
}
|
||||
destination := filepath.Join(recoveryDir, fmt.Sprintf("%s-%s", time.Now().Format("20060102-150405.000000000"), name))
|
||||
if err := os.Rename(tempPath, destination); err != nil {
|
||||
return tempPath, err
|
||||
}
|
||||
return destination, nil
|
||||
}
|
||||
136
internal/serve/http.go
Normal file
136
internal/serve/http.go
Normal file
@@ -0,0 +1,136 @@
|
||||
package serve
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
client *baidu.Client
|
||||
server *http.Server
|
||||
}
|
||||
|
||||
func New(client *baidu.Client, address string) *Server {
|
||||
s := &Server{client: client}
|
||||
s.server = &http.Server{Addr: address, Handler: s, ReadHeaderTimeout: 10 * time.Second}
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *Server) ListenAndServe(ctx context.Context) error {
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- s.server.ListenAndServe() }()
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if errors.Is(err, http.ErrServerClosed) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
return errors.Join(s.server.Shutdown(shutdownCtx), ctx.Err())
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodHead {
|
||||
w.Header().Set("Allow", "GET, HEAD")
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
entry, err := s.client.Stat(r.Context(), r.URL.Path)
|
||||
if err != nil {
|
||||
if errors.Is(err, baidu.ErrNotFound) {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
http.Error(w, err.Error(), http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
if entry.IsDirectory() {
|
||||
s.serveDirectory(w, r)
|
||||
return
|
||||
}
|
||||
s.serveFile(w, r, entry)
|
||||
}
|
||||
|
||||
func (s *Server) serveDirectory(w http.ResponseWriter, r *http.Request) {
|
||||
entries, err := s.client.List(r.Context(), r.URL.Path)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
if r.Method == http.MethodHead {
|
||||
return
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(entries)
|
||||
}
|
||||
|
||||
func (s *Server) serveFile(w http.ResponseWriter, r *http.Request, file baidu.File) {
|
||||
offset, length, partial, err := parseRange(r.Header.Get("Range"), file.Size)
|
||||
if err != nil {
|
||||
w.Header().Set("Content-Range", fmt.Sprintf("bytes */%d", file.Size))
|
||||
http.Error(w, err.Error(), http.StatusRequestedRangeNotSatisfiable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Accept-Ranges", "bytes")
|
||||
w.Header().Set("Last-Modified", file.ModTime().UTC().Format(http.TimeFormat))
|
||||
if r.Method == http.MethodHead || length == 0 {
|
||||
w.Header().Set("Content-Length", strconv.FormatInt(length, 10))
|
||||
if partial {
|
||||
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", offset, offset+length-1, file.Size))
|
||||
w.WriteHeader(http.StatusPartialContent)
|
||||
}
|
||||
return
|
||||
}
|
||||
body, err := s.client.Open(r.Context(), file, offset, length)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
defer body.Close()
|
||||
w.Header().Set("Content-Length", strconv.FormatInt(length, 10))
|
||||
if partial {
|
||||
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", offset, offset+length-1, file.Size))
|
||||
w.WriteHeader(http.StatusPartialContent)
|
||||
}
|
||||
_, _ = io.CopyN(w, body, length)
|
||||
}
|
||||
|
||||
func parseRange(value string, size int64) (offset, length int64, partial bool, err error) {
|
||||
if value == "" {
|
||||
return 0, size, false, nil
|
||||
}
|
||||
if !strings.HasPrefix(value, "bytes=") || strings.Contains(value, ",") {
|
||||
return 0, 0, false, errors.New("only one bytes range is supported")
|
||||
}
|
||||
parts := strings.SplitN(strings.TrimPrefix(value, "bytes="), "-", 2)
|
||||
if len(parts) != 2 || parts[0] == "" {
|
||||
return 0, 0, false, errors.New("suffix ranges are not supported")
|
||||
}
|
||||
start, parseErr := strconv.ParseInt(parts[0], 10, 64)
|
||||
if parseErr != nil || start < 0 || start >= size {
|
||||
return 0, 0, false, errors.New("invalid range start")
|
||||
}
|
||||
end := size - 1
|
||||
if parts[1] != "" {
|
||||
end, parseErr = strconv.ParseInt(parts[1], 10, 64)
|
||||
if parseErr != nil || end < start {
|
||||
return 0, 0, false, errors.New("invalid range end")
|
||||
}
|
||||
if end >= size {
|
||||
end = size - 1
|
||||
}
|
||||
}
|
||||
return start, end - start + 1, true, nil
|
||||
}
|
||||
29
internal/serve/http_test.go
Normal file
29
internal/serve/http_test.go
Normal file
@@ -0,0 +1,29 @@
|
||||
package serve
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseRange(t *testing.T) {
|
||||
tests := []struct {
|
||||
value string
|
||||
offset, length int64
|
||||
partial, shouldReject bool
|
||||
}{
|
||||
{"", 0, 10, false, false},
|
||||
{"bytes=2-5", 2, 4, true, false},
|
||||
{"bytes=7-", 7, 3, true, false},
|
||||
{"bytes=7-99", 7, 3, true, false},
|
||||
{"bytes=-3", 0, 0, false, true},
|
||||
{"bytes=10-11", 0, 0, false, true},
|
||||
{"items=0-1", 0, 0, false, true},
|
||||
}
|
||||
for _, test := range tests {
|
||||
offset, length, partial, err := parseRange(test.value, 10)
|
||||
if (err != nil) != test.shouldReject {
|
||||
t.Errorf("parseRange(%q) err=%v", test.value, err)
|
||||
continue
|
||||
}
|
||||
if err == nil && (offset != test.offset || length != test.length || partial != test.partial) {
|
||||
t.Errorf("parseRange(%q) = (%d,%d,%t), want (%d,%d,%t)", test.value, offset, length, partial, test.offset, test.length, test.partial)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user