diff --git a/.gitignore b/.gitignore index 06129e0a49..ae123dc183 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,7 @@ .DS_Store output/ /dist/ +OpenList # Binaries for programs and plugins *.exe @@ -33,4 +34,4 @@ output/ !/public/dist/README.md .VSCodeCounter -*.syso \ No newline at end of file +*.syso diff --git a/build.sh b/build.sh old mode 100644 new mode 100755 diff --git a/drivers/115_open/driver.go b/drivers/115_open/driver.go index 65fc5eefbd..4a538a39e1 100644 --- a/drivers/115_open/driver.go +++ b/drivers/115_open/driver.go @@ -343,6 +343,9 @@ func (d *Open115) Put(ctx context.Context, dstDir model.Obj, file model.FileStre return err } if resp.Status == 2 { + if err := d.finishUpload(ctx, dstDir, file, resp.FileID); err != nil { + return err + } up(100) return nil } @@ -378,6 +381,9 @@ func (d *Open115) Put(ctx context.Context, dstDir model.Obj, file model.FileStre return err } if resp.Status == 2 { + if err := d.finishUpload(ctx, dstDir, file, resp.FileID); err != nil { + return err + } up(100) return nil } @@ -388,10 +394,42 @@ func (d *Open115) Put(ctx context.Context, dstDir model.Obj, file model.FileStre return err } // 4. upload - err = d.multpartUpload(ctx, file, up, tokenResp, resp) + callbackResp, err := d.multpartUpload(ctx, file, up, tokenResp, resp) if err != nil { return err } + if err := d.finishUpload(ctx, dstDir, file, callbackResp.Data.FileID); err != nil { + return err + } + up(100) + return nil +} + +func (d *Open115) finishUpload(ctx context.Context, dstDir model.Obj, file model.FileStreamer, fileID string) error { + if fileID == "" { + return errors.New("upload returned empty file id") + } + + // The 115 Open upload API always creates a file and has no overwrite flag. + // Keep the old file until the new file ID is confirmed, then remove the old + // file by ID so overwrites do not depend on path lookup or rename semantics. + existing := file.GetExist() + if existing == nil || existing.GetID() == fileID { + return nil + } + if existing.GetID() == "" { + return errors.New("existing file has empty file id") + } + if err := d.WaitLimit(ctx); err != nil { + return err + } + _, err := d.client.DelFile(ctx, &sdk.DelFileReq{ + FileIDs: existing.GetID(), + ParentID: dstDir.GetID(), + }) + if err != nil { + return fmt.Errorf("failed to remove overwritten file %q: %w", existing.GetName(), err) + } return nil } diff --git a/drivers/115_open/driver_test.go b/drivers/115_open/driver_test.go new file mode 100644 index 0000000000..20b14ef986 --- /dev/null +++ b/drivers/115_open/driver_test.go @@ -0,0 +1,257 @@ +package _115_open + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "slices" + "strings" + "testing" + + sdk "github.com/OpenListTeam/115-sdk-go" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/stream" + "github.com/OpenListTeam/OpenList/v4/pkg/utils" +) + +func TestPutRapidUploadOverwritesExistingFileByID(t *testing.T) { + var calls []string + var deleteForm url.Values + d, closeServer := newTestOpen115(t, func(w http.ResponseWriter, r *http.Request) { + calls = append(calls, r.URL.Path) + switch r.URL.Path { + case "/open/upload/init": + writeJSON(t, w, `{"state":true,"data":{"status":2,"file_id":"new-id","pick_code":"new-pick"}}`) + case "/open/ufile/delete": + if err := r.ParseForm(); err != nil { + t.Errorf("parse delete form: %v", err) + } + deleteForm = r.Form + writeJSON(t, w, `{"state":true,"data":[]}`) + default: + http.NotFound(w, r) + } + }) + defer closeServer() + + file := newTestUploadStream() + defer file.Close() + file.SetExist(&Obj{Fid: "old-id", Pid: "stale-parent", Fn: file.GetName(), Fc: "1"}) + dstDir := &Obj{Fid: "parent-id", Fn: "parent", Fc: "0"} + progress := 0.0 + + err := d.Put(context.Background(), dstDir, file, func(p float64) { progress = p }) + if err != nil { + t.Fatalf("Put failed: %v", err) + } + if got, want := calls, []string{"/open/upload/init", "/open/ufile/delete"}; !slices.Equal(got, want) { + t.Fatalf("unexpected request order: got %v, want %v", got, want) + } + if got := deleteForm.Get("file_ids"); got != "old-id" { + t.Errorf("deleted file id = %q, want old-id", got) + } + if got := deleteForm.Get("parent_id"); got != "parent-id" { + t.Errorf("delete parent id = %q, want parent-id", got) + } + if progress != 100 { + t.Errorf("progress = %v, want 100", progress) + } +} + +func TestPutRapidUploadWithoutExistingFileDoesNotDelete(t *testing.T) { + var calls []string + d, closeServer := newTestOpen115(t, func(w http.ResponseWriter, r *http.Request) { + calls = append(calls, r.URL.Path) + if r.URL.Path != "/open/upload/init" { + http.NotFound(w, r) + return + } + writeJSON(t, w, `{"state":true,"data":{"status":2,"file_id":"new-id","pick_code":"new-pick"}}`) + }) + defer closeServer() + + file := newTestUploadStream() + defer file.Close() + dstDir := &Obj{Fid: "parent-id", Fn: "parent", Fc: "0"} + + if err := d.Put(context.Background(), dstDir, file, func(float64) {}); err != nil { + t.Fatalf("Put failed: %v", err) + } + if got, want := calls, []string{"/open/upload/init"}; !slices.Equal(got, want) { + t.Fatalf("unexpected requests: got %v, want %v", got, want) + } +} + +func TestPutSecondVerificationOverwritesExistingFile(t *testing.T) { + var calls []string + initCalls := 0 + d, closeServer := newTestOpen115(t, func(w http.ResponseWriter, r *http.Request) { + calls = append(calls, r.URL.Path) + switch r.URL.Path { + case "/open/upload/init": + initCalls++ + if initCalls == 1 { + writeJSON(t, w, `{"state":true,"data":{"status":7,"sign_key":"challenge","sign_check":"0-2"}}`) + return + } + if err := r.ParseForm(); err != nil { + t.Errorf("parse verification form: %v", err) + } + if got := r.Form.Get("sign_key"); got != "challenge" { + t.Errorf("sign_key = %q, want challenge", got) + } + wantSignVal := strings.ToUpper(utils.HashData(utils.SHA1, []byte("rep"))) + if got := r.Form.Get("sign_val"); got != wantSignVal { + t.Errorf("sign_val = %q, want %q", got, wantSignVal) + } + writeJSON(t, w, `{"state":true,"data":{"status":2,"file_id":"new-id","pick_code":"new-pick"}}`) + case "/open/ufile/delete": + writeJSON(t, w, `{"state":true,"data":[]}`) + default: + http.NotFound(w, r) + } + }) + defer closeServer() + + file := newTestUploadStream() + defer file.Close() + file.SetExist(&Obj{Fid: "old-id", Pid: "parent-id", Fn: file.GetName(), Fc: "1"}) + dstDir := &Obj{Fid: "parent-id", Fn: "parent", Fc: "0"} + + if err := d.Put(context.Background(), dstDir, file, func(float64) {}); err != nil { + t.Fatalf("Put failed: %v", err) + } + want := []string{"/open/upload/init", "/open/upload/init", "/open/ufile/delete"} + if !slices.Equal(calls, want) { + t.Fatalf("unexpected requests: got %v, want %v", calls, want) + } +} + +func TestPutDoesNotRemoveExistingFileWhenUploadFails(t *testing.T) { + var calls []string + d, closeServer := newTestOpen115(t, func(w http.ResponseWriter, r *http.Request) { + calls = append(calls, r.URL.Path) + writeJSON(t, w, `{"state":false,"code":500,"message":"upload failed"}`) + }) + defer closeServer() + + file := newTestUploadStream() + defer file.Close() + file.SetExist(&Obj{Fid: "old-id", Pid: "parent-id", Fn: file.GetName(), Fc: "1"}) + dstDir := &Obj{Fid: "parent-id", Fn: "parent", Fc: "0"} + + if err := d.Put(context.Background(), dstDir, file, func(float64) {}); err == nil { + t.Fatal("Put succeeded, want upload error") + } + if got, want := calls, []string{"/open/upload/init"}; !slices.Equal(got, want) { + t.Fatalf("unexpected requests: got %v, want %v", got, want) + } +} + +func TestPutReturnsErrorWhenOverwrittenFileCannotBeRemoved(t *testing.T) { + var calls []string + d, closeServer := newTestOpen115(t, func(w http.ResponseWriter, r *http.Request) { + calls = append(calls, r.URL.Path) + switch r.URL.Path { + case "/open/upload/init": + writeJSON(t, w, `{"state":true,"data":{"status":2,"file_id":"new-id","pick_code":"new-pick"}}`) + case "/open/ufile/delete": + writeJSON(t, w, `{"state":false,"code":500,"message":"delete failed"}`) + default: + http.NotFound(w, r) + } + }) + defer closeServer() + + file := newTestUploadStream() + defer file.Close() + file.SetExist(&Obj{Fid: "old-id", Pid: "parent-id", Fn: file.GetName(), Fc: "1"}) + dstDir := &Obj{Fid: "parent-id", Fn: "parent", Fc: "0"} + + err := d.Put(context.Background(), dstDir, file, func(float64) {}) + if err == nil || !strings.Contains(err.Error(), "failed to remove overwritten file") { + t.Fatalf("Put error = %v, want overwritten-file removal error", err) + } + if got, want := calls, []string{"/open/upload/init", "/open/ufile/delete"}; !slices.Equal(got, want) { + t.Fatalf("unexpected requests: got %v, want %v", got, want) + } +} + +func TestParseCallbackResult(t *testing.T) { + t.Run("success", func(t *testing.T) { + result, err := parseCallbackResult([]byte(`{"state":true,"data":{"file_id":"new-id","file_name":"file.bin"}}`)) + if err != nil { + t.Fatalf("parseCallbackResult failed: %v", err) + } + if result.Data.FileID != "new-id" { + t.Fatalf("file id = %q, want new-id", result.Data.FileID) + } + }) + + for _, tc := range []struct { + name string + body string + }{ + {name: "callback error", body: `{"state":false,"code":500,"message":"failed"}`}, + {name: "missing file id", body: `{"state":true,"data":{}}`}, + {name: "invalid json", body: `{`}, + } { + t.Run(tc.name, func(t *testing.T) { + if _, err := parseCallbackResult([]byte(tc.body)); err == nil { + t.Fatal("parseCallbackResult succeeded, want error") + } + }) + } +} + +func newTestUploadStream() *stream.FileStream { + data := []byte("replacement content") + return &stream.FileStream{ + Obj: &model.Object{ + Name: "same.bin", + Size: int64(len(data)), + HashInfo: utils.NewHashInfo(utils.SHA1, utils.HashData(utils.SHA1, data)), + }, + Reader: bytes.NewReader(data), + } +} + +func newTestOpen115(t *testing.T, handler http.HandlerFunc) (*Open115, func()) { + t.Helper() + server := httptest.NewServer(handler) + target, err := url.Parse(server.URL) + if err != nil { + server.Close() + t.Fatalf("parse test server URL: %v", err) + } + client := sdk.New().SetAccessToken("test-token") + client.SetHttpClient(&http.Client{Transport: &rewriteTransport{ + target: target, + base: server.Client().Transport, + }}) + return &Open115{client: client}, server.Close +} + +type rewriteTransport struct { + target *url.URL + base http.RoundTripper +} + +func (t *rewriteTransport) RoundTrip(req *http.Request) (*http.Response, error) { + clone := req.Clone(req.Context()) + clone.URL.Scheme = t.target.Scheme + clone.URL.Host = t.target.Host + clone.Host = t.target.Host + return t.base.RoundTrip(clone) +} + +func writeJSON(t *testing.T, w http.ResponseWriter, body string) { + t.Helper() + w.Header().Set("Content-Type", "application/json") + if _, err := io.WriteString(w, body); err != nil { + t.Errorf("write response: %v", err) + } +} diff --git a/drivers/115_open/upload.go b/drivers/115_open/upload.go index d02640e2c4..8ee0a52862 100644 --- a/drivers/115_open/upload.go +++ b/drivers/115_open/upload.go @@ -3,6 +3,8 @@ package _115_open import ( "context" "encoding/base64" + "encoding/json" + "fmt" "io" "time" @@ -54,42 +56,56 @@ func (d *Open115) singleUpload(ctx context.Context, tempF model.File, tokenResp return err } -// type CallbackResult struct { -// State bool `json:"state"` -// Code int `json:"code"` -// Message string `json:"message"` -// Data struct { -// PickCode string `json:"pick_code"` -// FileName string `json:"file_name"` -// FileSize int64 `json:"file_size"` -// FileID string `json:"file_id"` -// ThumbURL string `json:"thumb_url"` -// Sha1 string `json:"sha1"` -// Aid int `json:"aid"` -// Cid string `json:"cid"` -// } `json:"data"` -// } - -func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, up driver.UpdateProgress, tokenResp *sdk.UploadGetTokenResp, initResp *sdk.UploadInitResp) error { +type callbackResult struct { + State bool `json:"state"` + Code int `json:"code"` + Message string `json:"message"` + Data struct { + PickCode string `json:"pick_code"` + FileName string `json:"file_name"` + FileSize int64 `json:"file_size"` + FileID string `json:"file_id"` + ThumbURL string `json:"thumb_url"` + Sha1 string `json:"sha1"` + Aid int `json:"aid"` + Cid string `json:"cid"` + } `json:"data"` +} + +func parseCallbackResult(body []byte) (*callbackResult, error) { + var result callbackResult + if err := json.Unmarshal(body, &result); err != nil { + return nil, fmt.Errorf("failed to parse upload callback: %w", err) + } + if !result.State { + return nil, fmt.Errorf("upload callback failed: code=%d, message=%s", result.Code, result.Message) + } + if result.Data.FileID == "" { + return nil, fmt.Errorf("upload callback returned empty file id") + } + return &result, nil +} + +func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, up driver.UpdateProgress, tokenResp *sdk.UploadGetTokenResp, initResp *sdk.UploadInitResp) (*callbackResult, error) { ossClient, err := netutil.NewOSSClient(tokenResp.Endpoint, tokenResp.AccessKeyId, tokenResp.AccessKeySecret, oss.SecurityToken(tokenResp.SecurityToken)) if err != nil { - return err + return nil, err } bucket, err := ossClient.Bucket(initResp.Bucket) if err != nil { - return err + return nil, err } imur, err := bucket.InitiateMultipartUpload(initResp.Object, oss.Sequential()) if err != nil { - return err + return nil, err } fileSize := stream.GetSize() chunkSize := calPartSize(fileSize) ss, err := streamPkg.NewStreamSectionReader(stream, int(chunkSize), &up) if err != nil { - return err + return nil, err } partNum := (stream.GetSize() + chunkSize - 1) / chunkSize @@ -97,7 +113,7 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, offset := int64(0) for i := int64(1); i <= partNum; i++ { if utils.IsCanceled(ctx) { - return ctx.Err() + return nil, ctx.Err() } partSize := chunkSize @@ -106,7 +122,7 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, } rd, err := ss.GetSectionReader(offset, partSize) if err != nil { - return err + return nil, err } err = retry.Do(func() error { rd.Seek(0, io.SeekStart) @@ -123,7 +139,7 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, retry.Delay(time.Second)) ss.FreeSectionReader(rd) if err != nil { - return err + return nil, err } if i == partNum { @@ -134,17 +150,17 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, up(float64(offset) * 100 / float64(fileSize)) } - // callbackRespBytes := make([]byte, 1024) + var callbackRespBytes []byte _, err = bucket.CompleteMultipartUpload( imur, parts, oss.Callback(base64.StdEncoding.EncodeToString([]byte(initResp.Callback.Value.Callback))), oss.CallbackVar(base64.StdEncoding.EncodeToString([]byte(initResp.Callback.Value.CallbackVar))), - // oss.CallbackResult(&callbackRespBytes), + oss.CallbackResult(&callbackRespBytes), ) if err != nil { - return err + return nil, err } - return nil + return parseCallbackResult(callbackRespBytes) } diff --git a/internal/bootstrap/config.go b/internal/bootstrap/config.go index 8304468080..207f53071c 100644 --- a/internal/bootstrap/config.go +++ b/internal/bootstrap/config.go @@ -11,6 +11,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/drivers/base" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/net" + "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/caarlos0/env/v9" "github.com/shirou/gopsutil/v4/mem" @@ -194,12 +195,22 @@ func initURL() { } func CleanTempDir() { + if tool.CleanupTaskManager == nil { + log.Warn("offline cleanup manager is not initialized, skip temp cleanup") + return + } files, err := os.ReadDir(conf.Conf.TempDir) if err != nil { log.Errorln("failed list temp file: ", err) + return } for _, file := range files { - if err := os.RemoveAll(filepath.Join(conf.Conf.TempDir, file.Name())); err != nil { + entryPath := filepath.Join(conf.Conf.TempDir, file.Name()) + if tool.CleanupTaskManager.Protects(entryPath) { + log.Infof("skip protected temp path: %s", entryPath) + continue + } + if err := os.RemoveAll(entryPath); err != nil { log.Errorln("failed delete temp file: ", err) } } diff --git a/internal/bootstrap/data/task.go b/internal/bootstrap/data/task.go index 9898faebd5..c1658f9b76 100644 --- a/internal/bootstrap/data/task.go +++ b/internal/bootstrap/data/task.go @@ -25,6 +25,7 @@ func InitialTasks() []model.TaskItem { {Key: "move", PersistData: "[]"}, {Key: "download", PersistData: "[]"}, {Key: "transfer", PersistData: "[]"}, + {Key: "offline_cleanup", PersistData: "[]"}, } return initialTaskItems } diff --git a/internal/bootstrap/run.go b/internal/bootstrap/run.go index 6740dba657..edb712d80a 100644 --- a/internal/bootstrap/run.go +++ b/internal/bootstrap/run.go @@ -15,6 +15,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/db" "github.com/OpenListTeam/OpenList/v4/internal/fs" + "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/OpenListTeam/OpenList/v4/server" "github.com/OpenListTeam/OpenList/v4/server/middlewares" @@ -40,6 +41,9 @@ func Init() { } func Release() { + if tool.CleanupTaskManager != nil { + tool.CleanupTaskManager.Close() + } db.Close() } diff --git a/internal/bootstrap/task.go b/internal/bootstrap/task.go index 47e0b59ebf..9a716dc1a4 100644 --- a/internal/bootstrap/task.go +++ b/internal/bootstrap/task.go @@ -7,6 +7,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/setting" + "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/OpenListTeam/tache" ) @@ -30,17 +31,27 @@ func InitTaskManager() { op.RegisterSettingChangingCallback(func() { fs.MoveTaskManager.SetWorkersNumActive(taskFilterNegative(setting.GetInt(conf.TaskMoveThreadsNum, conf.Conf.Tasks.Move.Workers))) }) - tool.DownloadTaskManager = tache.NewManager[*tool.DownloadTask](tache.WithWorks(setting.GetInt(conf.TaskOfflineDownloadThreadsNum, conf.Conf.Tasks.Download.Workers)), tache.WithPersistFunction(db.GetTaskDataFunc("download", conf.Conf.Tasks.Download.TaskPersistant), db.UpdateTaskDataFunc("download", conf.Conf.Tasks.Download.TaskPersistant)), tache.WithMaxRetry(conf.Conf.Tasks.Download.MaxRetry)) + tool.DownloadTaskManager = tache.NewManager[*tool.DownloadTask](tache.WithWorks(setting.GetInt(conf.TaskOfflineDownloadThreadsNum, conf.Conf.Tasks.Download.Workers)), tache.WithPersistFunction(db.GetTaskDataFunc("download", conf.Conf.Tasks.Download.TaskPersistant), db.UpdateTaskDataFunc("download", conf.Conf.Tasks.Download.TaskPersistant)), tache.WithMaxRetry(conf.Conf.Tasks.Download.MaxRetry), tache.WithRunning(false)) op.RegisterSettingChangingCallback(func() { tool.DownloadTaskManager.SetWorkersNumActive(taskFilterNegative(setting.GetInt(conf.TaskOfflineDownloadThreadsNum, conf.Conf.Tasks.Download.Workers))) }) - tool.TransferTaskManager = tache.NewManager[*tool.TransferTask](tache.WithWorks(setting.GetInt(conf.TaskOfflineDownloadTransferThreadsNum, conf.Conf.Tasks.Transfer.Workers)), tache.WithPersistFunction(db.GetTaskDataFunc("transfer", conf.Conf.Tasks.Transfer.TaskPersistant), db.UpdateTaskDataFunc("transfer", conf.Conf.Tasks.Transfer.TaskPersistant)), tache.WithMaxRetry(conf.Conf.Tasks.Transfer.MaxRetry)) + tool.TransferTaskManager = tache.NewManager[*tool.TransferTask](tache.WithWorks(setting.GetInt(conf.TaskOfflineDownloadTransferThreadsNum, conf.Conf.Tasks.Transfer.Workers)), tache.WithPersistFunction(db.GetTaskDataFunc("transfer", conf.Conf.Tasks.Transfer.TaskPersistant), db.UpdateTaskDataFunc("transfer", conf.Conf.Tasks.Transfer.TaskPersistant)), tache.WithMaxRetry(conf.Conf.Tasks.Transfer.MaxRetry), tache.WithRunning(false)) op.RegisterSettingChangingCallback(func() { tool.TransferTaskManager.SetWorkersNumActive(taskFilterNegative(setting.GetInt(conf.TaskOfflineDownloadTransferThreadsNum, conf.Conf.Tasks.Transfer.Workers))) }) - if len(tool.TransferTaskManager.GetAll()) == 0 { //prevent offline downloaded files from being deleted - CleanTempDir() + if err := tool.InitCleanupManager(db.GetTaskDataFunc("offline_cleanup", true), db.UpdateTaskDataFunc("offline_cleanup", true)); err != nil { + utils.Log.Fatalf("failed to initialize offline cleanup manager: %v", err) } + if err := tool.CleanupTaskManager.ReconcileDownloads(tool.DownloadTaskManager.GetAll()); err != nil { + utils.Log.Fatalf("failed to reconcile offline download cleanup jobs: %v", err) + } + if err := tool.CleanupTaskManager.ReconcileTransfers(tool.TransferTaskManager.GetAll()); err != nil { + utils.Log.Fatalf("failed to reconcile offline transfer cleanup jobs: %v", err) + } + CleanTempDir() + tool.CleanupTaskManager.Start() + tool.TransferTaskManager.Start() + tool.DownloadTaskManager.Start() fs.ArchiveDownloadTaskManager = tache.NewManager[*fs.ArchiveDownloadTask](tache.WithWorks(setting.GetInt(conf.TaskDecompressDownloadThreadsNum, conf.Conf.Tasks.Decompress.Workers)), tache.WithPersistFunction(db.GetTaskDataFunc("decompress", conf.Conf.Tasks.Decompress.TaskPersistant), db.UpdateTaskDataFunc("decompress", conf.Conf.Tasks.Decompress.TaskPersistant)), tache.WithMaxRetry(conf.Conf.Tasks.Decompress.MaxRetry)) op.RegisterSettingChangingCallback(func() { fs.ArchiveDownloadTaskManager.SetWorkersNumActive(taskFilterNegative(setting.GetInt(conf.TaskDecompressDownloadThreadsNum, conf.Conf.Tasks.Decompress.Workers))) diff --git a/internal/offline_download/qbit/fetch.go b/internal/offline_download/qbit/fetch.go new file mode 100644 index 0000000000..9368b7ddc4 --- /dev/null +++ b/internal/offline_download/qbit/fetch.go @@ -0,0 +1,207 @@ +package qbit + +import ( + "context" + "crypto/tls" + "fmt" + "io" + "net" + "net/http" + "net/netip" + "net/url" + "strings" + "time" + + "github.com/OpenListTeam/OpenList/v4/drivers/base" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" + "github.com/OpenListTeam/OpenList/v4/pkg/torrent" +) + +const ( + maxQbittorrentTorrentSize = 10 * 1024 * 1024 + torrentFetchTimeout = 30 * time.Second + maxTorrentRedirects = 10 +) + +var nonPublicTorrentNetworks = []netip.Prefix{ + netip.MustParsePrefix("100.64.0.0/10"), // Shared address space and common private overlays. + netip.MustParsePrefix("192.0.0.0/24"), // IETF protocol assignments. + netip.MustParsePrefix("192.0.2.0/24"), // Documentation. + netip.MustParsePrefix("192.88.99.0/24"), // Deprecated 6to4 relay anycast. + netip.MustParsePrefix("198.18.0.0/15"), // Benchmarking networks, often routed internally. + netip.MustParsePrefix("198.51.100.0/24"), + netip.MustParsePrefix("203.0.113.0/24"), + netip.MustParsePrefix("240.0.0.0/4"), + netip.MustParsePrefix("64:ff9b::/96"), // NAT64 can translate to private IPv4 destinations. + netip.MustParsePrefix("64:ff9b:1::/48"), + netip.MustParsePrefix("2001::/32"), // Teredo embeds IPv4 destinations. + netip.MustParsePrefix("2001:db8::/32"), + netip.MustParsePrefix("2002::/16"), // 6to4 embeds IPv4 destinations. +} + +type torrentResolver interface { + LookupNetIP(context.Context, string, string) ([]netip.Addr, error) +} + +type torrentDialer interface { + DialContext(context.Context, string, string) (net.Conn, error) +} + +func isRemoteTorrentURL(rawURL string) bool { + u, err := url.Parse(strings.TrimSpace(rawURL)) + return err == nil && (strings.EqualFold(u.Scheme, "http") || strings.EqualFold(u.Scheme, "https")) +} + +func isMagnetURL(rawURL string) bool { + u, err := url.Parse(strings.TrimSpace(rawURL)) + return err == nil && strings.EqualFold(u.Scheme, "magnet") +} + +func fetchTorrentDataFromURL(args *tool.AddUrlArgs) ([]byte, error) { + return fetchTorrentData(args.Ctx, strings.TrimSpace(args.Url), newPublicTorrentHTTPClient()) +} + +func fetchTorrentData(ctx context.Context, rawURL string, client *http.Client) ([]byte, error) { + u, err := url.Parse(rawURL) + if err != nil { + return nil, fmt.Errorf("invalid torrent URL: %w", err) + } + if err := validateRemoteTorrentURL(u); err != nil { + return nil, err + } + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil) + if err != nil { + return nil, fmt.Errorf("failed to create torrent request: %w", err) + } + req.Header.Set("User-Agent", base.UserAgent) + + resp, err := client.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to fetch torrent URL: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return nil, fmt.Errorf("failed to fetch torrent URL: unexpected HTTP status %s", resp.Status) + } + if resp.ContentLength > maxQbittorrentTorrentSize { + return nil, fmt.Errorf("torrent data is too large, maximum size is 10MB") + } + + data, err := io.ReadAll(io.LimitReader(resp.Body, maxQbittorrentTorrentSize+1)) + if err != nil { + return nil, fmt.Errorf("failed to read torrent URL: %w", err) + } + if len(data) > maxQbittorrentTorrentSize { + return nil, fmt.Errorf("torrent data is too large, maximum size is 10MB") + } + if _, err := torrent.Decode(data); err != nil { + return nil, fmt.Errorf("remote URL did not return a valid torrent: %w", err) + } + return data, nil +} + +func newPublicTorrentHTTPClient() *http.Client { + tlsInsecureSkipVerify := conf.Conf != nil && conf.Conf.TlsInsecureSkipVerify + dialer := &net.Dialer{ + Timeout: 10 * time.Second, + KeepAlive: 30 * time.Second, + } + transport := &http.Transport{ + // A proxy could resolve the hostname again and bypass the pinned public IP. + Proxy: nil, + DialContext: publicTorrentDialContext(net.DefaultResolver, dialer), + TLSClientConfig: &tls.Config{InsecureSkipVerify: tlsInsecureSkipVerify}, + TLSHandshakeTimeout: 10 * time.Second, + ResponseHeaderTimeout: 15 * time.Second, + IdleConnTimeout: 30 * time.Second, + DisableKeepAlives: true, + } + client := &http.Client{ + Transport: transport, + Timeout: torrentFetchTimeout, + } + client.CheckRedirect = func(req *http.Request, via []*http.Request) error { + if len(via) >= maxTorrentRedirects { + return fmt.Errorf("stopped after %d redirects", maxTorrentRedirects) + } + return validateRemoteTorrentURL(req.URL) + } + return client +} + +func publicTorrentDialContext(resolver torrentResolver, dialer torrentDialer) func(context.Context, string, string) (net.Conn, error) { + return func(ctx context.Context, network, address string) (net.Conn, error) { + host, port, err := net.SplitHostPort(address) + if err != nil { + return nil, fmt.Errorf("invalid torrent host address %q: %w", address, err) + } + addrs, err := resolveTorrentHost(ctx, resolver, host) + if err != nil { + return nil, err + } + for _, addr := range addrs { + if !isPublicTorrentAddress(addr) { + return nil, fmt.Errorf("torrent URL resolves to a non-public address") + } + } + + var lastErr error + for _, addr := range addrs { + conn, err := dialer.DialContext(ctx, network, net.JoinHostPort(addr.String(), port)) + if err == nil { + return conn, nil + } + lastErr = err + } + return nil, fmt.Errorf("failed to connect to torrent host: %w", lastErr) + } +} + +func resolveTorrentHost(ctx context.Context, resolver torrentResolver, host string) ([]netip.Addr, error) { + if addr, err := netip.ParseAddr(host); err == nil { + return []netip.Addr{addr}, nil + } + addrs, err := resolver.LookupNetIP(ctx, "ip", host) + if err != nil { + return nil, fmt.Errorf("failed to resolve torrent host: %w", err) + } + if len(addrs) == 0 { + return nil, fmt.Errorf("torrent host resolved to no addresses") + } + return addrs, nil +} + +func validateRemoteTorrentURL(u *url.URL) error { + if u == nil || (!strings.EqualFold(u.Scheme, "http") && !strings.EqualFold(u.Scheme, "https")) { + return fmt.Errorf("torrent URL must use HTTP or HTTPS") + } + host := strings.TrimSuffix(strings.ToLower(u.Hostname()), ".") + if host == "" { + return fmt.Errorf("torrent URL host is required") + } + if host == "localhost" || strings.HasSuffix(host, ".localhost") { + return fmt.Errorf("torrent URL must not target localhost") + } + if addr, err := netip.ParseAddr(host); err == nil && !isPublicTorrentAddress(addr) { + return fmt.Errorf("torrent URL must target a public address") + } + return nil +} + +func isPublicTorrentAddress(addr netip.Addr) bool { + if !addr.IsValid() { + return false + } + addr = addr.Unmap() + if !addr.IsGlobalUnicast() || addr.IsPrivate() || addr.IsLoopback() || addr.IsLinkLocalUnicast() || addr.IsLinkLocalMulticast() || addr.IsMulticast() || addr.IsUnspecified() { + return false + } + for _, prefix := range nonPublicTorrentNetworks { + if prefix.Contains(addr) { + return false + } + } + return true +} diff --git a/internal/offline_download/qbit/fetch_test.go b/internal/offline_download/qbit/fetch_test.go new file mode 100644 index 0000000000..526a47c212 --- /dev/null +++ b/internal/offline_download/qbit/fetch_test.go @@ -0,0 +1,220 @@ +package qbit + +import ( + "context" + "errors" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "strings" + "sync/atomic" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" + "github.com/OpenListTeam/OpenList/v4/pkg/qbittorrent" + "github.com/OpenListTeam/OpenList/v4/pkg/torrent" +) + +type recordingClient struct { + qbittorrent.Client + linkCalls int + torrentCalls int + link string + torrentData []byte +} + +func (c *recordingClient) AddFromLink(link, _, _ string) error { + c.linkCalls++ + c.link = link + return nil +} + +func (c *recordingClient) AddFromTorrent(data []byte, _, _ string) error { + c.torrentCalls++ + c.torrentData = append([]byte(nil), data...) + return nil +} + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +type staticResolver struct { + addrs []netip.Addr +} + +func (r staticResolver) LookupNetIP(context.Context, string, string) ([]netip.Addr, error) { + return r.addrs, nil +} + +type recordingDialer struct { + addresses []string +} + +func (d *recordingDialer) DialContext(_ context.Context, _, address string) (net.Conn, error) { + d.addresses = append(d.addresses, address) + return nil, errors.New("test dial stopped") +} + +func TestIsPublicTorrentAddress(t *testing.T) { + tests := []struct { + address string + want bool + }{ + {address: "8.8.8.8", want: true}, + {address: "2606:4700:4700::1111", want: true}, + {address: "0.0.0.0", want: false}, + {address: "127.0.0.1", want: false}, + {address: "10.0.0.1", want: false}, + {address: "172.16.0.1", want: false}, + {address: "192.168.0.1", want: false}, + {address: "169.254.169.254", want: false}, + {address: "100.100.100.200", want: false}, + {address: "198.18.0.1", want: false}, + {address: "::1", want: false}, + {address: "fc00::1", want: false}, + {address: "fe80::1", want: false}, + {address: "::ffff:127.0.0.1", want: false}, + {address: "64:ff9b::7f00:1", want: false}, + } + + for _, tt := range tests { + t.Run(tt.address, func(t *testing.T) { + if got := isPublicTorrentAddress(netip.MustParseAddr(tt.address)); got != tt.want { + t.Fatalf("isPublicTorrentAddress(%s) = %v, want %v", tt.address, got, tt.want) + } + }) + } +} + +func TestQBittorrentCapabilities(t *testing.T) { + capabilities := (&QBittorrent{}).Capabilities() + if !capabilities.TorrentData { + t.Fatal("qBittorrent did not advertise torrent data support") + } +} + +func TestQBittorrentRejectsPrivateTorrentURLWithoutFallback(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + requests.Add(1) + })) + defer server.Close() + + client := &recordingClient{} + qbitTool := New(client) + _, err := qbitTool.AddURL(&tool.AddUrlArgs{ + Ctx: context.Background(), + Url: server.URL + "/file.torrent", + UID: "task-id", + }) + if err == nil || !strings.Contains(err.Error(), "public address") { + t.Fatalf("AddURL() error = %v, want public-address rejection", err) + } + if requests.Load() != 0 { + t.Fatalf("private server received %d requests", requests.Load()) + } + if client.linkCalls != 0 || client.torrentCalls != 0 { + t.Fatalf("qBittorrent calls = links:%d torrents:%d, want none", client.linkCalls, client.torrentCalls) + } +} + +func TestPublicTorrentDialRejectsPrivateDNSResult(t *testing.T) { + dialer := &recordingDialer{} + dial := publicTorrentDialContext(staticResolver{addrs: []netip.Addr{ + netip.MustParseAddr("93.184.216.34"), + netip.MustParseAddr("127.0.0.1"), + }}, dialer) + + if _, err := dial(context.Background(), "tcp", "example.com:80"); err == nil || !strings.Contains(err.Error(), "non-public") { + t.Fatalf("dial error = %v, want non-public rejection", err) + } + if len(dialer.addresses) != 0 { + t.Fatalf("dialer called with %v before all DNS results were validated", dialer.addresses) + } +} + +func TestPublicTorrentDialPinsValidatedIPAddress(t *testing.T) { + dialer := &recordingDialer{} + dial := publicTorrentDialContext(staticResolver{addrs: []netip.Addr{ + netip.MustParseAddr("93.184.216.34"), + }}, dialer) + + _, err := dial(context.Background(), "tcp", "example.com:443") + if err == nil || !strings.Contains(err.Error(), "test dial stopped") { + t.Fatalf("dial error = %v", err) + } + if len(dialer.addresses) != 1 || dialer.addresses[0] != "93.184.216.34:443" { + t.Fatalf("dial addresses = %v, want pinned IP", dialer.addresses) + } +} + +func TestQBittorrentKeepsMagnetLinkFlow(t *testing.T) { + client := &recordingClient{} + qbitTool := New(client) + const magnet = "magnet:?xt=urn:btih:test" + + id, err := qbitTool.AddURL(&tool.AddUrlArgs{Ctx: context.Background(), Url: magnet, UID: "task-id"}) + if err != nil { + t.Fatal(err) + } + if id != "task-id" || client.linkCalls != 1 || client.link != magnet || client.torrentCalls != 0 { + t.Fatalf("unexpected magnet result: id=%q links=%d link=%q torrents=%d", id, client.linkCalls, client.link, client.torrentCalls) + } +} + +func TestQBittorrentRejectsUnsupportedURLScheme(t *testing.T) { + client := &recordingClient{} + qbitTool := New(client) + + _, err := qbitTool.AddURL(&tool.AddUrlArgs{Ctx: context.Background(), Url: "ftp://127.0.0.1/file.torrent", UID: "task-id"}) + if err == nil || !strings.Contains(err.Error(), "only supports") { + t.Fatalf("AddURL() error = %v, want unsupported-scheme rejection", err) + } + if client.linkCalls != 0 || client.torrentCalls != 0 { + t.Fatalf("qBittorrent calls = links:%d torrents:%d, want none", client.linkCalls, client.torrentCalls) + } +} + +func TestFetchTorrentDataAcceptsValidPublicResponse(t *testing.T) { + torrentData, err := torrent.NewTorrent("test.bin", 1, "").Encode() + if err != nil { + t.Fatal(err) + } + client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if req.URL.Host != "example.com" { + t.Fatalf("request host = %q", req.URL.Host) + } + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Body: io.NopCloser(strings.NewReader(string(torrentData))), + ContentLength: int64(len(torrentData)), + Header: make(http.Header), + Request: req, + }, nil + })} + + got, err := fetchTorrentData(context.Background(), "https://example.com/file.torrent", client) + if err != nil { + t.Fatal(err) + } + if string(got) != string(torrentData) { + t.Fatal("fetched torrent data changed") + } +} + +func TestTorrentRedirectRejectsLocalhost(t *testing.T) { + client := newPublicTorrentHTTPClient() + req, err := http.NewRequest(http.MethodGet, "http://localhost/internal", nil) + if err != nil { + t.Fatal(err) + } + if err := client.CheckRedirect(req, nil); err == nil { + t.Fatal("localhost redirect was accepted") + } +} diff --git a/internal/offline_download/qbit/qbit.go b/internal/offline_download/qbit/qbit.go index aac3a2c72b..5a3d0f076a 100644 --- a/internal/offline_download/qbit/qbit.go +++ b/internal/offline_download/qbit/qbit.go @@ -14,6 +14,10 @@ type QBittorrent struct { client qbittorrent.Client } +func New(client qbittorrent.Client) *QBittorrent { + return &QBittorrent{client: client} +} + func (a *QBittorrent) Run(task *tool.DownloadTask) error { return errs.NotSupport } @@ -22,6 +26,10 @@ func (a *QBittorrent) Name() string { return "qBittorrent" } +func (*QBittorrent) Capabilities() tool.Capabilities { + return tool.Capabilities{TorrentData: true} +} + func (a *QBittorrent) Items() []model.SettingItem { // qBittorrent settings return []model.SettingItem{ @@ -46,7 +54,20 @@ func (a *QBittorrent) IsReady() bool { } func (a *QBittorrent) AddURL(args *tool.AddUrlArgs) (string, error) { - err := a.client.AddFromLink(args.Url, args.TempDir, args.UID) + var err error + if len(args.TorrentData) > 0 { + err = a.client.AddFromTorrent(args.TorrentData, args.TempDir, args.UID) + } else if isRemoteTorrentURL(args.Url) { + var torrentData []byte + torrentData, err = fetchTorrentDataFromURL(args) + if err == nil { + err = a.client.AddFromTorrent(torrentData, args.TempDir, args.UID) + } + } else if isMagnetURL(args.Url) { + err = a.client.AddFromLink(args.Url, args.TempDir, args.UID) + } else { + err = errors.New("qBittorrent only supports magnet links, public HTTP(S) torrent URLs, or uploaded torrent files") + } if err != nil { return "", err } @@ -63,9 +84,21 @@ func (a *QBittorrent) Status(task *tool.DownloadTask) (*tool.Status, error) { if err != nil { return nil, err } + if info.SavePath != "" { + task.TempDir = info.SavePath + if task.CleanupID != "" { + if err := tool.CleanupTaskManager.SetTempDir(task.CleanupID, info.SavePath); err != nil { + return nil, err + } + } + } s := &tool.Status{} s.TotalBytes = info.Size - s.Progress = float64(info.Completed) / float64(info.Size) * 100 + if info.Size > 0 { + s.Progress = float64(info.Completed) / float64(info.Size) * 100 + } else { + s.Progress = info.Progress * 100 + } switch info.State { case qbittorrent.UPLOADING, qbittorrent.PAUSEDUP, qbittorrent.QUEUEDUP, qbittorrent.STALLEDUP, qbittorrent.FORCEDUP, qbittorrent.CHECKINGUP: s.Completed = true diff --git a/internal/offline_download/tool/add.go b/internal/offline_download/tool/add.go index 12160dcfee..3ee1058fec 100644 --- a/internal/offline_download/tool/add.go +++ b/internal/offline_download/tool/add.go @@ -37,11 +37,13 @@ const ( DeleteOnUploadFailed DeletePolicy = "delete_on_upload_failed" DeleteNever DeletePolicy = "delete_never" DeleteAlways DeletePolicy = "delete_always" + DeleteAfterSeeding DeletePolicy = "delete_after_seeding" UploadDownloadStream DeletePolicy = "upload_download_stream" ) type AddURLArgs struct { URL string + TorrentData []byte DstDirPath string Tool string DeletePolicy DeletePolicy @@ -70,7 +72,7 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro } } // try putting url - if args.Tool == "SimpleHttp" && !isEd2kURL(args.URL) { + if len(args.TorrentData) == 0 && args.Tool == "SimpleHttp" && !isEd2kURL(args.URL) { if isSimpleHttpSchemeUnsupported(args.URL) { return nil, fmt.Errorf("SimpleHttp tool does not support this URL scheme, please use aria2 or other tools for magnet/ed2k links") } @@ -82,7 +84,7 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro } // ed2k 链接自动路由:如果当前工具不支持 ed2k,自动尝试使用迅雷系工具 - if isEd2kURL(args.URL) { + if len(args.TorrentData) == 0 && isEd2kURL(args.URL) { if !isEd2kCapableTool(args.Tool) { if storageTool := ed2kToolForStorage(storage); storageTool != "" { // Prefer the matching native tool when the destination storage supports ed2k. @@ -104,6 +106,9 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro if err != nil { return nil, errors.Wrapf(err, "failed get offline download tool") } + if len(args.TorrentData) > 0 && !CapabilitiesOf(tool).TorrentData { + return nil, fmt.Errorf("%s does not support uploaded torrent files", args.Tool) + } // check tool is ready if !tool.IsReady() { // try to init tool @@ -187,12 +192,28 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro ApiUrl: common.GetApiUrl(ctx), }, Url: args.URL, + TorrentData: args.TorrentData, DstDirPath: args.DstDirPath, TempDir: tempDir, DeletePolicy: deletePolicy, Toolname: args.Tool, tool: tool, } + if deletePolicy == DeleteAfterSeeding && isSeedingToolName(args.Tool) { + if CleanupTaskManager == nil { + return nil, fmt.Errorf("offline cleanup manager is not initialized") + } + t.CleanupID = uid + t.SetID(uid) + if err := CleanupTaskManager.Register(CleanupJob{ + ID: uid, + DownloadTaskID: uid, + TempDir: tempDir, + Toolname: args.Tool, + }); err != nil { + return nil, errors.WithMessage(err, "failed to register offline cleanup") + } + } DownloadTaskManager.Add(t) return t, nil } diff --git a/internal/offline_download/tool/base.go b/internal/offline_download/tool/base.go index 823bac5266..41db8a488f 100644 --- a/internal/offline_download/tool/base.go +++ b/internal/offline_download/tool/base.go @@ -7,11 +7,12 @@ import ( ) type AddUrlArgs struct { - Url string - UID string - TempDir string - Signal chan int - Ctx context.Context + Url string + TorrentData []byte + UID string + TempDir string + Signal chan int + Ctx context.Context } type Status struct { @@ -23,6 +24,26 @@ type Status struct { Err error } +// Capabilities describes optional input features supported by an offline download tool. +// The zero value represents a tool with no optional capabilities. +type Capabilities struct { + TorrentData bool +} + +// CapabilityProvider is implemented by tools that support optional capabilities. +type CapabilityProvider interface { + Capabilities() Capabilities +} + +// CapabilitiesOf returns the optional capabilities advertised by a tool. +func CapabilitiesOf(downloadTool Tool) Capabilities { + provider, ok := downloadTool.(CapabilityProvider) + if !ok { + return Capabilities{} + } + return provider.Capabilities() +} + type Tool interface { Name() string // Items return the setting items the tool need diff --git a/internal/offline_download/tool/base_test.go b/internal/offline_download/tool/base_test.go new file mode 100644 index 0000000000..212644dde8 --- /dev/null +++ b/internal/offline_download/tool/base_test.go @@ -0,0 +1,41 @@ +package tool + +import "testing" + +type testCapabilityProvider struct { + Tool + capabilities Capabilities +} + +func (t testCapabilityProvider) Capabilities() Capabilities { + return t.capabilities +} + +func TestCapabilitiesOf(t *testing.T) { + tests := []struct { + name string + downloadTool Tool + want Capabilities + }{ + { + name: "tool without provider", + downloadTool: struct{ Tool }{}, + want: Capabilities{}, + }, + { + name: "tool with torrent data support", + downloadTool: testCapabilityProvider{ + capabilities: Capabilities{TorrentData: true}, + }, + want: Capabilities{TorrentData: true}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := CapabilitiesOf(tt.downloadTool); got != tt.want { + t.Fatalf("CapabilitiesOf() = %+v, want %+v", got, tt.want) + } + }) + } +} diff --git a/internal/offline_download/tool/cleanup.go b/internal/offline_download/tool/cleanup.go new file mode 100644 index 0000000000..8e495d449d --- /dev/null +++ b/internal/offline_download/tool/cleanup.go @@ -0,0 +1,478 @@ +package tool + +import ( + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/task" + "github.com/OpenListTeam/tache" + log "github.com/sirupsen/logrus" +) + +type CleanupPhase string + +const ( + CleanupDownloading CleanupPhase = "downloading" + CleanupTransferring CleanupPhase = "transferring" + CleanupBlocked CleanupPhase = "blocked" + CleanupWaitingSeeding CleanupPhase = "waiting_seeding" + CleanupRunning CleanupPhase = "cleaning" +) + +type CleanupJob struct { + ID string `json:"id"` + DownloadTaskID string `json:"download_task_id,omitempty"` + TempDir string `json:"temp_dir"` + Toolname string `json:"toolname"` + GID string `json:"gid,omitempty"` + DeleteAfterTime time.Time `json:"delete_after_time,omitempty"` + Phase CleanupPhase `json:"phase"` + TransferSetupDone bool `json:"transfer_setup_done"` + PendingTransfers int `json:"pending_transfers"` + FailedTransfers int `json:"failed_transfers"` + LastError string `json:"last_error,omitempty"` +} + +type CleanupExecutor func(context.Context, CleanupJob) error + +type CleanupManager struct { + mu sync.Mutex + jobs map[string]CleanupJob + write func([]byte) error + execute CleanupExecutor + wake chan struct{} + stop chan struct{} + stopped chan struct{} + start sync.Once + close sync.Once +} + +var CleanupTaskManager *CleanupManager + +func NewCleanupManager(read func() ([]byte, error), write func([]byte) error, execute CleanupExecutor) (*CleanupManager, error) { + m := &CleanupManager{ + jobs: make(map[string]CleanupJob), + write: write, + execute: execute, + wake: make(chan struct{}, 1), + stop: make(chan struct{}), + stopped: make(chan struct{}), + } + if m.execute == nil { + m.execute = executeCleanupJob + } + if read == nil || write == nil { + return nil, fmt.Errorf("cleanup persistence is not configured") + } + data, err := read() + if err != nil { + return nil, err + } + if len(data) > 0 && string(data) != "null" { + var jobs []CleanupJob + if err := json.Unmarshal(data, &jobs); err != nil { + return nil, err + } + for _, job := range jobs { + if job.ID == "" { + continue + } + if job.Phase == CleanupRunning { + job.Phase = CleanupWaitingSeeding + } + m.jobs[job.ID] = job + } + } + return m, nil +} + +func InitCleanupManager(read func() ([]byte, error), write func([]byte) error) error { + m, err := NewCleanupManager(read, write, nil) + if err != nil { + return err + } + CleanupTaskManager = m + return nil +} + +func (m *CleanupManager) Start() { + m.start.Do(func() { + go m.loop() + m.notify() + }) +} + +func (m *CleanupManager) Close() { + m.close.Do(func() { + close(m.stop) + <-m.stopped + }) +} + +func (m *CleanupManager) loop() { + defer close(m.stopped) + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + for { + select { + case <-ticker.C: + m.RunDue(context.Background(), time.Now()) + case <-m.wake: + m.RunDue(context.Background(), time.Now()) + case <-m.stop: + return + } + } +} + +func (m *CleanupManager) notify() { + select { + case m.wake <- struct{}{}: + default: + } +} + +func (m *CleanupManager) persistLocked() error { + jobs := make([]CleanupJob, 0, len(m.jobs)) + for _, job := range m.jobs { + jobs = append(jobs, job) + } + sort.Slice(jobs, func(i, j int) bool { return jobs[i].ID < jobs[j].ID }) + data, err := json.Marshal(jobs) + if err != nil { + return err + } + return m.write(data) +} + +func (m *CleanupManager) Register(job CleanupJob) error { + if job.ID == "" || job.TempDir == "" { + return fmt.Errorf("cleanup job id and temp dir are required") + } + m.mu.Lock() + defer m.mu.Unlock() + if _, ok := m.jobs[job.ID]; ok { + return nil + } + job.Phase = CleanupDownloading + m.jobs[job.ID] = job + if err := m.persistLocked(); err != nil { + delete(m.jobs, job.ID) + return err + } + m.notify() + return nil +} + +func (m *CleanupManager) SetDownloadTaskID(id, taskID string) error { + return m.update(id, func(job *CleanupJob) { + job.DownloadTaskID = taskID + }) +} + +func (m *CleanupManager) SetGID(id, gid string) error { + return m.update(id, func(job *CleanupJob) { + job.GID = gid + }) +} + +func (m *CleanupManager) SetTempDir(id, tempDir string) error { + if tempDir == "" { + return nil + } + return m.update(id, func(job *CleanupJob) { + job.TempDir = tempDir + }) +} + +func (m *CleanupManager) BeginTransfer(id, gid string, deleteAfterTime time.Time) error { + return m.update(id, func(job *CleanupJob) { + job.GID = gid + job.DeleteAfterTime = deleteAfterTime + job.Phase = CleanupTransferring + job.TransferSetupDone = false + job.LastError = "" + }) +} + +func (m *CleanupManager) AddTransfer(id string) error { + return m.update(id, func(job *CleanupJob) { + job.PendingTransfers++ + job.Phase = CleanupTransferring + job.LastError = "" + }) +} + +func (m *CleanupManager) FinishTransferSetup(id string) error { + return m.update(id, func(job *CleanupJob) { + job.TransferSetupDone = true + evaluateCleanupJob(job) + }) +} + +func (m *CleanupManager) TransferSucceeded(id string) error { + return m.update(id, func(job *CleanupJob) { + if job.PendingTransfers > 0 { + job.PendingTransfers-- + } + evaluateCleanupJob(job) + }) +} + +func (m *CleanupManager) TransferFailed(id string, taskErr error) error { + return m.update(id, func(job *CleanupJob) { + if job.PendingTransfers > 0 { + job.PendingTransfers-- + } + job.FailedTransfers++ + job.Phase = CleanupBlocked + if taskErr != nil { + job.LastError = taskErr.Error() + } + }) +} + +func (m *CleanupManager) RetryTransfer(id string) error { + return m.update(id, func(job *CleanupJob) { + if job.FailedTransfers > 0 { + job.FailedTransfers-- + } + job.PendingTransfers++ + job.Phase = CleanupTransferring + job.LastError = "" + }) +} + +func (m *CleanupManager) DownloadFailed(id string, taskErr error) error { + return m.update(id, func(job *CleanupJob) { + job.Phase = CleanupBlocked + if taskErr != nil { + job.LastError = taskErr.Error() + } + }) +} + +func (m *CleanupManager) RetryDownload(id string) error { + return m.update(id, func(job *CleanupJob) { + job.Phase = CleanupDownloading + job.LastError = "" + }) +} + +func (m *CleanupManager) ReconcileDownloads(tasks []*DownloadTask) error { + m.mu.Lock() + defer m.mu.Unlock() + changed := false + for _, downloadTask := range tasks { + if downloadTask.CleanupID == "" { + continue + } + job, ok := m.jobs[downloadTask.CleanupID] + if !ok { + continue + } + previous := job + job.DownloadTaskID = downloadTask.GetID() + job.TempDir = downloadTask.TempDir + job.GID = downloadTask.GID + if !downloadTask.DeleteAfterTime.IsZero() { + job.DeleteAfterTime = downloadTask.DeleteAfterTime + } + if job != previous { + m.jobs[job.ID] = job + changed = true + } + } + if !changed { + return nil + } + return m.persistLocked() +} + +func (m *CleanupManager) ReconcileTransfers(tasks []*TransferTask) error { + type transferState struct { + found bool + pending int + failed int + } + states := make(map[string]transferState) + for _, transferTask := range tasks { + if transferTask.CleanupID == "" { + continue + } + state := states[transferTask.CleanupID] + state.found = true + switch transferTask.GetState() { + case tache.StateSucceeded: + case tache.StateFailed, tache.StateCanceled: + state.failed++ + default: + state.pending++ + } + states[transferTask.CleanupID] = state + } + + m.mu.Lock() + defer m.mu.Unlock() + changed := false + for id, state := range states { + if !state.found { + continue + } + job, ok := m.jobs[id] + if !ok { + continue + } + previous := job + job.PendingTransfers = state.pending + job.FailedTransfers = state.failed + evaluateCleanupJob(&job) + if job != previous { + m.jobs[id] = job + changed = true + } + } + if !changed { + return nil + } + return m.persistLocked() +} + +func evaluateCleanupJob(job *CleanupJob) { + if !job.TransferSetupDone || job.PendingTransfers > 0 { + return + } + if job.FailedTransfers > 0 { + job.Phase = CleanupBlocked + return + } + job.Phase = CleanupWaitingSeeding +} + +func (m *CleanupManager) update(id string, update func(*CleanupJob)) error { + if id == "" { + return nil + } + m.mu.Lock() + job, ok := m.jobs[id] + if !ok { + m.mu.Unlock() + return fmt.Errorf("cleanup job %s not found", id) + } + previous := job + update(&job) + if job == previous { + m.mu.Unlock() + return nil + } + m.jobs[id] = job + err := m.persistLocked() + if err != nil { + m.jobs[id] = previous + } + m.mu.Unlock() + if err == nil { + m.notify() + } + return err +} + +func (m *CleanupManager) RunDue(ctx context.Context, now time.Time) { + for { + m.mu.Lock() + var due CleanupJob + found := false + for id, job := range m.jobs { + if job.Phase != CleanupWaitingSeeding || job.DeleteAfterTime.IsZero() || job.DeleteAfterTime.After(now) { + continue + } + job.Phase = CleanupRunning + m.jobs[id] = job + due = job + found = true + break + } + if found { + if err := m.persistLocked(); err != nil { + log.Errorf("failed to persist cleanup job before execution: %v", err) + due.Phase = CleanupWaitingSeeding + m.jobs[due.ID] = due + m.mu.Unlock() + return + } + } + m.mu.Unlock() + if !found { + return + } + + err := m.execute(ctx, due) + m.mu.Lock() + if err != nil { + due.Phase = CleanupBlocked + due.LastError = err.Error() + m.jobs[due.ID] = due + } else { + delete(m.jobs, due.ID) + } + if persistErr := m.persistLocked(); persistErr != nil { + log.Errorf("failed to persist cleanup result: %v", persistErr) + } + m.mu.Unlock() + } +} + +func (m *CleanupManager) Protects(path string) bool { + path = filepath.Clean(path) + m.mu.Lock() + defer m.mu.Unlock() + for _, job := range m.jobs { + rel, err := filepath.Rel(path, filepath.Clean(job.TempDir)) + if err == nil && rel != ".." && !strings.HasPrefix(rel, ".."+string(filepath.Separator)) { + return true + } + } + return false +} + +func (m *CleanupManager) Get(id string) (CleanupJob, bool) { + m.mu.Lock() + defer m.mu.Unlock() + job, ok := m.jobs[id] + return job, ok +} + +func executeCleanupJob(ctx context.Context, job CleanupJob) error { + if job.GID == "" { + return fmt.Errorf("cleanup job %s has no download gid", job.ID) + } + tempDir := filepath.Clean(job.TempDir) + root := filepath.Clean(conf.Conf.TempDir) + rel, err := filepath.Rel(root, tempDir) + if err != nil || rel == "." || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { + return fmt.Errorf("refusing to remove temp dir outside configured root: %s", tempDir) + } + downloadTool, err := Tools.Get(job.Toolname) + if err != nil { + return err + } + t := &DownloadTask{ + TaskExtension: task.TaskExtension{}, + TempDir: tempDir, + Toolname: job.Toolname, + GID: job.GID, + } + t.SetCtx(ctx) + if err := downloadTool.Remove(t); err != nil { + return err + } + return os.RemoveAll(tempDir) +} diff --git a/internal/offline_download/tool/download.go b/internal/offline_download/tool/download.go index ca87768397..4318e22ceb 100644 --- a/internal/offline_download/tool/download.go +++ b/internal/offline_download/tool/download.go @@ -21,13 +21,17 @@ import ( type DownloadTask struct { task.TaskExtension Url string `json:"url"` + TorrentData []byte `json:"-"` DstDirPath string `json:"dst_dir_path"` TempDir string `json:"temp_dir"` DeletePolicy DeletePolicy `json:"delete_policy"` + DeleteAfterTime time.Time `json:"delete_after_time,omitempty"` Toolname string `json:"toolname"` + CleanupID string `json:"cleanup_id,omitempty"` + TransferStarted bool `json:"transfer_started,omitempty"` Status string `json:"-"` Signal chan int `json:"-"` - GID string `json:"-"` + GID string `json:"gid,omitempty"` tool Tool callStatusRetried int } @@ -45,26 +49,41 @@ func (t *DownloadTask) Run() error { } if err := t.tool.Run(t); !errs.IsNotSupportError(err) { if err == nil { - return t.Transfer() + return t.startTransfer() } return err } + if t.TransferStarted { + return nil + } t.Signal = make(chan int) defer func() { t.Signal = nil }() - gid, err := t.tool.AddURL(&AddUrlArgs{ - Ctx: t.Ctx(), - Url: t.Url, - UID: t.ID, - TempDir: t.TempDir, - Signal: t.Signal, - }) - if err != nil { - return err + if t.GID == "" { + gid, err := t.tool.AddURL(&AddUrlArgs{ + Ctx: t.Ctx(), + Url: t.Url, + TorrentData: t.TorrentData, + UID: t.ID, + TempDir: t.TempDir, + Signal: t.Signal, + }) + if err != nil { + return err + } + t.GID = gid + t.Persist() + if t.CleanupID != "" { + if err := CleanupTaskManager.SetGID(t.CleanupID, gid); err != nil { + return errors.WithMessage(err, "failed to persist cleanup gid") + } + } } - t.GID = gid - var ok bool + var ( + ok bool + err error + ) outer: for { select { @@ -86,6 +105,9 @@ outer: if err != nil { return err } + if err := t.startTransfer(); err != nil { + return err + } if t.tool.Name() == "Pikpak" { return nil } @@ -116,29 +138,35 @@ outer: if t.tool.Name() == "123 Open" { return nil } + if t.CleanupID != "" { + t.Status = "offline download completed, cleanup lifecycle is waiting for transfers and seeding" + return nil + } t.Status = "offline download completed, maybe transferring" // hack for qBittorrent if t.tool.Name() == "qBittorrent" { - seedTime := setting.GetInt(conf.QbittorrentSeedtime, 0) - if seedTime >= 0 { + if seedDuration, ok := t.seedingDuration(); ok { t.Status = "offline download completed, waiting for seeding" - <-time.After(time.Minute * time.Duration(seedTime)) - err := t.tool.Remove(t) - if err != nil { - log.Errorln(err.Error()) + <-time.After(seedDuration) + if t.shouldRemoveTaskAfterSeeding() { + err := t.tool.Remove(t) + if err != nil { + log.Errorln(err.Error()) + } } } } if t.tool.Name() == "Transmission" { // hack for transmission - seedTime := setting.GetInt(conf.TransmissionSeedtime, 0) - if seedTime >= 0 { + if seedDuration, ok := t.seedingDuration(); ok { t.Status = "offline download completed, waiting for seeding" - <-time.After(time.Minute * time.Duration(seedTime)) - err := t.tool.Remove(t) - if err != nil { - log.Errorln(err.Error()) + <-time.After(seedDuration) + if t.shouldRemoveTaskAfterSeeding() { + err := t.tool.Remove(t) + if err != nil { + log.Errorln(err.Error()) + } } } } @@ -167,8 +195,8 @@ func (t *DownloadTask) Update() (bool, error) { } // if download completed if info.Completed { - err := t.Transfer() - return true, errors.WithMessage(err, "failed to transfer file") + t.setDeleteAfterTime() + return true, nil } // if download failed if info.Err != nil { @@ -177,16 +205,105 @@ func (t *DownloadTask) Update() (bool, error) { return false, nil } +func (t *DownloadTask) setDeleteAfterTime() { + if t.DeletePolicy != DeleteAfterSeeding || !t.isSeedingTool() || !t.DeleteAfterTime.IsZero() { + return + } + seedDuration, ok := t.seedingDuration() + if !ok { + return + } + t.DeleteAfterTime = time.Now().Add(seedDuration) + t.Persist() +} + +func (t *DownloadTask) transferDeletePolicy() DeletePolicy { + if t.DeletePolicy == DeleteAfterSeeding && !t.isSeedingTool() { + return DeleteNever + } + return t.DeletePolicy +} + +func (t *DownloadTask) isSeedingTool() bool { + return isSeedingToolName(t.tool.Name()) +} + +func isSeedingToolName(toolName string) bool { + return toolName == "qBittorrent" || toolName == "Transmission" +} + +func (t *DownloadTask) shouldRemoveTaskAfterSeeding() bool { + return t.DeletePolicy != DeleteNever +} + +func (t *DownloadTask) seedingDuration() (time.Duration, bool) { + var seedTime int + switch t.tool.Name() { + case "qBittorrent": + seedTime = setting.GetInt(conf.QbittorrentSeedtime, 0) + case "Transmission": + seedTime = setting.GetInt(conf.TransmissionSeedtime, 0) + default: + return 0, false + } + if seedTime < 0 { + return 0, false + } + return time.Minute * time.Duration(seedTime), true +} + +func (t *DownloadTask) startTransfer() error { + if t.TransferStarted { + return nil + } + if t.CleanupID != "" { + if err := CleanupTaskManager.BeginTransfer(t.CleanupID, t.GID, t.DeleteAfterTime); err != nil { + return errors.WithMessage(err, "failed to begin cleanup transfer lifecycle") + } + } + if err := t.Transfer(); err != nil { + if t.CleanupID != "" { + _ = CleanupTaskManager.DownloadFailed(t.CleanupID, err) + } + return errors.WithMessage(err, "failed to transfer file") + } + t.TransferStarted = true + t.Persist() + if t.CleanupID != "" { + if err := CleanupTaskManager.FinishTransferSetup(t.CleanupID); err != nil { + return errors.WithMessage(err, "failed to finish cleanup transfer setup") + } + } + return nil +} + +func (t *DownloadTask) OnFailed() { + if t.CleanupID != "" { + if err := CleanupTaskManager.DownloadFailed(t.CleanupID, t.GetErr()); err != nil { + log.Errorf("failed to block cleanup lifecycle: %v", err) + } + } +} + +func (t *DownloadTask) OnBeforeRetry() { + if t.CleanupID != "" { + if err := CleanupTaskManager.RetryDownload(t.CleanupID); err != nil { + log.Errorf("failed to resume cleanup lifecycle: %v", err) + } + } +} + func (t *DownloadTask) Transfer() error { toolName := t.tool.Name() + deletePolicy := t.transferDeletePolicy() if toolName == "115 Cloud" || toolName == "115 Open" || toolName == "123 Open" || toolName == "123Pan" || toolName == "PikPak" || toolName == "Thunder" || toolName == "ThunderX" || toolName == "ThunderBrowser" || toolName == "GuangYaPan" { // 如果不是直接下载到目标路径,则进行转存 if t.TempDir != t.DstDirPath { - return transferObj(t.Ctx(), t.TempDir, t.DstDirPath, t.DeletePolicy) + return transferObj(t.Ctx(), t.TempDir, t.DstDirPath, deletePolicy, t.DeleteAfterTime, t.CleanupID) } return nil } - if t.DeletePolicy == UploadDownloadStream { + if deletePolicy == UploadDownloadStream { dstStorage, dstDirActualPath, err := op.GetStorageAndActualPath(t.DstDirPath) if err != nil { return errors.WithMessage(err, "failed get dst storage") @@ -203,16 +320,17 @@ func (t *DownloadTask) Transfer() error { DstStorage: dstStorage, DstStorageMp: dstStorage.GetStorage().MountPath, }, - DeletePolicy: t.DeletePolicy, - Url: t.Url, + DeletePolicy: deletePolicy, + DeleteAfterTime: t.DeleteAfterTime, + Url: t.Url, + CleanupID: t.CleanupID, } tsk.SetTotalBytes(t.GetTotalBytes()) tsk.groupID = path.Join(tsk.DstStorageMp, tsk.DstActualPath) task_group.TransferCoordinator.AddTask(tsk.groupID, nil) - TransferTaskManager.Add(tsk) - return nil + return addTransferTask(tsk) } - return transferStd(t.Ctx(), t.TempDir, t.DstDirPath, t.DeletePolicy) + return transferStd(t.Ctx(), t.TempDir, t.DstDirPath, deletePolicy, t.DeleteAfterTime, t.CleanupID) } func (t *DownloadTask) GetName() string { diff --git a/internal/offline_download/tool/transfer.go b/internal/offline_download/tool/transfer.go index 7109669ee7..0562910170 100644 --- a/internal/offline_download/tool/transfer.go +++ b/internal/offline_download/tool/transfer.go @@ -28,9 +28,12 @@ import ( type TransferTask struct { fs.TaskData - DeletePolicy DeletePolicy `json:"delete_policy"` - Url string `json:"url"` - groupID string `json:"-"` + DeletePolicy DeletePolicy `json:"delete_policy"` + DeleteAfterTime time.Time `json:"delete_after_time,omitempty"` + Url string `json:"url"` + CleanupID string `json:"cleanup_id,omitempty"` + CleanupFailed bool `json:"cleanup_failed,omitempty"` + groupID string `json:"-"` } func (t *TransferTask) Run() error { @@ -90,12 +93,24 @@ func (t *TransferTask) GetName() string { } func (t *TransferTask) OnSucceeded() { - if t.DeletePolicy == DeleteOnUploadSucceed || t.DeletePolicy == DeleteAlways { + switch t.DeletePolicy { + case DeleteOnUploadSucceed, DeleteAlways: if t.SrcStorage == nil { removeStdTemp(t) } else { removeObjTemp(t) } + case DeleteAfterSeeding: + if t.CleanupID == "" { + t.removeTempAfterSeeding() + } + } + if t.CleanupID != "" { + if err := CleanupTaskManager.TransferSucceeded(t.CleanupID); err != nil { + log.Errorf("failed to update cleanup lifecycle after transfer succeeded: %v", err) + } + t.CleanupFailed = false + t.Persist() } task_group.TransferCoordinator.Done(context.WithoutCancel(t.Ctx()), t.groupID, true) } @@ -108,9 +123,28 @@ func (t *TransferTask) OnFailed() { removeObjTemp(t) } } + if t.CleanupID != "" { + if err := CleanupTaskManager.TransferFailed(t.CleanupID, t.GetErr()); err != nil { + log.Errorf("failed to block cleanup lifecycle after transfer failed: %v", err) + } + t.CleanupFailed = true + t.Persist() + } task_group.TransferCoordinator.Done(context.WithoutCancel(t.Ctx()), t.groupID, false) } +func (t *TransferTask) OnBeforeRetry() { + if t.CleanupID == "" || !t.CleanupFailed { + return + } + if err := CleanupTaskManager.RetryTransfer(t.CleanupID); err != nil { + log.Errorf("failed to resume cleanup lifecycle for transfer retry: %v", err) + return + } + t.CleanupFailed = false + t.Persist() +} + func (t *TransferTask) SetRetry(retry int, maxRetry int) { if retry == 0 && (len(t.groupID) == 0 || // 重启恢复 @@ -121,11 +155,43 @@ func (t *TransferTask) SetRetry(retry int, maxRetry int) { t.TaskData.SetRetry(retry, maxRetry) } +func (t *TransferTask) removeTempAfterSeeding() { + if t.DeleteAfterTime.IsZero() { + return + } + remove := func() { + if t.SrcStorage == nil { + removeStdTemp(t) + } else { + removeObjTemp(t) + } + } + delay := time.Until(t.DeleteAfterTime) + if delay <= 0 { + remove() + return + } + time.AfterFunc(delay, remove) +} + var ( TransferTaskManager *tache.Manager[*TransferTask] ) -func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy DeletePolicy) error { +func addTransferTask(t *TransferTask) error { + if t.CleanupID != "" { + if CleanupTaskManager == nil { + return fmt.Errorf("offline cleanup manager is not initialized") + } + if err := CleanupTaskManager.AddTransfer(t.CleanupID); err != nil { + return err + } + } + TransferTaskManager.Add(t) + return nil +} + +func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy DeletePolicy, deleteAfterTime time.Time, cleanupID string) error { dstStorage, dstDirActualPath, err := op.GetStorageAndActualPath(dstDirPath) if err != nil { return errors.WithMessage(err, "failed get dst storage") @@ -147,11 +213,15 @@ func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy D DstStorage: dstStorage, DstStorageMp: dstStorage.GetStorage().MountPath, }, - DeletePolicy: deletePolicy, + DeletePolicy: deletePolicy, + DeleteAfterTime: deleteAfterTime, + CleanupID: cleanupID, } t.groupID = path.Join(t.DstStorageMp, t.DstActualPath) task_group.TransferCoordinator.AddTask(t.groupID, nil) - TransferTaskManager.Add(t) + if err := addTransferTask(t); err != nil { + return err + } } return nil } @@ -169,6 +239,9 @@ func transferStdPath(t *TransferTask) error { return err } dstDirActualPath := stdpath.Join(t.DstActualPath, info.Name()) + if err := op.MakeDir(t.Ctx(), t.DstStorage, dstDirActualPath); err != nil { + return errors.WithMessagef(err, "failed to make dst dir [%s]", dstDirActualPath) + } task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.DstPathToHook(dstDirActualPath)) for _, entry := range entries { srcRawPath := stdpath.Join(t.SrcActualPath, entry.Name()) @@ -184,11 +257,15 @@ func transferStdPath(t *TransferTask) error { SrcStorageMp: t.SrcStorageMp, DstStorageMp: t.DstStorageMp, }, - groupID: t.groupID, - DeletePolicy: t.DeletePolicy, + groupID: t.groupID, + DeletePolicy: t.DeletePolicy, + DeleteAfterTime: t.DeleteAfterTime, + CleanupID: t.CleanupID, } task_group.TransferCoordinator.AddTask(t.groupID, nil) - TransferTaskManager.Add(task) + if err := addTransferTask(task); err != nil { + return err + } } t.Status = "src object is dir, added all transfer tasks of files" return nil @@ -257,7 +334,7 @@ func removeStdTemp(t *TransferTask) { } } -func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy DeletePolicy) error { +func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy DeletePolicy, deleteAfterTime time.Time, cleanupID string) error { srcStorage, srcObjActualPath, err := op.GetStorageAndActualPath(tempDir) if err != nil { return errors.WithMessage(err, "failed get src storage") @@ -285,11 +362,15 @@ func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy D SrcStorageMp: srcStorage.GetStorage().MountPath, DstStorageMp: dstStorage.GetStorage().MountPath, }, - DeletePolicy: deletePolicy, + DeletePolicy: deletePolicy, + DeleteAfterTime: deleteAfterTime, + CleanupID: cleanupID, } t.groupID = path.Join(t.DstStorageMp, t.DstActualPath) task_group.TransferCoordinator.AddTask(t.groupID, nil) - TransferTaskManager.Add(t) + if err := addTransferTask(t); err != nil { + return err + } } return nil } @@ -307,6 +388,9 @@ func transferObjPath(t *TransferTask) error { return errors.WithMessagef(err, "failed list src [%s] objs", t.SrcActualPath) } dstDirActualPath := stdpath.Join(t.DstActualPath, srcObj.GetName()) + if err := op.MakeDir(t.Ctx(), t.DstStorage, dstDirActualPath); err != nil { + return errors.WithMessagef(err, "failed to make dst dir [%s]", dstDirActualPath) + } task_group.TransferCoordinator.AppendPayload(t.groupID, task_group.DstPathToHook(dstDirActualPath)) for _, obj := range objs { if utils.IsCanceled(t.Ctx()) { @@ -314,7 +398,7 @@ func transferObjPath(t *TransferTask) error { } srcObjPath := stdpath.Join(t.SrcActualPath, obj.GetName()) task_group.TransferCoordinator.AddTask(t.groupID, nil) - TransferTaskManager.Add(&TransferTask{ + childTask := &TransferTask{ TaskData: fs.TaskData{ TaskExtension: task.TaskExtension{ Creator: t.Creator, @@ -327,9 +411,14 @@ func transferObjPath(t *TransferTask) error { SrcStorageMp: t.SrcStorageMp, DstStorageMp: t.DstStorageMp, }, - groupID: t.groupID, - DeletePolicy: t.DeletePolicy, - }) + groupID: t.groupID, + DeletePolicy: t.DeletePolicy, + DeleteAfterTime: t.DeleteAfterTime, + CleanupID: t.CleanupID, + } + if err := addTransferTask(childTask); err != nil { + return err + } } t.Status = "src object is dir, added all transfer tasks of objs" return nil diff --git a/pkg/qbittorrent/client.go b/pkg/qbittorrent/client.go index e4c12db6fc..740d1ea864 100644 --- a/pkg/qbittorrent/client.go +++ b/pkg/qbittorrent/client.go @@ -15,6 +15,7 @@ import ( type Client interface { AddFromLink(link string, savePath string, id string) error + AddFromTorrent(torrentData []byte, savePath string, id string) error GetInfo(id string) (TorrentInfo, error) GetFiles(id string) ([]FileInfo, error) Delete(id string, deleteFiles bool) error @@ -142,6 +143,26 @@ func (c *client) post(path string, data url.Values) (*http.Response, error) { } func (c *client) AddFromLink(link string, savePath string, id string) error { + return c.addTorrent(func(writer *multipart.Writer) error { + return writer.WriteField("urls", link) + }, savePath, id, link) +} + +func (c *client) AddFromTorrent(torrentData []byte, savePath string, id string) error { + if len(torrentData) == 0 { + return errors.New("empty torrent data") + } + return c.addTorrent(func(writer *multipart.Writer) error { + part, err := writer.CreateFormFile("torrents", id+".torrent") + if err != nil { + return err + } + _, err = part.Write(torrentData) + return err + }, savePath, id, "torrent file") +} + +func (c *client) addTorrent(writeContent func(*multipart.Writer) error, savePath string, id string, description string) error { err := c.checkAuthorization() if err != nil { return err @@ -150,22 +171,19 @@ func (c *client) AddFromLink(link string, savePath string, id string) error { buf := new(bytes.Buffer) writer := multipart.NewWriter(buf) - addField := func(name string, value string) { - if err != nil { - return - } - err = writer.WriteField(name, value) + if err := writeContent(writer); err != nil { + return err } - addField("urls", link) - addField("savepath", savePath) - addField("tags", "openlist-"+id) - addField("autoTMM", "false") - if err != nil { + if err := writer.WriteField("savepath", savePath); err != nil { return err } - - err = writer.Close() - if err != nil { + if err := writer.WriteField("tags", "openlist-"+id); err != nil { + return err + } + if err := writer.WriteField("autoTMM", "false"); err != nil { + return err + } + if err := writer.Close(); err != nil { return err } @@ -182,21 +200,17 @@ func (c *client) AddFromLink(link string, savePath string, id string) error { return err } defer resp.Body.Close() - // qBittorrent 5.2.0 returns 204 on success. - if resp.StatusCode != http.StatusNoContent { + + if resp.StatusCode >= 200 && resp.StatusCode < 300 { return nil } - // check result - body := make([]byte, 2) - _, err = resp.Body.Read(body) - if err != nil { - return err - } - if resp.StatusCode != 200 || string(body) != "Ok" { - return errors.New("failed to add qBittorrent task: " + link) + body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024)) + msg := "failed to add qBittorrent task: " + description + ", status: " + resp.Status + if body = bytes.TrimSpace(body); len(body) > 0 { + msg += ", response: " + string(body) } - return nil + return errors.New(msg) } type TorrentStatus string diff --git a/server/handles/offline_download.go b/server/handles/offline_download.go index 4d4167ba8a..2f033a38cc 100644 --- a/server/handles/offline_download.go +++ b/server/handles/offline_download.go @@ -1,6 +1,9 @@ package handles import ( + "encoding/base64" + "encoding/json" + "fmt" "strings" _115 "github.com/OpenListTeam/OpenList/v4/drivers/115" @@ -18,6 +21,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/task" + "github.com/OpenListTeam/OpenList/v4/pkg/torrent" "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/gin-gonic/gin" "github.com/pkg/errors" @@ -523,10 +527,74 @@ func OfflineDownloadTools(c *gin.Context) { } type AddOfflineDownloadReq struct { - Urls []string `json:"urls"` - Path string `json:"path"` - Tool string `json:"tool"` - DeletePolicy string `json:"delete_policy"` + Urls []string `json:"urls"` + TorrentData torrentDataList `json:"torrent_data"` + Path string `json:"path"` + Tool string `json:"tool"` + DeletePolicy string `json:"delete_policy"` +} + +type torrentDataList []string + +func (l *torrentDataList) UnmarshalJSON(data []byte) error { + var list []string + if err := json.Unmarshal(data, &list); err == nil { + *l = list + return nil + } + + var single string + if err := json.Unmarshal(data, &single); err == nil { + if strings.TrimSpace(single) == "" { + *l = nil + } else { + *l = []string{single} + } + return nil + } + + return fmt.Errorf("torrent_data must be a base64 string or string array") +} + +func decodeOfflineDownloadTorrentData(encoded string) ([]byte, string, error) { + encoded = strings.TrimSpace(encoded) + if encoded == "" { + return nil, "", nil + } + if strings.HasPrefix(strings.ToLower(encoded), "data:") { + comma := strings.Index(encoded, ",") + if comma < 0 { + return nil, "", fmt.Errorf("invalid torrent data URL") + } + meta := strings.ToLower(encoded[:comma]) + if !strings.Contains(meta, ";base64") { + return nil, "", fmt.Errorf("torrent data URL must use base64 encoding") + } + encoded = encoded[comma+1:] + } + + if len(encoded) > maxTorrentBase64Len { + return nil, "", fmt.Errorf("torrent data is too large, maximum size is 10MB") + } + + torrentData, err := base64.StdEncoding.DecodeString(encoded) + if err != nil { + return nil, "", fmt.Errorf("invalid torrent base64 encoding: %w", err) + } + + t, err := torrent.Decode(torrentData) + if err != nil { + return nil, "", fmt.Errorf("failed to parse torrent: %w", err) + } + if err := validateParsedTorrent(t); err != nil { + return nil, "", err + } + + name := t.Info.Name + if name == "" { + name = t.GetInfoHashHex() + } + return torrentData, name, nil } func AddOfflineDownload(c *gin.Context) { @@ -577,6 +645,31 @@ func AddOfflineDownload(c *gin.Context) { tasks = append(tasks, t) } } + for _, encodedTorrent := range req.TorrentData { + torrentData, name, err := decodeOfflineDownloadTorrentData(encodedTorrent) + if err != nil { + common.ErrorResp(c, err, 400) + return + } + if len(torrentData) == 0 { + continue + } + + t, err := tool.AddURL(c, &tool.AddURLArgs{ + URL: name, + TorrentData: torrentData, + DstDirPath: reqPath, + Tool: req.Tool, + DeletePolicy: tool.DeletePolicy(req.DeletePolicy), + }) + if err != nil { + common.ErrorResp(c, err, 500) + return + } + if t != nil { + tasks = append(tasks, t) + } + } common.SuccessResp(c, gin.H{ "tasks": getTaskInfos(tasks), }) diff --git a/tests/offline_cleanup_test.go b/tests/offline_cleanup_test.go new file mode 100644 index 0000000000..1cbcfb44f9 --- /dev/null +++ b/tests/offline_cleanup_test.go @@ -0,0 +1,180 @@ +package tests + +import ( + "context" + "encoding/json" + "errors" + "os" + "path/filepath" + "testing" + "time" + + "github.com/OpenListTeam/OpenList/v4/internal/bootstrap" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" +) + +type cleanupStore struct { + data []byte +} + +func (s *cleanupStore) read() ([]byte, error) { + if len(s.data) == 0 { + return []byte("[]"), nil + } + return append([]byte(nil), s.data...), nil +} + +func (s *cleanupStore) write(data []byte) error { + s.data = append(s.data[:0], data...) + return nil +} + +func TestOfflineCleanupBlocksUntilFailedTransferSucceeds(t *testing.T) { + store := &cleanupStore{} + executed := 0 + manager, err := tool.NewCleanupManager(store.read, store.write, func(context.Context, tool.CleanupJob) error { + executed++ + return nil + }) + if err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(-time.Minute) + job := tool.CleanupJob{ + ID: "cleanup-1", + TempDir: "/tmp/openlist/qBittorrent/task-1", + Toolname: "qBittorrent", + } + if err := manager.Register(job); err != nil { + t.Fatal(err) + } + if err := manager.BeginTransfer(job.ID, "gid-1", deadline); err != nil { + t.Fatal(err) + } + if err := manager.AddTransfer(job.ID); err != nil { + t.Fatal(err) + } + if err := manager.FinishTransferSetup(job.ID); err != nil { + t.Fatal(err) + } + if err := manager.TransferFailed(job.ID, errors.New("upload failed")); err != nil { + t.Fatal(err) + } + + manager.RunDue(context.Background(), time.Now()) + if executed != 0 { + t.Fatalf("cleanup executed %d times while transfer was failed", executed) + } + blocked, ok := manager.Get(job.ID) + if !ok || blocked.Phase != tool.CleanupBlocked || blocked.FailedTransfers != 1 { + t.Fatalf("unexpected blocked cleanup job: %+v, exists=%v", blocked, ok) + } + + if err := manager.RetryTransfer(job.ID); err != nil { + t.Fatal(err) + } + if err := manager.TransferSucceeded(job.ID); err != nil { + t.Fatal(err) + } + manager.RunDue(context.Background(), time.Now()) + if executed != 1 { + t.Fatalf("cleanup executed %d times after retry succeeded, want 1", executed) + } + if _, ok := manager.Get(job.ID); ok { + t.Fatal("completed cleanup job was not removed") + } + + var persisted []tool.CleanupJob + if err := json.Unmarshal(store.data, &persisted); err != nil { + t.Fatal(err) + } + if len(persisted) != 0 { + t.Fatalf("persisted cleanup jobs = %+v, want none", persisted) + } +} + +func TestOfflineCleanupRestoresSeedingDeadlineAfterRestart(t *testing.T) { + store := &cleanupStore{} + deadline := time.Now().Add(time.Hour) + first, err := tool.NewCleanupManager(store.read, store.write, func(context.Context, tool.CleanupJob) error { return nil }) + if err != nil { + t.Fatal(err) + } + job := tool.CleanupJob{ID: "cleanup-restore", TempDir: "/tmp/openlist/qBittorrent/task-restore", Toolname: "qBittorrent"} + if err := first.Register(job); err != nil { + t.Fatal(err) + } + if err := first.BeginTransfer(job.ID, "gid-restore", deadline); err != nil { + t.Fatal(err) + } + if err := first.AddTransfer(job.ID); err != nil { + t.Fatal(err) + } + if err := first.FinishTransferSetup(job.ID); err != nil { + t.Fatal(err) + } + if err := first.TransferSucceeded(job.ID); err != nil { + t.Fatal(err) + } + + executed := 0 + restored, err := tool.NewCleanupManager(store.read, store.write, func(context.Context, tool.CleanupJob) error { + executed++ + return nil + }) + if err != nil { + t.Fatal(err) + } + restoredJob, ok := restored.Get(job.ID) + if !ok || restoredJob.Phase != tool.CleanupWaitingSeeding || !restoredJob.DeleteAfterTime.Equal(deadline) { + t.Fatalf("unexpected restored cleanup job: %+v, exists=%v", restoredJob, ok) + } + restored.RunDue(context.Background(), deadline.Add(time.Second)) + if executed != 1 { + t.Fatalf("restored cleanup executed %d times, want 1", executed) + } +} + +func TestGlobalTempCleanupPreservesReferencedPaths(t *testing.T) { + root := t.TempDir() + protected := filepath.Join(root, "qBittorrent", "task-1") + orphan := filepath.Join(root, "aria2", "orphan") + if err := os.MkdirAll(protected, 0o700); err != nil { + t.Fatal(err) + } + if err := os.MkdirAll(orphan, 0o700); err != nil { + t.Fatal(err) + } + + store := &cleanupStore{} + manager, err := tool.NewCleanupManager(store.read, store.write, func(context.Context, tool.CleanupJob) error { return nil }) + if err != nil { + t.Fatal(err) + } + if err := manager.Register(tool.CleanupJob{ + ID: "cleanup-1", + TempDir: protected, + Toolname: "qBittorrent", + }); err != nil { + t.Fatal(err) + } + + oldConf := conf.Conf + oldManager := tool.CleanupTaskManager + conf.Conf = conf.DefaultConfig(root) + conf.Conf.TempDir = root + tool.CleanupTaskManager = manager + t.Cleanup(func() { + conf.Conf = oldConf + tool.CleanupTaskManager = oldManager + }) + + bootstrap.CleanTempDir() + if _, err := os.Stat(protected); err != nil { + t.Fatalf("protected temp path was removed: %v", err) + } + if _, err := os.Stat(filepath.Join(root, "aria2")); !os.IsNotExist(err) { + t.Fatalf("orphan temp path still exists or stat failed: %v", err) + } +} diff --git a/tests/qbit_test.go b/tests/qbit_test.go new file mode 100644 index 0000000000..1fda999016 --- /dev/null +++ b/tests/qbit_test.go @@ -0,0 +1,40 @@ +package tests + +import ( + "testing" + + qbit "github.com/OpenListTeam/OpenList/v4/internal/offline_download/qbit" + "github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool" + "github.com/OpenListTeam/OpenList/v4/pkg/qbittorrent" +) + +type statusClient struct { + qbittorrent.Client + info qbittorrent.TorrentInfo +} + +func (c *statusClient) GetInfo(string) (qbittorrent.TorrentInfo, error) { + return c.info, nil +} + +func TestQbittorrentStatusUsesSavePath(t *testing.T) { + client := &statusClient{info: qbittorrent.TorrentInfo{ + SavePath: "/downloads/existing-task", + Size: 42, + Completed: 42, + State: qbittorrent.UPLOADING, + }} + qbitTool := qbit.New(client) + task := &tool.DownloadTask{TempDir: "/downloads/new-task"} + + status, err := qbitTool.Status(task) + if err != nil { + t.Fatal(err) + } + if task.TempDir != client.info.SavePath { + t.Fatalf("TempDir = %q, want %q", task.TempDir, client.info.SavePath) + } + if !status.Completed { + t.Fatal("completed torrent was not reported as completed") + } +} diff --git a/tests/qbittorrent_client_test.go b/tests/qbittorrent_client_test.go new file mode 100644 index 0000000000..4bf3337a4b --- /dev/null +++ b/tests/qbittorrent_client_test.go @@ -0,0 +1,141 @@ +package tests + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/OpenListTeam/OpenList/v4/pkg/qbittorrent" +) + +func TestQbittorrentAddFromTorrentUploadsTorrentFile(t *testing.T) { + torrentData := []byte("torrent-content") + addCalled := false + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v2/app/version": + w.WriteHeader(http.StatusOK) + case "/api/v2/torrents/add": + addCalled = true + if err := r.ParseMultipartForm(1024 * 1024); err != nil { + t.Fatalf("ParseMultipartForm() error = %v", err) + } + if got := r.MultipartForm.Value["savepath"]; len(got) != 1 || got[0] != "/downloads" { + t.Fatalf("savepath = %v", got) + } + if got := r.MultipartForm.Value["tags"]; len(got) != 1 || got[0] != "openlist-task-id" { + t.Fatalf("tags = %v", got) + } + if got := r.MultipartForm.Value["autoTMM"]; len(got) != 1 || got[0] != "false" { + t.Fatalf("autoTMM = %v", got) + } + if got := r.MultipartForm.Value["urls"]; len(got) != 0 { + t.Fatalf("urls field should be absent, got %v", got) + } + + files := r.MultipartForm.File["torrents"] + if len(files) != 1 { + t.Fatalf("torrents files = %d", len(files)) + } + if files[0].Filename != "task-id.torrent" { + t.Fatalf("torrent filename = %q", files[0].Filename) + } + file, err := files[0].Open() + if err != nil { + t.Fatalf("Open() error = %v", err) + } + defer file.Close() + got, err := io.ReadAll(file) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if string(got) != string(torrentData) { + t.Fatalf("torrent data = %q", got) + } + + w.WriteHeader(http.StatusNoContent) + default: + t.Fatalf("unexpected path %s", r.URL.Path) + } + })) + defer server.Close() + + client, err := qbittorrent.New(server.URL) + if err != nil { + t.Fatalf("New() error = %v", err) + } + if err := client.AddFromTorrent(torrentData, "/downloads", "task-id"); err != nil { + t.Fatalf("AddFromTorrent() error = %v", err) + } + if !addCalled { + t.Fatal("qBittorrent add endpoint was not called") + } +} + +func TestQbittorrentAddFromLinkUsesUrlsFieldAndAcceptsHTTP200(t *testing.T) { + addCalled := false + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v2/app/version": + w.WriteHeader(http.StatusOK) + case "/api/v2/torrents/add": + addCalled = true + if err := r.ParseMultipartForm(1024 * 1024); err != nil { + t.Fatalf("ParseMultipartForm() error = %v", err) + } + if got := r.MultipartForm.Value["urls"]; len(got) != 1 || got[0] != "magnet:?xt=urn:btih:test" { + t.Fatalf("urls = %v", got) + } + if got := r.MultipartForm.File["torrents"]; len(got) != 0 { + t.Fatalf("torrents field should be absent, got %v", got) + } + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("Ok.")) + default: + t.Fatalf("unexpected path %s", r.URL.Path) + } + })) + defer server.Close() + + client, err := qbittorrent.New(server.URL) + if err != nil { + t.Fatalf("New() error = %v", err) + } + if err := client.AddFromLink("magnet:?xt=urn:btih:test", "/downloads", "task-id"); err != nil { + t.Fatalf("AddFromLink() error = %v", err) + } + if !addCalled { + t.Fatal("qBittorrent add endpoint was not called") + } +} + +func TestQbittorrentAddFromLinkReportsNon2xxResponse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v2/app/version": + w.WriteHeader(http.StatusOK) + case "/api/v2/torrents/add": + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte("bad torrent")) + default: + t.Fatalf("unexpected path %s", r.URL.Path) + } + })) + defer server.Close() + + client, err := qbittorrent.New(server.URL) + if err != nil { + t.Fatalf("New() error = %v", err) + } + err = client.AddFromLink("magnet:?xt=urn:btih:test", "/downloads", "task-id") + if err == nil { + t.Fatal("AddFromLink() error = nil") + } + if !strings.Contains(err.Error(), "400 Bad Request") || !strings.Contains(err.Error(), "bad torrent") { + t.Fatalf("AddFromLink() error = %v", err) + } +}