From d6d1956050e92448d831c6e79d72ccaa0605a68f Mon Sep 17 00:00:00 2001 From: youbin Date: Thu, 13 Aug 2026 23:02:25 +0800 Subject: [PATCH] Initial release of bdrclone --- .gitignore | 4 + Makefile | 10 + README.md | 151 +++++++++++ cmd/bdrclone/main.go | 390 ++++++++++++++++++++++++++++ go.mod | 15 ++ go.sum | 18 ++ internal/auth/auth.go | 174 +++++++++++++ internal/auth/auth_test.go | 25 ++ internal/baidu/client.go | 327 +++++++++++++++++++++++ internal/baidu/client_test.go | 211 +++++++++++++++ internal/baidu/files.go | 186 +++++++++++++ internal/baidu/types.go | 91 +++++++ internal/baidu/upload.go | 321 +++++++++++++++++++++++ internal/config/config.go | 127 +++++++++ internal/config/config_test.go | 32 +++ internal/mount/cmount.go | 341 ++++++++++++++++++++++++ internal/mount/fs.go | 368 ++++++++++++++++++++++++++ internal/mount/mount.go | 37 +++ internal/mount/mount_unsupported.go | 14 + internal/mount/options.go | 6 + internal/mount/recovery.go | 29 +++ internal/serve/http.go | 136 ++++++++++ internal/serve/http_test.go | 29 +++ 23 files changed, 3042 insertions(+) create mode 100644 .gitignore create mode 100644 Makefile create mode 100644 README.md create mode 100644 cmd/bdrclone/main.go create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal/auth/auth.go create mode 100644 internal/auth/auth_test.go create mode 100644 internal/baidu/client.go create mode 100644 internal/baidu/client_test.go create mode 100644 internal/baidu/files.go create mode 100644 internal/baidu/types.go create mode 100644 internal/baidu/upload.go create mode 100644 internal/config/config.go create mode 100644 internal/config/config_test.go create mode 100644 internal/mount/cmount.go create mode 100644 internal/mount/fs.go create mode 100644 internal/mount/mount.go create mode 100644 internal/mount/mount_unsupported.go create mode 100644 internal/mount/options.go create mode 100644 internal/mount/recovery.go create mode 100644 internal/serve/http.go create mode 100644 internal/serve/http_test.go diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..c8b89b7 --- /dev/null +++ b/.gitignore @@ -0,0 +1,4 @@ +/bdrclone +/dist/ +*.log +.DS_Store diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..0752525 --- /dev/null +++ b/Makefile @@ -0,0 +1,10 @@ +.PHONY: build test clean + +build: + go build -o bdrclone ./cmd/bdrclone + +test: + go test ./... + +clean: + rm -f bdrclone diff --git a/README.md b/README.md new file mode 100644 index 0000000..b40653d --- /dev/null +++ b/README.md @@ -0,0 +1,151 @@ +# bdrclone + +`bdrclone` 是面向百度网盘官方开放 API 的命令行和 FUSE 挂载客户端。它参考了 +rclone 的命令/VFS 语义和 OpenList 的百度驱动流程,但代码独立实现,只调用百度官方接口。 + +已实现: + +- OAuth 授权码登录、过期自动刷新、token 原子持久化 +- `ls`、`stat`、`cat`、`mkdir`、`rm`、`mv`、`cp`、`quota` +- 官方 dlink 下载、`User-Agent: pan.baidu.com`、HTTP Range 随机读取 +- MD5 秒传尝试、4/16/32 MiB 自动分片、多分片并发上传 +- Linux FUSE 挂载;macOS 使用 macFUSE/cgofuse 可选构建 +- 只读 HTTP 服务,支持单 Range 请求 + +## 准备百度应用 + +1. 在[百度网盘开放平台](https://pan.baidu.com/union/)完成开发者认证并创建“软件”应用。 +2. 在应用安全设置中加入回调地址:`http://127.0.0.1:53682/callback`。百度说明修改后最长约 + 1 小时生效。 +3. 记下应用的 AppKey (`client_id`) 和 SecretKey (`client_secret`)。 + +不需要把应用提交上线审核即可授权自己的测试账号,但可用能力最终以百度控制台给应用开通的 +权限为准。 + +## 构建 + +需要 Go 1.24 或更高版本。 + +```bash +go build -o bdrclone ./cmd/bdrclone +go test ./... +``` + +### macOS 挂载构建 + +先安装 [macFUSE](https://osxfuse.github.io/),再让 cgo 找到头文件和动态库: + +```bash +brew install --cask macfuse +CGO_CFLAGS="-I/usr/local/include/fuse" \ +CGO_LDFLAGS="-L/usr/local/lib -lfuse" \ +go build -tags cmount -o bdrclone ./cmd/bdrclone +``` + +Apple Silicon 上 macFUSE 的实际 include/lib 路径可能随版本变化;以安装包提供的路径为准。 +未加 `cmount` 标签的 macOS 构建仍包含全部 CLI/HTTP 功能,执行 `mount` 时会给出依赖提示。 + +Linux 直接构建即可得到 FUSE 后端,运行时需要系统已安装 FUSE 设备和挂载工具。 + +## 配置与授权 + +```bash +./bdrclone config \ + --client-id '你的 AppKey' \ + --client-secret '你的 SecretKey' + +./bdrclone auth +``` + +如果控制台尚未接受本机回调地址,可使用百度官方 OOB 模式。授权后,将百度页面显示的授权码 +粘贴回终端: + +```bash +./bdrclone auth --oob +``` + +OOB 模式使用 `redirect_uri=oob`,无需在本机监听端口。 + +默认配置文件: + +- macOS: `~/Library/Application Support/bdrclone/config.json` +- Linux: `~/.config/bdrclone/config.json` + +配置文件以 `0600` 权限保存。可以用 `--config /path/config.json` 指定其他位置。 + +只暴露网盘中的一个子目录: + +```bash +./bdrclone config --root /我的资料 +``` + +## 使用 + +```bash +./bdrclone ls / +./bdrclone stat /文档/report.pdf +./bdrclone download /文档/report.pdf ./report.pdf +./bdrclone upload ./photo.jpg /备份/photo.jpg +./bdrclone mkdir /备份/新目录 +./bdrclone mv /备份/a.txt /备份/b.txt +./bdrclone cp /备份/b.txt /副本/b.txt +./bdrclone rm /副本/b.txt +# 目录删除必须显式确认递归语义;配置的根目录始终禁止删除 +./bdrclone rm --recursive /旧备份 +./bdrclone quota +``` + +挂载: + +```bash +mkdir -p ~/BaiduNetdisk +./bdrclone mount ~/BaiduNetdisk +# 只读模式 +./bdrclone mount --read-only ~/BaiduNetdisk +``` + +挂载写入使用本地临时文件,文件关闭/flush 时整体分片上传。随机读取映射为官方 dlink 的 Range +请求;随机写不是云端原地修改,而是“下载旧文件到缓存、修改、重新上传”。 +写回失败的临时内容会保留到用户缓存目录的 `bdrclone/failed-writes`,并在挂载进程的标准错误中 +打印恢复路径。 + +HTTP 服务默认只监听本机: + +```bash +./bdrclone serve --addr 127.0.0.1:8080 +curl http://127.0.0.1:8080/文档/report.pdf -o report.pdf +``` + +目录 URL 返回 JSON。不要在没有鉴权或反向代理保护的情况下监听公网地址。 + +## 百度 API 限制 + +- 百度网盘官方接口不允许创建空文件;挂载下的空文件在 flush 时会失败。 +- 大文件下载必须使用 `User-Agent: pan.baidu.com`,客户端已统一设置。 +- 分片大小由会员等级决定:普通用户 4 MiB、会员 16 MiB、超级会员 32 MiB;百度限制分片数, + 因此不同等级的单文件上限不同。 +- `--upload-parts` 控制并发(默认 3)。低上行带宽遇到超时时建议设为 1。 +- dlink 是临时链接,每次打开文件会重新获取;不会把链接写入长期缓存。 +- 本程序只走官方接口,不实现 OpenList 文档中已经失效的 `crack`/`crack_video` 接口,也不会绕过 + 百度会员限速或开放平台权限。 + +## 参考资料 + +- [rclone](https://github.com/rclone/rclone) +- [OpenList](https://github.com/OpenListTeam/openlist) +- [OpenList 百度网盘驱动说明](https://doc.oplist.org/guide/drivers/baidu) +- [百度授权码模式](https://pan.baidu.com/union/doc/%E4%BD%BF%E7%94%A8%E5%85%A5%E9%97%A8/%E6%8E%A5%E5%85%A5%E6%8E%88%E6%9D%83/%E6%8E%88%E6%9D%83%E7%A0%81%E6%A8%A1%E5%BC%8F/) +- [百度文件列表](https://pan.baidu.com/union/doc/%E5%9F%BA%E7%A1%80%E7%BD%91%E7%9B%98%E6%9C%8D%E5%8A%A1/%E8%8E%B7%E5%8F%96%E6%96%87%E4%BB%B6%E4%BF%A1%E6%81%AF/%E8%8E%B7%E5%8F%96%E6%96%87%E4%BB%B6%E5%88%97%E8%A1%A8/) +- [百度上传接口](https://pan.baidu.com/union/doc/%E5%9F%BA%E7%A1%80%E7%BD%91%E7%9B%98%E6%9C%8D%E5%8A%A1/%E4%B8%8A%E4%BC%A0/%E9%A2%84%E4%B8%8A%E4%BC%A0/) +- [百度下载接口](https://pan.baidu.com/union/doc/%E5%9F%BA%E7%A1%80%E7%BD%91%E7%9B%98%E6%9C%8D%E5%8A%A1/%E4%B8%8B%E8%BD%BD/) + +## 开发验证 + +测试使用本地模拟百度端点,不需要真实账号: + +```bash +go test -race ./... +go vet ./... +``` + +真实账号集成测试需要你自己的 AppKey/SecretKey 和授权 token,仓库不会保存这些凭据。 diff --git a/cmd/bdrclone/main.go b/cmd/bdrclone/main.go new file mode 100644 index 0000000..3640713 --- /dev/null +++ b/cmd/bdrclone/main.go @@ -0,0 +1,390 @@ +package main + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/signal" + "path/filepath" + "strconv" + "syscall" + "time" + + "gitea.dddbg.com/youbin/bdrclone/internal/auth" + "gitea.dddbg.com/youbin/bdrclone/internal/baidu" + "gitea.dddbg.com/youbin/bdrclone/internal/config" + "gitea.dddbg.com/youbin/bdrclone/internal/mount" + "gitea.dddbg.com/youbin/bdrclone/internal/serve" + "github.com/spf13/cobra" +) + +var version = "dev" + +type application struct { + configPath string +} + +func main() { + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + app := &application{} + root := app.command() + if err := root.ExecuteContext(ctx); err != nil { + fmt.Fprintln(os.Stderr, "错误:", err) + os.Exit(1) + } +} + +func (a *application) command() *cobra.Command { + defaultConfig, _ := config.DefaultPath() + cmd := &cobra.Command{ + Use: "bdrclone", + Short: "百度网盘命令行和 FUSE 挂载客户端", + SilenceUsage: true, + SilenceErrors: true, + Version: version, + } + cmd.PersistentFlags().StringVar(&a.configPath, "config", defaultConfig, "配置文件路径") + cmd.AddCommand( + a.configCommand(), a.authCommand(), a.lsCommand(), a.statCommand(), a.catCommand(), + a.mkdirCommand(), a.rmCommand(), a.mvCommand(), a.cpCommand(), a.uploadCommand(), + a.downloadCommand(), a.quotaCommand(), a.mountCommand(), a.serveCommand(), + ) + return cmd +} + +func (a *application) client() (*baidu.Client, *config.Config, error) { + cfg, err := config.Load(a.configPath) + if err != nil { + return nil, nil, err + } + if err := cfg.Validate(true); err != nil { + return nil, nil, err + } + client := baidu.New(cfg, baidu.WithTokenSaver(func(token baidu.Token, expiresAt time.Time) error { + cfg.AccessToken = token.AccessToken + cfg.RefreshToken = token.RefreshToken + cfg.ExpiresAt = expiresAt + return config.Save(a.configPath, cfg) + })) + return client, cfg, nil +} + +func (a *application) configCommand() *cobra.Command { + var clientID, secret, redirect, root, userAgent string + var uploadParts int + var partSize int64 + cmd := &cobra.Command{ + Use: "config", + Short: "创建或更新百度开放平台配置", + RunE: func(cmd *cobra.Command, _ []string) error { + cfg := &config.Config{} + if old, err := config.Load(a.configPath); err == nil { + cfg = old + } + if clientID != "" { + cfg.ClientID = clientID + } + if secret != "" { + cfg.ClientSecret = secret + } + if redirect != "" { + cfg.RedirectURI = redirect + } + if root != "" { + cfg.Root = root + } + if userAgent != "" { + cfg.UserAgent = userAgent + } + if uploadParts != 0 { + cfg.UploadParts = uploadParts + } + if cmd.Flags().Changed("part-size") { + cfg.PartSize = partSize + } + if err := config.Save(a.configPath, cfg); err != nil { + return err + } + fmt.Println("配置已保存到", a.configPath) + return nil + }, + } + cmd.Flags().StringVar(&clientID, "client-id", "", "百度开放平台 AppKey") + cmd.Flags().StringVar(&secret, "client-secret", "", "百度开放平台 SecretKey") + cmd.Flags().StringVar(&redirect, "redirect-uri", "", "OAuth 回调地址") + cmd.Flags().StringVar(&root, "root", "", "挂载的网盘根路径") + cmd.Flags().StringVar(&userAgent, "user-agent", "", "下载 User-Agent") + cmd.Flags().IntVar(&uploadParts, "upload-parts", 0, "并发上传分片数 (1-32)") + cmd.Flags().Int64Var(&partSize, "part-size", 0, "上传分片字节数,0 表示按会员等级自动") + return cmd +} + +func (a *application) authCommand() *cobra.Command { + var noOpen, oob bool + cmd := &cobra.Command{ + Use: "auth", + Short: "通过 OAuth 授权百度网盘", + RunE: func(cmd *cobra.Command, _ []string) error { + cfg, err := config.Load(a.configPath) + if err != nil { + return err + } + client := baidu.New(cfg, baidu.WithTokenSaver(func(token baidu.Token, expiresAt time.Time) error { + cfg.AccessToken, cfg.RefreshToken, cfg.ExpiresAt = token.AccessToken, token.RefreshToken, expiresAt + return config.Save(a.configPath, cfg) + })) + if oob { + err = auth.AuthorizeOOB(cmd.Context(), client, !noOpen, cmd.InOrStdin(), cmd.OutOrStdout()) + } else { + err = auth.Authorize(cmd.Context(), client, cfg.RedirectURI, !noOpen) + } + if err != nil { + return err + } + fmt.Println("授权成功") + return nil + }, + } + cmd.Flags().BoolVar(&noOpen, "no-open", false, "不自动打开浏览器") + cmd.Flags().BoolVar(&oob, "oob", false, "使用百度页面显示授权码,不启动本机回调服务") + return cmd +} + +func (a *application) lsCommand() *cobra.Command { + var asJSON bool + cmd := &cobra.Command{Use: "ls [远端目录]", Args: cobra.MaximumNArgs(1), Short: "列出远端目录", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + name := "/" + if len(args) > 0 { + name = args[0] + } + entries, err := client.List(cmd.Context(), name) + if err != nil { + return err + } + if asJSON { + return printJSON(entries) + } + for _, entry := range entries { + kind := "-" + if entry.IsDirectory() { + kind = "d" + } + fmt.Printf("%s %12d %s %s\n", kind, entry.Size, entry.ModTime().Format("2006-01-02 15:04:05"), entry.Name()) + } + return nil + }} + cmd.Flags().BoolVar(&asJSON, "json", false, "输出 JSON") + return cmd +} + +func (a *application) statCommand() *cobra.Command { + return &cobra.Command{Use: "stat <远端路径>", Args: cobra.ExactArgs(1), Short: "查看远端文件元数据", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + entry, err := client.Stat(cmd.Context(), args[0]) + if err != nil { + return err + } + return printJSON(entry) + }} +} + +func (a *application) catCommand() *cobra.Command { + return &cobra.Command{Use: "cat <远端文件>", Args: cobra.ExactArgs(1), Short: "输出远端文件", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + entry, err := client.Stat(cmd.Context(), args[0]) + if err != nil { + return err + } + body, err := client.Open(cmd.Context(), entry, 0, 0) + if err != nil { + return err + } + defer body.Close() + _, err = io.Copy(os.Stdout, body) + return err + }} +} + +func (a *application) mkdirCommand() *cobra.Command { + return &cobra.Command{Use: "mkdir <远端目录>", Args: cobra.ExactArgs(1), Short: "创建远端目录", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + _, err = client.Mkdir(cmd.Context(), args[0]) + return err + }} +} + +func (a *application) rmCommand() *cobra.Command { + var recursive bool + cmd := &cobra.Command{Use: "rm <远端路径>", Args: cobra.ExactArgs(1), Short: "删除远端文件或目录", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + entry, err := client.Stat(cmd.Context(), args[0]) + if err != nil { + return err + } + if entry.IsDirectory() && !recursive { + return errors.New("target is a directory; pass --recursive to delete it") + } + return client.Delete(cmd.Context(), args[0]) + }} + cmd.Flags().BoolVarP(&recursive, "recursive", "r", false, "递归删除目录") + return cmd +} + +func (a *application) mvCommand() *cobra.Command { + return a.manageCommand("mv", "移动或重命名远端路径", func(ctx context.Context, c *baidu.Client, from, to string) error { return c.Move(ctx, from, to) }) +} +func (a *application) cpCommand() *cobra.Command { + return a.manageCommand("cp", "复制远端路径", func(ctx context.Context, c *baidu.Client, from, to string) error { return c.Copy(ctx, from, to) }) +} + +func (a *application) manageCommand(use, short string, fn func(context.Context, *baidu.Client, string, string) error) *cobra.Command { + return &cobra.Command{Use: use + " <源路径> <目标路径>", Args: cobra.ExactArgs(2), Short: short, RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + return fn(cmd.Context(), client, args[0], args[1]) + }} +} + +func (a *application) uploadCommand() *cobra.Command { + return &cobra.Command{Use: "upload <本地文件> <远端文件>", Args: cobra.ExactArgs(2), Short: "分片上传本地文件", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + _, err = client.UploadFile(cmd.Context(), args[0], args[1], func(done, total int64) { + fmt.Fprintf(os.Stderr, "\r上传 %d/%d bytes (%d%%)", done, total, done*100/total) + }) + if err == nil { + fmt.Fprintln(os.Stderr) + } + return err + }} +} + +func (a *application) downloadCommand() *cobra.Command { + return &cobra.Command{Use: "download <远端文件> <本地文件>", Args: cobra.ExactArgs(2), Short: "下载远端文件", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + entry, err := client.Stat(cmd.Context(), args[0]) + if err != nil { + return err + } + body, err := client.Open(cmd.Context(), entry, 0, 0) + if err != nil { + return err + } + defer body.Close() + if err := os.MkdirAll(filepath.Dir(args[1]), 0o755); err != nil { + return err + } + tmp, err := os.CreateTemp(filepath.Dir(args[1]), ".bdrclone-download-*") + if err != nil { + return err + } + tmpName := tmp.Name() + defer os.Remove(tmpName) + _, copyErr := io.Copy(tmp, body) + closeErr := tmp.Close() + if err := errors.Join(copyErr, closeErr); err != nil { + return err + } + return os.Rename(tmpName, args[1]) + }} +} + +func (a *application) quotaCommand() *cobra.Command { + return &cobra.Command{Use: "quota", Short: "查看网盘容量", RunE: func(cmd *cobra.Command, _ []string) error { + client, _, err := a.client() + if err != nil { + return err + } + quota, err := client.Quota(cmd.Context()) + if err != nil { + return err + } + percent := int64(0) + if quota.Total > 0 { + percent = quota.Used * 100 / quota.Total + } + fmt.Printf("已用 %s / 总计 %s (%d%%)\n", formatBytes(quota.Used), formatBytes(quota.Total), percent) + return nil + }} +} + +func (a *application) mountCommand() *cobra.Command { + var readOnly bool + cmd := &cobra.Command{Use: "mount <挂载点>", Args: cobra.ExactArgs(1), Short: "通过 FUSE 挂载百度网盘", RunE: func(cmd *cobra.Command, args []string) error { + client, _, err := a.client() + if err != nil { + return err + } + if err := os.MkdirAll(args[0], 0o755); err != nil { + return err + } + fmt.Println("正在挂载", args[0], ",按 Ctrl-C 卸载") + return mount.Mount(cmd.Context(), client, args[0], mount.Options{ReadOnly: readOnly}) + }} + cmd.Flags().BoolVar(&readOnly, "read-only", false, "只读挂载") + return cmd +} + +func (a *application) serveCommand() *cobra.Command { + var address string + cmd := &cobra.Command{Use: "serve", Short: "启动只读 HTTP 文件服务", RunE: func(cmd *cobra.Command, _ []string) error { + client, _, err := a.client() + if err != nil { + return err + } + fmt.Println("HTTP 服务监听 http://" + address) + err = serve.New(client, address).ListenAndServe(cmd.Context()) + if errors.Is(err, context.Canceled) { + return nil + } + return err + }} + cmd.Flags().StringVar(&address, "addr", "127.0.0.1:8080", "监听地址") + return cmd +} + +func printJSON(value any) error { + encoder := json.NewEncoder(os.Stdout) + encoder.SetIndent("", " ") + return encoder.Encode(value) +} + +func formatBytes(value int64) string { + const unit = int64(1024) + if value < unit { + return strconv.FormatInt(value, 10) + " B" + } + div, exp := unit, 0 + for n := value / unit; n >= unit && exp < 5; n /= unit { + div *= unit + exp++ + } + return fmt.Sprintf("%.1f %ciB", float64(value)/float64(div), "KMGTPE"[exp]) +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..181c68b --- /dev/null +++ b/go.mod @@ -0,0 +1,15 @@ +module gitea.dddbg.com/youbin/bdrclone + +go 1.24 + +require ( + bazil.org/fuse v0.0.0-20230120002735-62a210ff1fd5 + github.com/spf13/cobra v1.10.1 + github.com/winfsp/cgofuse v1.6.1-0.20260126094232-f2c4fccdb286 +) + +require ( + github.com/inconshreveable/mousetrap v1.1.0 // indirect + github.com/spf13/pflag v1.0.9 // indirect + golang.org/x/sys v0.4.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..b9ed1c2 --- /dev/null +++ b/go.sum @@ -0,0 +1,18 @@ +bazil.org/fuse v0.0.0-20230120002735-62a210ff1fd5 h1:A0NsYy4lDBZAC6QiYeJ4N+XuHIKBpyhAVRMHRQZKTeQ= +bazil.org/fuse v0.0.0-20230120002735-62a210ff1fd5/go.mod h1:gG3RZAMXCa/OTes6rr9EwusmR1OH1tDDy+cg9c5YliY= +github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= +github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= +github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= +github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= +github.com/spf13/cobra v1.10.1 h1:lJeBwCfmrnXthfAupyUTzJ/J4Nc1RsHC/mSRU2dll/s= +github.com/spf13/cobra v1.10.1/go.mod h1:7SmJGaTHFVBY0jW4NXGluQoLvhqFQM+6XSKD+P4XaB0= +github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY= +github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= +github.com/tv42/httpunix v0.0.0-20191220191345-2ba4b9c3382c h1:u6SKchux2yDvFQnDHS3lPnIRmfVJ5Sxy3ao2SIdysLQ= +github.com/tv42/httpunix v0.0.0-20191220191345-2ba4b9c3382c/go.mod h1:hzIxponao9Kjc7aWznkXaL4U4TWaDSs8zcsY4Ka08nM= +github.com/winfsp/cgofuse v1.6.1-0.20260126094232-f2c4fccdb286 h1:tw5GqRXqExB/xghPoPLtVujBe9w9Pg1G78tvXCJNJAA= +github.com/winfsp/cgofuse v1.6.1-0.20260126094232-f2c4fccdb286/go.mod h1:uxjoF2jEYT3+x+vC2KJddEGdk/LU8pRowXmyVMHSV5I= +golang.org/x/sys v0.4.0 h1:Zr2JFtRQNX3BCZ8YtxRE9hNJYC8J6I1MVbMg6owUp18= +golang.org/x/sys v0.4.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/auth/auth.go b/internal/auth/auth.go new file mode 100644 index 0000000..5112ec2 --- /dev/null +++ b/internal/auth/auth.go @@ -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("bdrclone

授权成功,可以关闭此页面。

")) + } + 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() +} diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go new file mode 100644 index 0000000..b69273e --- /dev/null +++ b/internal/auth/auth_test.go @@ -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) + } + } +} diff --git a/internal/baidu/client.go b/internal/baidu/client.go new file mode 100644 index 0000000..3386b3b --- /dev/null +++ b/internal/baidu/client.go @@ -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<= 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) } diff --git a/internal/baidu/client_test.go b/internal/baidu/client_test.go new file mode 100644 index 0000000..d33a17f --- /dev/null +++ b/internal/baidu/client_test.go @@ -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, + } +} diff --git a/internal/baidu/files.go b/internal/baidu/files.go new file mode 100644 index 0000000..db8f661 --- /dev/null +++ b/internal/baidu/files.go @@ -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) +} diff --git a/internal/baidu/types.go b/internal/baidu/types.go new file mode 100644 index 0000000..29b2b39 --- /dev/null +++ b/internal/baidu/types.go @@ -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) +} diff --git a/internal/baidu/upload.go b/internal/baidu/upload.go new file mode 100644 index 0000000..f35783f --- /dev/null +++ b/internal/baidu/upload.go @@ -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< 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 + } +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..859cd45 --- /dev/null +++ b/internal/config/config_test.go @@ -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()) + } + } +} diff --git a/internal/mount/cmount.go b/internal/mount/cmount.go new file mode 100644 index 0000000..4b475b1 --- /dev/null +++ b/internal/mount/cmount.go @@ -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) diff --git a/internal/mount/fs.go b/internal/mount/fs.go new file mode 100644 index 0000000..f2764eb --- /dev/null +++ b/internal/mount/fs.go @@ -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) +) diff --git a/internal/mount/mount.go b/internal/mount/mount.go new file mode 100644 index 0000000..f25ba62 --- /dev/null +++ b/internal/mount/mount.go @@ -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 + } +} diff --git a/internal/mount/mount_unsupported.go b/internal/mount/mount_unsupported.go new file mode 100644 index 0000000..8eeb0a4 --- /dev/null +++ b/internal/mount/mount_unsupported.go @@ -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`") +} diff --git a/internal/mount/options.go b/internal/mount/options.go new file mode 100644 index 0000000..7a79559 --- /dev/null +++ b/internal/mount/options.go @@ -0,0 +1,6 @@ +package mount + +type Options struct { + ReadOnly bool + Name string +} diff --git a/internal/mount/recovery.go b/internal/mount/recovery.go new file mode 100644 index 0000000..4bdc6f5 --- /dev/null +++ b/internal/mount/recovery.go @@ -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 +} diff --git a/internal/serve/http.go b/internal/serve/http.go new file mode 100644 index 0000000..d7740e7 --- /dev/null +++ b/internal/serve/http.go @@ -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 +} diff --git a/internal/serve/http_test.go b/internal/serve/http_test.go new file mode 100644 index 0000000..7551893 --- /dev/null +++ b/internal/serve/http_test.go @@ -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) + } + } +}