diff --git a/internal/handler/download_client_handler.go b/internal/handler/download_client_handler.go index a2d7c7b..7c49948 100644 --- a/internal/handler/download_client_handler.go +++ b/internal/handler/download_client_handler.go @@ -100,13 +100,11 @@ func (h *DownloadClientHandler) Create(c *gin.Context) { return } - // 热插拔:加载新客户端 - go func() { - if initErr := h.svc.DownloadMgr.AddClient(ctx, client); initErr != nil { - h.log.Warn("failed to hot-add download client", zap.Error(initErr)) - } - }() - _ = h.svc.Downloads.ReloadConfig(ctx) + if err := h.svc.Downloads.ReloadConfig(ctx); err != nil { + h.log.Warn("failed to reload download clients after create", zap.Error(err)) + Error(c, http.StatusInternalServerError, ErrInternal, "客户端已保存,但运行时重载失败: "+err.Error()) + return + } Success(c, client) } @@ -207,13 +205,11 @@ func (h *DownloadClientHandler) Update(c *gin.Context) { } clearLegacyQBitSettingsIfNoDefault(c.Request.Context(), h.svc) - // 热更新适配器 - go func() { - if updateErr := h.svc.DownloadMgr.UpdateClient(ctx, client); updateErr != nil { - h.log.Warn("failed to hot-update download client", zap.Error(updateErr)) - } - }() - _ = h.svc.Downloads.ReloadConfig(ctx) + if err := h.svc.Downloads.ReloadConfig(ctx); err != nil { + h.log.Warn("failed to reload download clients after update", zap.Error(err)) + Error(c, http.StatusInternalServerError, ErrInternal, "客户端已更新,但运行时重载失败: "+err.Error()) + return + } Success(c, client) } @@ -235,10 +231,12 @@ func (h *DownloadClientHandler) Delete(c *gin.Context) { return } - // 热移除 - h.svc.DownloadMgr.RemoveClient(id) clearLegacyQBitSettingsIfNoDefault(c.Request.Context(), h.svc) - _ = h.svc.Downloads.ReloadConfig(ctx) + if err := h.svc.Downloads.ReloadConfig(ctx); err != nil { + h.log.Warn("failed to reload download clients after delete", zap.Error(err)) + Error(c, http.StatusInternalServerError, ErrInternal, "客户端已删除,但运行时重载失败: "+err.Error()) + return + } SuccessWithMessage(c, "已删除", nil) } diff --git a/internal/handler/download_clients.go b/internal/handler/download_clients.go index 0680c4e..308eeac 100644 --- a/internal/handler/download_clients.go +++ b/internal/handler/download_clients.go @@ -39,7 +39,10 @@ func createDownloadClientHandler(svc *service.Container) gin.HandlerFunc { } // 让真正发起下载的 DownloadService 立刻读到新的 qb 配置, // 避免保存后还要重启进程才能生效。 - _ = svc.Downloads.ReloadConfig(c.Request.Context()) + if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "client saved but runtime reload failed: " + err.Error()}) + return + } c.JSON(http.StatusCreated, row) } } @@ -56,7 +59,10 @@ func updateDownloadClientHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } - _ = svc.Downloads.ReloadConfig(c.Request.Context()) + if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "client updated but runtime reload failed: " + err.Error()}) + return + } c.JSON(http.StatusOK, row) } } @@ -67,7 +73,10 @@ func deleteDownloadClientHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } - _ = svc.Downloads.ReloadConfig(c.Request.Context()) + if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "client deleted but runtime reload failed: " + err.Error()}) + return + } c.Status(http.StatusNoContent) } } diff --git a/internal/handler/downloads.go b/internal/handler/downloads.go index b5e9a81..25e3917 100644 --- a/internal/handler/downloads.go +++ b/internal/handler/downloads.go @@ -138,11 +138,26 @@ func visibleLiveTorrents(rows []model.DownloadTask, live []service.QBitTorrent) } filtered := make([]service.QBitTorrent, 0, len(live)) for _, torrent := range live { + matchedByID := false + for _, row := range rows { + if strings.TrimSpace(row.ExternalID) != "" && strings.EqualFold(row.ExternalID, torrent.Hash) && + (strings.TrimSpace(row.DownloadClientID) == "" || row.DownloadClientID == torrent.ClientID) { + filtered = append(filtered, torrent) + matchedByID = true + break + } + } + if matchedByID { + continue + } torrentTitle := normalizeTitle(torrent.Name) if torrentTitle == "" { continue } for _, row := range rows { + if strings.TrimSpace(row.DownloadClientID) != "" && strings.TrimSpace(torrent.ClientID) != "" && row.DownloadClientID != torrent.ClientID { + continue + } rowTitle := normalizeTitle(row.Title) if rowTitle == "" { continue @@ -197,7 +212,7 @@ func deleteDownloadHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { hash := c.Param("hash") withFiles := c.Query("delete_files") == "true" - if err := svc.Downloads.Delete(c.Request.Context(), hash, withFiles); err != nil { + if err := svc.Downloads.Delete(c.Request.Context(), hash, withFiles, c.Query("client_id")); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } @@ -208,6 +223,7 @@ func deleteDownloadHandler(svc *service.Container) gin.HandlerFunc { type relocateDownloadReq struct { Hash string `json:"hash" binding:"required"` Location string `json:"location" binding:"required"` + ClientID string `json:"client_id"` } // relocateDownloadHandler moves a torrent's data to a new directory while @@ -219,8 +235,12 @@ func relocateDownloadHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - if err := svc.Downloads.RelocateTorrent(c.Request.Context(), req.Hash, req.Location); err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + if err := svc.Downloads.RelocateTorrent(c.Request.Context(), req.Hash, req.Location, req.ClientID); err != nil { + status := http.StatusInternalServerError + if errors.Is(err, service.ErrDownloadOperationUnsupported) { + status = http.StatusBadRequest + } + c.JSON(status, gin.H{"error": err.Error()}) return } c.JSON(http.StatusOK, gin.H{"hash": strings.TrimSpace(req.Hash), "location": strings.TrimSpace(req.Location)}) diff --git a/internal/handler/downloads_extra.go b/internal/handler/downloads_extra.go index 9946cda..c3c33c6 100644 --- a/internal/handler/downloads_extra.go +++ b/internal/handler/downloads_extra.go @@ -11,15 +11,9 @@ import ( "github.com/ShukeBta/MediaStationGo/internal/service" ) -// downloadPauseHandler is a thin alias — the underlying qBittorrent -// service exposes pause via the WebUI; we mark our local row too so -// the React UI shows the right state on next refresh. func downloadPauseHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { - if err := svc.Repo.DB.WithContext(c.Request.Context()). - Model(&model.DownloadTask{}). - Where("id = ?", c.Param("id")). - Update("status", "paused").Error; err != nil { + if err := svc.Downloads.PauseDownloadTask(c.Request.Context(), c.Param("id")); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } @@ -27,13 +21,9 @@ func downloadPauseHandler(svc *service.Container) gin.HandlerFunc { } } -// downloadResumeHandler marks the row as queued so the next poll picks it up. func downloadResumeHandler(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { - if err := svc.Repo.DB.WithContext(c.Request.Context()). - Model(&model.DownloadTask{}). - Where("id = ?", c.Param("id")). - Update("status", "queued").Error; err != nil { + if err := svc.Downloads.ResumeDownloadTask(c.Request.Context(), c.Param("id")); err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } diff --git a/internal/model/download_subscription.go b/internal/model/download_subscription.go index 3b230bd..50ad885 100644 --- a/internal/model/download_subscription.go +++ b/internal/model/download_subscription.go @@ -5,17 +5,19 @@ import "time" // DownloadTask 是待处理(或已完成)的 torrent / HTTP 下载。 type DownloadTask struct { Base - UserID string `gorm:"index;size:36" json:"user_id"` - SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"` - Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / http - URL string `gorm:"size:2048;not null" json:"-"` - Title string `gorm:"size:512" json:"title,omitempty"` - PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"` - BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"` - Overview string `gorm:"type:text" json:"overview,omitempty"` - SavePath string `gorm:"size:1024" json:"save_path"` - MediaType string `gorm:"size:16" json:"media_type,omitempty"` - MediaCategory string `gorm:"size:128" json:"media_category,omitempty"` + UserID string `gorm:"index;size:36" json:"user_id"` + SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"` + DownloadClientID string `gorm:"index;size:36" json:"download_client_id,omitempty"` + ExternalID string `gorm:"index;size:128" json:"external_id,omitempty"` + Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / aria2 / http + URL string `gorm:"size:2048;not null" json:"-"` + Title string `gorm:"size:512" json:"title,omitempty"` + PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"` + BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"` + Overview string `gorm:"type:text" json:"overview,omitempty"` + SavePath string `gorm:"size:1024" json:"save_path"` + MediaType string `gorm:"size:16" json:"media_type,omitempty"` + MediaCategory string `gorm:"size:128" json:"media_category,omitempty"` // 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。 OriginalName string `gorm:"size:512" json:"original_name,omitempty"` OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"` diff --git a/internal/repository/download_client_repo.go b/internal/repository/download_client_repo.go index b04bbd2..9883878 100644 --- a/internal/repository/download_client_repo.go +++ b/internal/repository/download_client_repo.go @@ -35,7 +35,10 @@ func (r *DownloadClientRepository) FindByID(ctx context.Context, id string) (*mo // FindDefault returns the default download client, or (nil, nil). func (r *DownloadClientRepository) FindDefault(ctx context.Context) (*model.DownloadClient, error) { var c model.DownloadClient - err := r.db.WithContext(ctx).Where("is_default = ? AND enabled = ?", true, true).First(&c).Error + err := r.db.WithContext(ctx). + Where("is_default = ? AND enabled = ?", true, true). + Order("created_at asc"). + First(&c).Error if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } diff --git a/internal/service/aria2_adp.go b/internal/service/aria2_adp.go index ade3ca8..01af2e1 100644 --- a/internal/service/aria2_adp.go +++ b/internal/service/aria2_adp.go @@ -6,7 +6,9 @@ package service import ( "context" + "encoding/base64" "encoding/json" + "errors" "fmt" "net/http" "strings" @@ -14,6 +16,8 @@ import ( "time" ) +var errAria2ListUnavailable = errors.New("aria2 task list unavailable") + // Aria2Adapter 是 Aria2 的 DownloadAdapter 实现。 type Aria2Adapter struct { mu sync.Mutex @@ -57,6 +61,30 @@ func (a *Aria2Adapter) AddMagnet(ctx context.Context, magnet, savePath string) ( return a.AddTorrent(ctx, magnet, savePath) } +// AddTorrentFile submits application-fetched .torrent bytes using +// aria2.addTorrent. The empty URI list matches aria2's RPC signature. +func (a *Aria2Adapter) AddTorrentFile(ctx context.Context, data []byte, _ string, savePath string) (string, error) { + a.mu.Lock() + defer a.mu.Unlock() + options := map[string]string{} + if savePath != "" { + options["dir"] = savePath + } + result, err := a.rpcLocked(ctx, "aria2.addTorrent", []interface{}{ + base64.StdEncoding.EncodeToString(data), + []string{}, + options, + }) + if err != nil { + return "", err + } + var gid string + if err := json.Unmarshal(result, &gid); err != nil { + return "", err + } + return gid, nil +} + // Pause 暂停下载任务(通过 GID)。 func (a *Aria2Adapter) Pause(ctx context.Context, hash string) error { a.mu.Lock() @@ -77,12 +105,13 @@ func (a *Aria2Adapter) Resume(ctx context.Context, hash string) error { func (a *Aria2Adapter) Remove(ctx context.Context, hash string, deleteFiles bool) error { a.mu.Lock() defer a.mu.Unlock() - if deleteFiles { - _, err := a.rpcLocked(ctx, "aria2.removeDownloadResult", []interface{}{hash}) - return err + _, removeErr := a.rpcLocked(ctx, "aria2.remove", []interface{}{hash}) + _, resultErr := a.rpcLocked(ctx, "aria2.removeDownloadResult", []interface{}{hash}) + if removeErr != nil && resultErr != nil { + return errors.Join(removeErr, resultErr) } - _, err := a.rpcLocked(ctx, "aria2.remove", []interface{}{hash}) - return err + _ = deleteFiles // aria2 RPC removes the task/result but has no delete-local-data flag. + return nil } // List 列出所有活动/等待/已停止的任务。 @@ -91,34 +120,48 @@ func (a *Aria2Adapter) List(ctx context.Context, filter string) ([]TorrentInfo, defer a.mu.Unlock() var allResults []TorrentInfo + var listErrs []error + var successfulCalls int // 获取活动任务 active, err := a.rpcLocked(ctx, "aria2.tellActive", []interface{}{ - []string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"}, + []string{"gid", "bittorrent", "files", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"}, }) if err == nil && active != nil { + successfulCalls++ items := a.parseAria2Items(active) allResults = append(allResults, items...) + } else if err != nil { + listErrs = append(listErrs, err) } // 获取等待中的任务 waiting, err := a.rpcLocked(ctx, "aria2.tellWaiting", []interface{}{ 0, 100, - []string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"}, + []string{"gid", "bittorrent", "files", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"}, }) if err == nil && waiting != nil { + successfulCalls++ items := a.parseAria2Items(waiting) allResults = append(allResults, items...) + } else if err != nil { + listErrs = append(listErrs, err) } // 获取已停止的任务 stopped, err := a.rpcLocked(ctx, "aria2.tellStopped", []interface{}{ 0, 100, - []string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"}, + []string{"gid", "bittorrent", "files", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"}, }) if err == nil && stopped != nil { + successfulCalls++ items := a.parseAria2Items(stopped) allResults = append(allResults, items...) + } else if err != nil { + listErrs = append(listErrs, err) + } + if successfulCalls == 0 { + return nil, fmt.Errorf("%w: %w", errAria2ListUnavailable, errors.Join(listErrs...)) } if filter != "" { @@ -128,10 +171,10 @@ func (a *Aria2Adapter) List(ctx context.Context, filter string) ([]TorrentInfo, filtered = append(filtered, item) } } - return filtered, nil + return filtered, errors.Join(listErrs...) } - return allResults, nil + return allResults, errors.Join(listErrs...) } // GetInfo 获取单个任务信息。 @@ -141,7 +184,7 @@ func (a *Aria2Adapter) GetInfo(ctx context.Context, hash string) (*TorrentInfo, result, err := a.rpcLocked(ctx, "aria2.tellStatus", []interface{}{ hash, - []string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections"}, + []string{"gid", "bittorrent", "files", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections"}, }) if err != nil { return nil, err diff --git a/internal/service/aria2_parse.go b/internal/service/aria2_parse.go index b2ba067..8976851 100644 --- a/internal/service/aria2_parse.go +++ b/internal/service/aria2_parse.go @@ -35,27 +35,24 @@ func (a *Aria2Adapter) parseSingleItem(item map[string]interface{}) *TorrentInfo connections := int(toInt64(item["connections"])) var name string - var hash string + var contentPath string // 尝试从 bittorrent info 获取名称和 hash if bt, ok := item["bittorrent"].(map[string]interface{}); ok { if info, ok := bt["info"].(map[string]interface{}); ok { name = strVal(info["name"]) } - hash = strVal(bt["infoHash"]) + contentPath = downloaderPayloadPath(dir, name) } - // 如果没有 bittorrent 信息,使用 GID 作为 hash - if hash == "" { - hash = gid - } if name == "" { // 尝试从 files 获取文件名 if files, ok := item["files"].([]interface{}); ok && len(files) > 0 { if f, ok := files[0].(map[string]interface{}); ok { - paths, ok := f["path"].([]interface{}) - if ok && len(paths) > 0 { - name = strVal(paths[len(paths)-1]) + filePath := strVal(f["path"]) + if filePath != "" { + contentPath = filePath + name = downloaderPathBase(filePath) } if name == "" { name = strVal(f["uris"]) @@ -69,24 +66,25 @@ func (a *Aria2Adapter) parseSingleItem(item map[string]interface{}) *TorrentInfo var progress float64 if totalLength > 0 { - progress = float64(completedLength) / float64(totalLength) * 100 + progress = float64(completedLength) / float64(totalLength) } // Aria2 状态映射 - state := aria2StatusStr(status) + state := canonicalTorrentState(aria2StatusStr(status), progress) return &TorrentInfo{ - Hash: hash, - Name: name, - Size: totalLength, - Progress: progress, - DLSpeed: dlSpeed, - UPSpeed: upSpeed, - State: state, - SavePath: dir, - NumSeeds: numSeeders, - NumLeechs: aria2MaxInt(connections-numSeeders, 0), - AddedOn: time.Now(), + Hash: gid, + Name: name, + Size: totalLength, + Progress: progress, + DLSpeed: dlSpeed, + UPSpeed: upSpeed, + State: state, + SavePath: dir, + NumSeeds: numSeeders, + NumLeechs: aria2MaxInt(connections-numSeeders, 0), + AddedOn: time.Now(), + ContentPath: contentPath, } } diff --git a/internal/service/download_active_paths.go b/internal/service/download_active_paths.go index 86fe8db..f965e0c 100644 --- a/internal/service/download_active_paths.go +++ b/internal/service/download_active_paths.go @@ -10,14 +10,14 @@ import ( const activeDownloadSnapshotFallbackAge = 2 * time.Minute func (d *DownloadService) ActiveDownloadPaths(ctx context.Context) []string { - if d == nil || d.qb == nil { + if d == nil { return nil } - live, err := d.qb.List(ctx, "") - if err != nil { + live, err := d.listLiveTorrents(ctx, "") + if err != nil && len(live) == 0 { live = d.LiveTorrentSnapshot(activeDownloadSnapshotFallbackAge) if d.log != nil && len(live) == 0 { - d.log.Debug("active download guard could not list qbittorrent and has no fresh snapshot", zap.Error(err)) + d.log.Debug("active download guard could not list download clients and has no fresh snapshot", zap.Error(err)) } } return activeDownloadPathCandidates(live, d.downloadPathMappings(ctx)) diff --git a/internal/service/download_adapter.go b/internal/service/download_adapter.go index c68b010..275fae3 100644 --- a/internal/service/download_adapter.go +++ b/internal/service/download_adapter.go @@ -29,21 +29,42 @@ type DownloadAdapter interface { GetInfo(ctx context.Context, hash string) (*TorrentInfo, error) } +// TorrentFileDownloadAdapter is implemented by clients that can accept the +// application-fetched .torrent payload instead of fetching a private URL. +type TorrentFileDownloadAdapter interface { + AddTorrentFile(ctx context.Context, data []byte, name, savePath string) (string, error) +} + +// CategorizedTorrentDownloadAdapter is implemented by qBittorrent, whose +// native category is part of MediaStationGo's automatic classification flow. +type CategorizedTorrentDownloadAdapter interface { + AddTorrentWithCategory(ctx context.Context, url, savePath, category string) (string, error) + AddTorrentFileWithCategory(ctx context.Context, data []byte, name, savePath, category string) (string, error) +} + +// TorrentRelocateAdapter is intentionally qBittorrent-only: qB can move +// payload data while preserving its seeding task through setLocation. +type TorrentRelocateAdapter interface { + Relocate(ctx context.Context, hash, location string) error +} + // TorrentInfo 是各种下载客户端的种子信息的统一表示。 type TorrentInfo struct { - Hash string `json:"hash"` - Name string `json:"name"` - Size int64 `json:"size"` - Progress float64 `json:"progress"` - DLSpeed int64 `json:"dl_speed"` - UPSpeed int64 `json:"up_speed"` - State string `json:"state"` - SavePath string `json:"save_path"` - NumSeeds int `json:"num_seeds"` - NumLeechs int `json:"num_leechs"` - AddedOn time.Time `json:"added_on"` - Category string `json:"category"` - Tags string `json:"tags"` + Hash string `json:"hash"` + Name string `json:"name"` + Size int64 `json:"size"` + Progress float64 `json:"progress"` + DLSpeed int64 `json:"dl_speed"` + UPSpeed int64 `json:"up_speed"` + State string `json:"state"` + SavePath string `json:"save_path"` + NumSeeds int `json:"num_seeds"` + NumLeechs int `json:"num_leechs"` + AddedOn time.Time `json:"added_on"` + Category string `json:"category"` + Tags string `json:"tags"` + ContentPath string `json:"content_path"` + CompletionOn int64 `json:"completion_on"` } // DownloadClientConfig 是下载客户端的连接配置。 diff --git a/internal/service/download_adapter_file_test.go b/internal/service/download_adapter_file_test.go new file mode 100644 index 0000000..c7a9479 --- /dev/null +++ b/internal/service/download_adapter_file_test.go @@ -0,0 +1,219 @@ +package service + +import ( + "encoding/base64" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "testing" +) + +func TestAria2AdapterRemoveClearsTaskAndResult(t *testing.T) { + var methods []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req aria2Request + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode aria2 request: %v", err) + return + } + if req.Method != "aria2.getVersion" { + methods = append(methods, req.Method) + } + result := interface{}("OK") + if req.Method == "aria2.getVersion" { + result = map[string]interface{}{"version": "1.37"} + } + _ = json.NewEncoder(w).Encode(map[string]interface{}{"jsonrpc": "2.0", "id": req.ID, "result": result}) + })) + defer server.Close() + + adapter := NewAria2Adapter() + if err := adapter.Initialize(t.Context(), DownloadClientConfig{Host: server.URL}); err != nil { + t.Fatal(err) + } + if err := adapter.Remove(t.Context(), "aria2-gid", true); err != nil { + t.Fatal(err) + } + if len(methods) != 2 || methods[0] != "aria2.remove" || methods[1] != "aria2.removeDownloadResult" { + t.Fatalf("methods = %#v", methods) + } +} + +func TestAria2AdapterListReportsConnectionFailures(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req aria2Request + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode aria2 request: %v", err) + return + } + response := map[string]interface{}{"jsonrpc": "2.0", "id": req.ID, "result": map[string]interface{}{"version": "1.37"}} + if req.Method != "aria2.getVersion" { + response = map[string]interface{}{"jsonrpc": "2.0", "id": req.ID, "error": map[string]interface{}{"code": 1, "message": "offline"}} + } + _ = json.NewEncoder(w).Encode(response) + })) + defer server.Close() + + adapter := NewAria2Adapter() + if err := adapter.Initialize(t.Context(), DownloadClientConfig{Host: server.URL}); err != nil { + t.Fatal(err) + } + _, err := adapter.List(t.Context(), "") + if err == nil || !errors.Is(err, errAria2ListUnavailable) { + t.Fatalf("err = %v", err) + } +} + +func TestTransmissionAdapterAddsTorrentFileAsMetainfo(t *testing.T) { + payload := []byte("d4:infod4:name5:movieee") + var added map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + if req.Method != "torrent-add" { + t.Errorf("transmission method = %q, want torrent-add", req.Method) + return + } + added = req.Arguments + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{ + Result: "success", + Arguments: map[string]interface{}{ + "torrent-added": map[string]interface{}{"hashString": "transmission-file-hash"}, + }, + }) + })) + defer server.Close() + + adapter := NewTransmissionAdapter() + if err := adapter.Initialize(t.Context(), DownloadClientConfig{Host: server.URL}); err != nil { + t.Fatal(err) + } + hash, err := adapter.AddTorrentFile(t.Context(), payload, "movie.torrent", "/downloads/movies") + if err != nil { + t.Fatal(err) + } + if hash != "transmission-file-hash" { + t.Fatalf("hash = %q", hash) + } + if added["metainfo"] != base64.StdEncoding.EncodeToString(payload) { + t.Fatalf("metainfo = %#v", added["metainfo"]) + } + if added["download-dir"] != "/downloads/movies" { + t.Fatalf("download-dir = %#v", added["download-dir"]) + } + if _, ok := added["filename"]; ok { + t.Fatalf("torrent file request unexpectedly included filename: %#v", added) + } +} + +func TestAria2AdapterAddsTorrentFileWithAddTorrent(t *testing.T) { + payload := []byte("d4:infod4:name5:movieee") + var addParams []interface{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req aria2Request + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode aria2 request: %v", err) + return + } + result := interface{}(map[string]interface{}{"version": "1.37"}) + if req.Method == "aria2.addTorrent" { + addParams = req.Params + result = "aria2-gid" + } + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "jsonrpc": "2.0", + "id": req.ID, + "result": result, + }) + })) + defer server.Close() + + adapter := NewAria2Adapter() + if err := adapter.Initialize(t.Context(), DownloadClientConfig{Host: server.URL, Password: "secret"}); err != nil { + t.Fatal(err) + } + gid, err := adapter.AddTorrentFile(t.Context(), payload, "movie.torrent", "/downloads/movies") + if err != nil { + t.Fatal(err) + } + if gid != "aria2-gid" { + t.Fatalf("gid = %q", gid) + } + if len(addParams) != 4 { + t.Fatalf("aria2 params = %#v", addParams) + } + if addParams[0] != "token:secret" || addParams[1] != base64.StdEncoding.EncodeToString(payload) { + t.Fatalf("aria2 params = %#v", addParams) + } + options, ok := addParams[3].(map[string]interface{}) + if !ok || options["dir"] != "/downloads/movies" { + t.Fatalf("aria2 options = %#v", addParams[3]) + } +} + +func TestQBitAdapterAddsTorrentFileWithCategory(t *testing.T) { + payload := []byte("d4:infod4:name5:movieee") + var gotCategory, gotSavePath, gotName string + var gotPayload []byte + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v2/auth/login": + _, _ = w.Write([]byte("Ok.")) + case "/api/v2/torrents/add": + reader, err := r.MultipartReader() + if err != nil { + t.Errorf("multipart reader: %v", err) + return + } + for { + part, err := reader.NextPart() + if err == io.EOF { + break + } + if err != nil { + t.Errorf("multipart next part: %v", err) + return + } + value, _ := io.ReadAll(part) + switch part.FormName() { + case "torrents": + gotName = part.FileName() + gotPayload = value + case "savepath": + gotSavePath = string(value) + case "category": + gotCategory = string(value) + } + } + _, _ = w.Write([]byte("Ok.")) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + adapter := NewQBitAdapter() + if err := adapter.Initialize(t.Context(), DownloadClientConfig{Host: server.URL}); err != nil { + t.Fatal(err) + } + hash, err := adapter.AddTorrentFileWithCategory(t.Context(), payload, "movie.torrent", "/downloads/movies", "Movies") + if err != nil { + t.Fatal(err) + } + if hash != torrentInfoHash(payload) { + t.Fatalf("hash = %q, want %q", hash, torrentInfoHash(payload)) + } + if gotName != "movie.torrent" || string(gotPayload) != string(payload) || gotSavePath != "/downloads/movies" || gotCategory != "Movies" { + t.Fatalf("multipart = name %q payload %q savepath %q category %q", gotName, gotPayload, gotSavePath, gotCategory) + } +} diff --git a/internal/service/download_add.go b/internal/service/download_add.go index 29659f5..96df9af 100644 --- a/internal/service/download_add.go +++ b/internal/service/download_add.go @@ -54,11 +54,16 @@ func (d *DownloadService) AddDownloadWithMeta(ctx context.Context, userID, urlSt return existing, ErrDownloadAlreadyExists } _ = d.ReloadConfig(ctx) - if !d.qb.IsConfigured() { + target, err := d.defaultDownloadTarget(ctx) + if err != nil { return nil, d.defaultDownloaderNotConfiguredError(ctx) } - if d.torrentExistsByIdentity(ctx, req) { - task, err := d.createTask(ctx, userID, urlStr, req.savePath, req.meta) + if liveTorrent, ok := d.findLiveTorrentByIdentity(ctx, urlStr, req); ok { + existingTarget := target + if strings.TrimSpace(liveTorrent.ClientID) != "" { + existingTarget = downloadTarget{clientID: liveTorrent.ClientID, typ: firstNonEmpty(liveTorrent.Source, target.typ)} + } + task, err := d.createTask(ctx, userID, urlStr, req.savePath, req.meta, existingTarget, liveTorrent.Hash) if err != nil { return nil, err } @@ -67,13 +72,14 @@ func (d *DownloadService) AddDownloadWithMeta(ctx context.Context, userID, urlSt } return task, ErrDownloadAlreadyExists } - if err := d.addPreparedDownloadToClient(ctx, urlStr, &req); err != nil { + externalID, err := d.addPreparedDownloadToClient(ctx, urlStr, &req, target) + if err != nil { if errors.Is(err, ErrDownloadAlreadyExists) && strings.TrimSpace(req.meta.SubscriptionID) != "" { - return d.createTask(ctx, userID, urlStr, req.savePath, req.meta) + return d.createTask(ctx, userID, urlStr, req.savePath, req.meta, target, externalID) } return nil, err } - return d.createTask(ctx, userID, urlStr, req.savePath, req.meta) + return d.createTask(ctx, userID, urlStr, req.savePath, req.meta, target, externalID) } func (d *DownloadService) prepareDownloadAdd(ctx context.Context, urlStr, savePath string, meta DownloadTaskMeta) (downloadAddRequest, error) { @@ -100,28 +106,70 @@ func (d *DownloadService) prepareDownloadAdd(ctx context.Context, urlStr, savePa }, nil } -func (d *DownloadService) addPreparedDownloadToClient(ctx context.Context, urlStr string, req *downloadAddRequest) error { +func (d *DownloadService) addPreparedDownloadToClient(ctx context.Context, urlStr string, req *downloadAddRequest, target downloadTarget) (string, error) { var siteFetchErr error if d.site != nil { if data, name, err := d.site.FetchTorrentFile(ctx, urlStr); err == nil { - if err := d.qb.AddTorrentFileWithCategory(ctx, data, name, req.savePath, req.qbitCategory); err != nil { - return err - } - if strings.TrimSpace(req.meta.Title) == "" { - req.meta.Title = strings.TrimSuffix(name, path.Ext(name)) - } - return nil + return d.addTorrentFileToTarget(ctx, data, name, req, target) } else { siteFetchErr = err } } - if err := d.qb.AddTorrentWithCategory(ctx, urlStr, req.savePath, req.qbitCategory); err != nil { - if siteFetchErr != nil && !strings.Contains(siteFetchErr.Error(), "no matching PT site") { - return errors.Join(err, siteFetchErr) - } - return err + externalID, err := d.addTorrentURLToTarget(ctx, urlStr, req, target) + if err != nil { + return externalID, joinTorrentFetchError(err, siteFetchErr) } - return nil + return externalID, nil +} + +func (d *DownloadService) addTorrentFileToTarget(ctx context.Context, data []byte, name string, req *downloadAddRequest, target downloadTarget) (string, error) { + if target.legacyQB { + if err := d.qb.AddTorrentFileWithCategory(ctx, data, name, req.savePath, req.qbitCategory); err != nil { + return "", err + } + return torrentInfoHash(data), nil + } + if categorized, ok := target.adapter.(CategorizedTorrentDownloadAdapter); ok { + externalID, err := categorized.AddTorrentFileWithCategory(ctx, data, name, req.savePath, req.qbitCategory) + setFetchedTorrentTitle(req, name, err) + return externalID, err + } + fileAdapter, ok := target.adapter.(TorrentFileDownloadAdapter) + if !ok { + return "", errors.New("configured downloader does not accept torrent files") + } + externalID, err := fileAdapter.AddTorrentFile(ctx, data, name, req.savePath) + setFetchedTorrentTitle(req, name, err) + return externalID, err +} + +func (d *DownloadService) addTorrentURLToTarget(ctx context.Context, urlStr string, req *downloadAddRequest, target downloadTarget) (string, error) { + if target.legacyQB { + if err := d.qb.AddTorrentWithCategory(ctx, urlStr, req.savePath, req.qbitCategory); err != nil { + return "", err + } + return torrentURLInfoHash(urlStr), nil + } + if categorized, ok := target.adapter.(CategorizedTorrentDownloadAdapter); ok { + return categorized.AddTorrentWithCategory(ctx, urlStr, req.savePath, req.qbitCategory) + } + if strings.HasPrefix(strings.ToLower(strings.TrimSpace(urlStr)), "magnet:") { + return target.adapter.AddMagnet(ctx, urlStr, req.savePath) + } + return target.adapter.AddTorrent(ctx, urlStr, req.savePath) +} + +func setFetchedTorrentTitle(req *downloadAddRequest, name string, addErr error) { + if addErr == nil && req != nil && strings.TrimSpace(req.meta.Title) == "" { + req.meta.Title = strings.TrimSuffix(name, path.Ext(name)) + } +} + +func joinTorrentFetchError(addErr, fetchErr error) error { + if fetchErr != nil && !strings.Contains(fetchErr.Error(), "no matching PT site") { + return errors.Join(addErr, fetchErr) + } + return addErr } func (d *DownloadService) resolveDownloadSavePath(ctx context.Context, explicitSavePath string, meta DownloadTaskMeta, autoClassify bool) (string, string) { @@ -150,7 +198,7 @@ func (d *DownloadService) resolveDownloadSavePath(ctx context.Context, explicitS return downloadSavePathCategoryRoot(base, sanitizeFilename(category)), category } -func (d *DownloadService) createTask(ctx context.Context, userID, urlStr, savePath string, meta DownloadTaskMeta) (*model.DownloadTask, error) { +func (d *DownloadService) createTask(ctx context.Context, userID, urlStr, savePath string, meta DownloadTaskMeta, target downloadTarget, externalID string) (*model.DownloadTask, error) { title := strings.TrimSpace(meta.Title) if title == "" { title = publicDownloadTitle(urlStr) @@ -158,7 +206,9 @@ func (d *DownloadService) createTask(ctx context.Context, userID, urlStr, savePa t := &model.DownloadTask{ UserID: userID, SubscriptionID: strings.TrimSpace(meta.SubscriptionID), - Source: "qbittorrent", + DownloadClientID: target.clientID, + ExternalID: strings.TrimSpace(externalID), + Source: target.typ, URL: urlStr, Title: title, PosterURL: meta.PosterURL, diff --git a/internal/service/download_add_dedup.go b/internal/service/download_add_dedup.go index eb488b9..081de44 100644 --- a/internal/service/download_add_dedup.go +++ b/internal/service/download_add_dedup.go @@ -42,10 +42,10 @@ func (d *DownloadService) findExistingDownloadTask(ctx context.Context, req down func (d *DownloadService) subscriptionDownloadTaskStillLive(ctx context.Context, row model.DownloadTask) bool { live, ok := d.liveTorrentSnapshot(30 * time.Second) - if !ok && d != nil && d.qb != nil && d.qb.IsConfigured() { + if !ok && d != nil { var err error - live, err = d.qb.List(ctx, "") - if err != nil { + live, err = d.listLiveTorrents(ctx, "") + if err != nil && len(live) == 0 { return true } ok = true @@ -62,6 +62,12 @@ func (d *DownloadService) subscriptionDownloadTaskStillLive(ctx context.Context, } func downloadTaskMatchesLiveTorrent(row model.DownloadTask, torrent QBitTorrent) bool { + if strings.TrimSpace(row.DownloadClientID) != "" && strings.TrimSpace(torrent.ClientID) != "" && row.DownloadClientID != torrent.ClientID { + return false + } + if strings.TrimSpace(row.ExternalID) != "" { + return strings.EqualFold(strings.TrimSpace(row.ExternalID), strings.TrimSpace(torrent.Hash)) + } torrentName := strings.TrimSpace(torrent.Name) if torrentName == "" { return false @@ -144,31 +150,37 @@ func downloadTaskInSubscriptionScope(row model.DownloadTask, req downloadAddRequ return sameOrChildPath(rowSavePath, requestSavePath) || sameOrChildPath(requestSavePath, rowSavePath) } -func (d *DownloadService) torrentExistsByIdentity(ctx context.Context, req downloadAddRequest) bool { +func (d *DownloadService) findLiveTorrentByIdentity(ctx context.Context, downloadURL string, req downloadAddRequest) (QBitTorrent, bool) { query := downloadTaskIdentityKey(req.title) - if query == "" { - return false + requestHash := torrentURLInfoHash(downloadURL) + if query == "" && requestHash == "" { + return QBitTorrent{}, false } - live, err := d.qb.List(ctx, "") + live, err := d.listLiveTorrents(ctx, "") if err != nil { - return false + if len(live) == 0 { + return QBitTorrent{}, false + } } for _, torrent := range live { if !torrentInDownloadRequestScope(torrent, req) { continue } + if requestHash != "" && strings.EqualFold(requestHash, strings.TrimSpace(torrent.Hash)) { + return torrent, true + } if downloadTaskCoversAddRequest(torrent.Name, req) { - return true + return torrent, true } current := downloadTaskIdentityKey(torrent.Name) if current == "" { continue } if current == query { - return true + return torrent, true } } - return false + return QBitTorrent{}, false } func torrentInDownloadRequestScope(torrent QBitTorrent, req downloadAddRequest) bool { diff --git a/internal/service/download_add_dedup_test.go b/internal/service/download_add_dedup_test.go index a291129..15bcace 100644 --- a/internal/service/download_add_dedup_test.go +++ b/internal/service/download_add_dedup_test.go @@ -197,6 +197,9 @@ func TestAddDownloadWithMetaTracksExistingQBTorrentForSubscription(t *testing.T) if task == nil || task.SubscriptionID != "sub-nanyang" { t.Fatalf("task = %#v, want subscription tracking task", task) } + if task.DownloadClientID != legacyQBitDownloadClientID || task.ExternalID != hash || task.Source != "qbittorrent" { + t.Fatalf("tracked downloader identity = %#v", task) + } if got := atomic.LoadInt32(&addCalls); got != 0 { t.Fatalf("qb add calls = %d, want 0 because infohash already exists", got) } diff --git a/internal/service/download_add_test.go b/internal/service/download_add_test.go index 10d5c04..d43cf95 100644 --- a/internal/service/download_add_test.go +++ b/internal/service/download_add_test.go @@ -20,6 +20,13 @@ func TestPublicDownloadTitleUsesMagnetDisplayName(t *testing.T) { } } +func TestTorrentURLInfoHashNormalizesBase32BTIH(t *testing.T) { + got := torrentURLInfoHash("magnet:?xt=urn:btih:AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA") + if got != "0000000000000000000000000000000000000000" { + t.Fatalf("hash = %q", got) + } +} + func configureTestDefaultQB(t *testing.T, repos *repository.Container, baseURL string) { t.Helper() if err := repos.DownloadClient.Create(t.Context(), &model.DownloadClient{ diff --git a/internal/service/download_completion.go b/internal/service/download_completion.go index 0caf285..05250f8 100644 --- a/internal/service/download_completion.go +++ b/internal/service/download_completion.go @@ -4,7 +4,6 @@ import ( "context" "errors" "os" - "path/filepath" "strings" "go.uber.org/zap" @@ -116,11 +115,13 @@ func (d *DownloadService) completedTorrentTask(ctx context.Context, torrent QBit return nil, false } taskByKey := tasksByTorrentIdentity(rows) - if task, ok := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey); ok { + if task, ok := findMatchingTaskForTorrent(torrent, taskByKey); ok { return &task, true } if strings.TrimSpace(torrent.ContentPath) != "" { - if task, ok := findMatchingTaskByTorrentIdentity(filepath.Base(torrent.ContentPath), taskByKey); ok { + pathTorrent := torrent + pathTorrent.Name = downloaderPathBase(torrent.ContentPath) + if task, ok := findMatchingTaskForTorrent(pathTorrent, taskByKey); ok { return &task, true } } @@ -151,7 +152,7 @@ func (d *DownloadService) completedTorrentSource(ctx context.Context, torrent QB mappings := d.downloadPathMappings(ctx) for _, candidate := range []string{ torrent.ContentPath, - filepath.Join(torrent.SavePath, torrent.Name), + downloaderPayloadPath(torrent.SavePath, torrent.Name), } { clean := strings.TrimSpace(candidate) if clean == "" || clean == "." { diff --git a/internal/service/download_completion_state.go b/internal/service/download_completion_state.go index e7e7889..db9da35 100644 --- a/internal/service/download_completion_state.go +++ b/internal/service/download_completion_state.go @@ -49,7 +49,7 @@ func qbitTorrentCompleted(torrent QBitTorrent) bool { } state := strings.ToLower(strings.TrimSpace(torrent.State)) switch state { - case "completed", "uploading", "stalledup", "pausedup", "queuedup", "forcedup": + case "completed", "complete", "seeding", "uploading", "stalledup", "pausedup", "queuedup", "forcedup": return true default: return false @@ -139,9 +139,16 @@ func completedTorrentNotifySettingKey(torrent QBitTorrent) string { func completedTorrentQueueKey(torrent QBitTorrent) string { hash := strings.ToLower(strings.TrimSpace(torrent.Hash)) if hash != "" { + owner := strings.ToLower(firstNonEmpty(torrent.ClientID, torrent.Source)) + if owner != "" { + return owner + "|" + hash + } return hash } parts := []string{torrent.Name, torrent.ContentPath, torrent.SavePath} + if owner := firstNonEmpty(torrent.ClientID, torrent.Source); owner != "" { + parts = append([]string{owner}, parts...) + } for i := range parts { parts[i] = strings.TrimSpace(parts[i]) } @@ -153,10 +160,10 @@ func completedTorrentQueueKey(torrent QBitTorrent) string { } func (d *DownloadService) syncDownloadTaskProgress(ctx context.Context, torrent QBitTorrent, taskByKey map[string]model.DownloadTask) { - if d == nil || d.repo == nil || d.repo.DB == nil || strings.TrimSpace(torrent.Name) == "" { + if d == nil || d.repo == nil || d.repo.DB == nil { return } - matched, ok := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey) + matched, ok := findMatchingTaskForTorrent(torrent, taskByKey) if !ok { return } @@ -174,6 +181,15 @@ func (d *DownloadService) syncDownloadTaskProgress(ctx context.Context, torrent if status != "" && status != matched.Status { updates["status"] = status } + if strings.TrimSpace(matched.DownloadClientID) == "" && strings.TrimSpace(torrent.ClientID) != "" { + updates["download_client_id"] = strings.TrimSpace(torrent.ClientID) + } + if strings.TrimSpace(matched.ExternalID) == "" && strings.TrimSpace(torrent.Hash) != "" { + updates["external_id"] = strings.TrimSpace(torrent.Hash) + } + if strings.TrimSpace(torrent.Source) != "" && matched.Source != strings.TrimSpace(torrent.Source) { + updates["source"] = strings.TrimSpace(torrent.Source) + } if len(updates) == 0 { return } @@ -192,16 +208,79 @@ func tasksByIdentity(rows []model.DownloadTask) map[string]model.DownloadTask { } func tasksByTorrentIdentity(rows []model.DownloadTask) map[string]model.DownloadTask { - out := make(map[string]model.DownloadTask, len(rows)) + out := make(map[string]model.DownloadTask, len(rows)*4) for _, row := range rows { key := normalizeTorrentName(row.Title) if key != "" { - out[key] = row + setDownloadTaskIndex(out, key, row) + setDownloadTaskIndex(out, downloadTaskClientTitleKey(row.DownloadClientID, key), row) + } + if externalID := strings.TrimSpace(row.ExternalID); externalID != "" { + setDownloadTaskIndex(out, downloadTaskExternalKey(row.DownloadClientID, externalID), row) + setDownloadTaskIndex(out, downloadTaskAnyExternalKey(externalID), row) } } return out } +func setDownloadTaskIndex(index map[string]model.DownloadTask, key string, row model.DownloadTask) { + if key == "" { + return + } + if _, exists := index[key]; !exists { + index[key] = row + } +} + +func downloadTaskExternalKey(clientID, externalID string) string { + clientID = strings.ToLower(strings.TrimSpace(clientID)) + externalID = strings.ToLower(strings.TrimSpace(externalID)) + if clientID == "" || externalID == "" { + return "" + } + return "\x00external:" + clientID + ":" + externalID +} + +func downloadTaskAnyExternalKey(externalID string) string { + externalID = strings.ToLower(strings.TrimSpace(externalID)) + if externalID == "" { + return "" + } + return "\x00external-any:" + externalID +} + +func downloadTaskClientTitleKey(clientID, titleKey string) string { + clientID = strings.ToLower(strings.TrimSpace(clientID)) + titleKey = strings.TrimSpace(titleKey) + if clientID == "" || titleKey == "" { + return "" + } + return "\x00client-title:" + clientID + ":" + titleKey +} + +func findMatchingTaskForTorrent(torrent QBitTorrent, taskByKey map[string]model.DownloadTask) (model.DownloadTask, bool) { + if row, ok := taskByKey[downloadTaskExternalKey(torrent.ClientID, torrent.Hash)]; ok { + return row, true + } + if row, ok := taskByKey[downloadTaskAnyExternalKey(torrent.Hash)]; ok { + if strings.TrimSpace(row.DownloadClientID) == "" || strings.TrimSpace(torrent.ClientID) == "" || row.DownloadClientID == torrent.ClientID { + return row, true + } + } + titleKey := normalizeTorrentName(torrent.Name) + if row, ok := taskByKey[downloadTaskClientTitleKey(torrent.ClientID, titleKey)]; ok { + return row, true + } + row, ok := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey) + if !ok { + return model.DownloadTask{}, false + } + if strings.TrimSpace(torrent.ClientID) != "" && strings.TrimSpace(row.DownloadClientID) != "" && row.DownloadClientID != torrent.ClientID { + return model.DownloadTask{}, false + } + return row, true +} + func findMatchingTaskByIdentity(title string, taskByKey map[string]model.DownloadTask) (model.DownloadTask, bool) { key := downloadTaskIdentityKey(title) if key == "" { @@ -211,6 +290,9 @@ func findMatchingTaskByIdentity(title string, taskByKey map[string]model.Downloa return row, true } for currentKey, row := range taskByKey { + if strings.HasPrefix(currentKey, "\x00") { + continue + } if strings.Contains(key, currentKey) || strings.Contains(currentKey, key) { return row, true } @@ -227,6 +309,9 @@ func findMatchingTaskByTorrentIdentity(title string, taskByKey map[string]model. return row, true } for currentKey, row := range taskByKey { + if strings.HasPrefix(currentKey, "\x00") { + continue + } if strings.Contains(key, currentKey) || strings.Contains(currentKey, key) { return row, true } diff --git a/internal/service/download_config_test.go b/internal/service/download_config_test.go index 17ec090..6414920 100644 --- a/internal/service/download_config_test.go +++ b/internal/service/download_config_test.go @@ -1,6 +1,7 @@ package service import ( + "encoding/json" "net/http" "net/http/httptest" "strings" @@ -166,6 +167,33 @@ func TestReloadConfigUsesSoleEnabledQBitWhenNoExplicitDefault(t *testing.T) { } } +func TestReloadConfigDoesNotOverrideExplicitTransmissionDefaultWithQBit(t *testing.T) { + db := newServiceTestDB(t, &model.DownloadClient{}, &model.Setting{}) + repos := repository.New(db) + transmission := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: "http://127.0.0.1:9091", IsDefault: true, Enabled: true} + qb := &model.DownloadClient{Name: "qB", Type: "qbittorrent", Host: "http://127.0.0.1:8080", Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), transmission); err != nil { + t.Fatal(err) + } + if err := repos.DownloadClient.Create(t.Context(), qb); err != nil { + t.Fatal(err) + } + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + if err := svc.ReloadConfig(t.Context()); err != nil { + t.Fatal(err) + } + if svc.qb.IsConfigured() { + t.Fatal("legacy qB client was configured despite explicit Transmission default") + } + selected, err := repos.DownloadClient.FindDefault(t.Context()) + if err != nil { + t.Fatal(err) + } + if selected == nil || selected.ID != transmission.ID { + t.Fatalf("default client = %#v", selected) + } +} + func TestAddDownloadWithMetaFailsClosedWhenNoDownloaderConfigured(t *testing.T) { db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) repos := repository.New(db) @@ -306,29 +334,57 @@ func TestReloadConfigManagedModeDoesNotFallbackToLegacyWithoutRows(t *testing.T) } } -func TestAddDownloadWithMetaExplainsUnsupportedEnabledDownloader(t *testing.T) { +func TestAddDownloadWithMetaUsesEnabledAria2Downloader(t *testing.T) { + var addCalls int32 + aria2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req aria2Request + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode aria2 request: %v", err) + return + } + var result interface{} = map[string]interface{}{"version": "1.37"} + switch req.Method { + case "aria2.tellActive", "aria2.tellWaiting", "aria2.tellStopped": + result = []interface{}{} + case "aria2.addUri": + atomic.AddInt32(&addCalls, 1) + result = "aria2-gid" + } + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "jsonrpc": "2.0", + "id": req.ID, + "result": result, + }) + })) + defer aria2.Close() + db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) repos := repository.New(db) if err := repos.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil { t.Fatal(err) } - if err := repos.DownloadClient.Create(t.Context(), &model.DownloadClient{ + client := &model.DownloadClient{ Name: "aria2", Type: "aria2", - Host: "http://127.0.0.1:6800", + Host: aria2.URL, Enabled: true, - }); err != nil { + } + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { t.Fatal(err) } svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) - _, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:ffffffffffffffffffffffffffffffffffffffff&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{ + svc.SetDownloadManager(NewDownloadManager(zap.NewNop(), repos, nil)) + task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:ffffffffffffffffffffffffffffffffffffffff&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{ Title: "Movie 2026 1080p", }) - if err == nil { - t.Fatal("expected unsupported downloader error") + if err != nil { + t.Fatal(err) } - if !strings.Contains(err.Error(), "订阅投递目前需要 qBittorrent") { - t.Fatalf("err = %v, want qBittorrent guidance", err) + if task.Source != "aria2" || task.DownloadClientID != client.ID || task.ExternalID != "aria2-gid" { + t.Fatalf("task = %#v", task) + } + if got := atomic.LoadInt32(&addCalls); got != 1 { + t.Fatalf("add calls = %d", got) } } diff --git a/internal/service/download_control_test.go b/internal/service/download_control_test.go new file mode 100644 index 0000000..dc4f49e --- /dev/null +++ b/internal/service/download_control_test.go @@ -0,0 +1,75 @@ +package service + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "reflect" + "testing" + + "go.uber.org/zap" + + "github.com/ShukeBta/MediaStationGo/internal/model" + "github.com/ShukeBta/MediaStationGo/internal/repository" +) + +func TestPauseAndResumeRouteThroughTaskDownloadClient(t *testing.T) { + var methods []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + methods = append(methods, req.Method) + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{Result: "success", Arguments: map[string]interface{}{}}) + })) + defer server.Close() + + db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{}) + repos := repository.New(db) + client := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: server.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + task := &model.DownloadTask{ + UserID: "u1", + Source: "transmission", + DownloadClientID: client.ID, + ExternalID: "transmission-hash", + URL: "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + Title: "Controlled Transmission Movie", + Status: "downloading", + } + if err := repos.Download.Create(t.Context(), task); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.SetDownloadManager(manager) + if err := svc.ReloadConfig(t.Context()); err != nil { + t.Fatal(err) + } + + if err := svc.PauseDownloadTask(t.Context(), task.ID); err != nil { + t.Fatal(err) + } + if err := svc.ResumeDownloadTask(t.Context(), task.ID); err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(methods, []string{"torrent-stop", "torrent-start"}) { + t.Fatalf("methods = %#v", methods) + } + var updated model.DownloadTask + if err := db.Where("id = ?", task.ID).First(&updated).Error; err != nil { + t.Fatal(err) + } + if updated.Status != "queued" { + t.Fatalf("status = %q", updated.Status) + } +} diff --git a/internal/service/download_controls.go b/internal/service/download_controls.go new file mode 100644 index 0000000..6da2e6f --- /dev/null +++ b/internal/service/download_controls.go @@ -0,0 +1,105 @@ +package service + +import ( + "context" + "errors" + "strings" + + "github.com/ShukeBta/MediaStationGo/internal/model" +) + +func (d *DownloadService) PauseDownloadTask(ctx context.Context, taskID string) error { + return d.controlDownloadTask(ctx, taskID, "paused", func(target downloadTarget, externalID string) error { + if target.legacyQB { + return d.qb.Pause(ctx, externalID) + } + return target.adapter.Pause(ctx, externalID) + }) +} + +func (d *DownloadService) ResumeDownloadTask(ctx context.Context, taskID string) error { + return d.controlDownloadTask(ctx, taskID, "queued", func(target downloadTarget, externalID string) error { + if target.legacyQB { + return d.qb.Resume(ctx, externalID) + } + return target.adapter.Resume(ctx, externalID) + }) +} + +func (d *DownloadService) controlDownloadTask(ctx context.Context, taskID, status string, operation func(downloadTarget, string) error) error { + taskID = strings.TrimSpace(taskID) + if taskID == "" { + return errors.New("task id is required") + } + var task model.DownloadTask + if d == nil || d.repo == nil || d.repo.DB == nil { + return errors.New("download repository is unavailable") + } + if err := d.repo.DB.WithContext(ctx).Where("id = ?", taskID).First(&task).Error; err != nil { + return err + } + clientID, externalID := d.resolveTaskDownloaderIdentity(ctx, &task) + if externalID == "" { + return errors.New("download task has no client task id") + } + target, err := d.downloadTargetByID(ctx, clientID) + if err != nil { + return err + } + if err := operation(target, externalID); err != nil { + return err + } + return d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}). + Where("id = ?", task.ID). + Updates(map[string]any{ + "status": status, + "download_client_id": clientID, + "external_id": externalID, + "source": firstNonEmpty(target.typ, task.Source), + }).Error +} + +func (d *DownloadService) resolveTaskDownloaderIdentity(ctx context.Context, task *model.DownloadTask) (string, string) { + if task == nil { + return "", "" + } + clientID := strings.TrimSpace(task.DownloadClientID) + persistedExternalID := strings.TrimSpace(task.ExternalID) + externalID := persistedExternalID + if externalID == "" { + externalID = torrentURLInfoHash(task.URL) + } + if clientID != "" && persistedExternalID != "" { + return clientID, externalID + } + live, _ := d.listLiveTorrents(ctx, "") + for _, torrent := range live { + if clientID != "" && torrent.ClientID != clientID { + continue + } + if externalID != "" && strings.EqualFold(torrent.Hash, externalID) { + clientID = firstNonEmpty(clientID, torrent.ClientID) + return clientID, torrent.Hash + } + if downloadTaskMatchesLiveTorrent(*task, torrent) { + clientID = firstNonEmpty(clientID, torrent.ClientID) + externalID = firstNonEmpty(externalID, torrent.Hash) + return clientID, externalID + } + } + if clientID == "" && d.manager != nil { + var matched []managedDownloadTarget + for _, target := range d.manager.targets() { + if strings.TrimSpace(task.Source) == "" || target.client.Type == task.Source { + matched = append(matched, target) + } + } + if len(matched) == 1 { + clientID = matched[0].client.ID + } + } + if clientID == "" && d.qb != nil && d.qb.IsConfigured() && (task.Source == "" || task.Source == "qbittorrent") { + clientID = legacyQBitDownloadClientID + } + return clientID, externalID +} diff --git a/internal/service/download_delete_test.go b/internal/service/download_delete_test.go index e3ada24..2bcbfe4 100644 --- a/internal/service/download_delete_test.go +++ b/internal/service/download_delete_test.go @@ -1,6 +1,7 @@ package service import ( + "encoding/json" "net/http" "net/http/httptest" "sync/atomic" @@ -67,6 +68,79 @@ func TestDeleteMarksMatchingDownloadTaskDeleted(t *testing.T) { } } +func TestDeleteRoutesToRequestedTransmissionClient(t *testing.T) { + var removed map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + arguments := map[string]interface{}{} + switch req.Method { + case "torrent-get": + arguments["torrents"] = []map[string]interface{}{{ + "hashString": "transmission-hash", + "name": "Delete Transmission Movie", + "percentDone": 0.5, + "status": 4, + }} + case "torrent-remove": + removed = req.Arguments + default: + t.Errorf("unexpected transmission method %q", req.Method) + } + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{Result: "success", Arguments: arguments}) + })) + defer server.Close() + + db := newServiceTestDB(t, &model.DownloadTask{}, &model.DownloadClient{}, &model.Setting{}) + repos := repository.New(db) + client := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: server.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + task := &model.DownloadTask{ + UserID: "u1", + Source: "transmission", + DownloadClientID: client.ID, + ExternalID: "transmission-hash", + URL: "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + Title: "Delete Transmission Movie", + Status: "downloading", + Progress: 0.5, + } + if err := repos.Download.Create(t.Context(), task); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.SetDownloadManager(manager) + if err := svc.ReloadConfig(t.Context()); err != nil { + t.Fatal(err) + } + + if err := svc.Delete(t.Context(), "transmission-hash", true, client.ID); err != nil { + t.Fatal(err) + } + ids, ok := removed["ids"].([]interface{}) + if !ok || len(ids) != 1 || ids[0] != "transmission-hash" || removed["delete-local-data"] != true { + t.Fatalf("remove arguments = %#v", removed) + } + var updated model.DownloadTask + if err := db.Where("id = ?", task.ID).First(&updated).Error; err != nil { + t.Fatal(err) + } + if updated.Status != "deleted" { + t.Fatalf("status = %q", updated.Status) + } +} + func TestDeleteMarksMagnetTaskDeletedWhenLiveTorrentNameMissing(t *testing.T) { const hash = "0123456789abcdef0123456789abcdef0123c0de" var deleteCalls int32 diff --git a/internal/service/download_manager_svc.go b/internal/service/download_manager_svc.go index cea39d7..e20abc2 100644 --- a/internal/service/download_manager_svc.go +++ b/internal/service/download_manager_svc.go @@ -8,6 +8,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "sync" "go.uber.org/zap" @@ -25,6 +26,8 @@ type DownloadManager struct { mu sync.RWMutex clients map[string]DownloadAdapter // clientID -> adapter configs map[string]DownloadClientConfig + models map[string]model.DownloadClient + order []string } // NewDownloadManager 创建新的下载管理器。 @@ -35,6 +38,7 @@ func NewDownloadManager(log *zap.Logger, repo *repository.Container, crypto *Cry crypto: crypto, clients: make(map[string]DownloadAdapter), configs: make(map[string]DownloadClientConfig), + models: make(map[string]model.DownloadClient), } } @@ -44,80 +48,106 @@ func (m *DownloadManager) LoadAll(ctx context.Context) error { if err != nil { return err } + if err := m.ensureEnabledDefault(ctx, dbClients); err != nil { + return err + } - m.mu.Lock() - defer m.mu.Unlock() - - // 清空现有 - m.clients = make(map[string]DownloadAdapter, len(dbClients)) - m.configs = make(map[string]DownloadClientConfig, len(dbClients)) + clients := make(map[string]DownloadAdapter, len(dbClients)) + configs := make(map[string]DownloadClientConfig, len(dbClients)) + models := make(map[string]model.DownloadClient, len(dbClients)) + order := make([]string, 0, len(dbClients)) for _, dc := range dbClients { - cfg, err := m.buildConfig(&dc) - if err != nil { - m.log.Warn("failed to build config for download client", - zap.String("id", dc.ID), - zap.String("name", dc.Name), - zap.Error(err), - ) + adapter, cfg, ok := m.initializeClient(ctx, dc) + if !ok { continue } - - adapter := AdapterFactory(dc.Type) - if adapter == nil { - m.log.Warn("unknown download client type", - zap.String("type", dc.Type), - zap.String("id", dc.ID), - ) - continue - } - - if initErr := adapter.Initialize(ctx, cfg); initErr != nil { - // 初始化失败通常是 Docker 启动顺序问题(qBittorrent 还没就绪)。 - // 适配器内部支持按需重新登录(403/未登录时透明重试),所以 - // 这里仍然注册适配器,等下载器上线后自动恢复;此前直接 continue - // 会让该客户端在应用重启前永久不可用,下载完成也无法整理入库。 - m.log.Warn("download client init failed; registered for lazy reconnect", - zap.String("id", dc.ID), - zap.String("name", dc.Name), - zap.Error(initErr), - ) - } - - m.clients[dc.ID] = adapter - m.configs[dc.ID] = cfg + clients[dc.ID] = adapter + configs[dc.ID] = cfg + models[dc.ID] = dc + order = append(order, dc.ID) m.log.Info("download client registered", zap.String("id", dc.ID), zap.String("name", dc.Name), zap.String("type", dc.Type), ) } + m.mu.Lock() + m.clients = clients + m.configs = configs + m.models = models + m.order = order + m.mu.Unlock() return nil } +func (m *DownloadManager) ensureEnabledDefault(ctx context.Context, clients []model.DownloadClient) error { + if len(clients) == 0 { + return nil + } + for i := range clients { + if clients[i].IsDefault { + return nil + } + } + if err := m.repo.DownloadClient.SetDefault(ctx, clients[0].ID); err != nil { + return err + } + clients[0].IsDefault = true + return nil +} + +func (m *DownloadManager) initializeClient(ctx context.Context, dc model.DownloadClient) (DownloadAdapter, DownloadClientConfig, bool) { + cfg, err := m.buildConfig(&dc) + if err != nil { + m.log.Warn("failed to build config for download client", + zap.String("id", dc.ID), + zap.String("name", dc.Name), + zap.Error(err)) + return nil, DownloadClientConfig{}, false + } + adapter := AdapterFactory(dc.Type) + if adapter == nil { + m.log.Warn("unknown download client type", + zap.String("type", dc.Type), + zap.String("id", dc.ID)) + return nil, DownloadClientConfig{}, false + } + if initErr := adapter.Initialize(ctx, cfg); initErr != nil { + // Register configured clients even when the external process is still + // starting; each operation can reconnect once it becomes reachable. + m.log.Warn("download client init failed; registered for lazy reconnect", + zap.String("id", dc.ID), + zap.String("name", dc.Name), + zap.Error(initErr)) + } + return adapter, cfg, true +} + // GetDefault 返回默认下载客户端适配器。 // 如果没有设置默认客户端,返回第一个可用的客户端。 -func (m *DownloadManager) GetDefault() (string, DownloadAdapter, error) { +func (m *DownloadManager) GetDefault(_ context.Context) (*model.DownloadClient, DownloadAdapter, error) { m.mu.RLock() defer m.mu.RUnlock() - // 首先找默认的 - defaultClient, err := m.repo.DownloadClient.FindDefault(context.Background()) - if err != nil { - return "", nil, err + for _, id := range m.order { + client := m.models[id] + if client.IsDefault { + if adapter, ok := m.clients[id]; ok { + copy := client + return ©, adapter, nil + } + } } - if defaultClient != nil { - if adapter, ok := m.clients[defaultClient.ID]; ok { - return defaultClient.ID, adapter, nil + for _, id := range m.order { + if adapter, ok := m.clients[id]; ok { + client := m.models[id] + copy := client + return ©, adapter, nil } } - // 返回第一个可用的 - for id, adapter := range m.clients { - return id, adapter, nil - } - - return "", nil, errors.New("no download client available") + return nil, nil, errors.New("no download client available") } // GetClient 返回指定 ID 的下载客户端适配器。 @@ -131,6 +161,55 @@ func (m *DownloadManager) GetClient(id string) (DownloadAdapter, error) { return adapter, nil } +type managedDownloadTarget struct { + client model.DownloadClient + adapter DownloadAdapter +} + +func (m *DownloadManager) getTarget(id string) (managedDownloadTarget, error) { + m.mu.RLock() + defer m.mu.RUnlock() + adapter, ok := m.clients[id] + if !ok { + return managedDownloadTarget{}, errors.New("download client not found or not initialized") + } + client, ok := m.models[id] + if !ok { + return managedDownloadTarget{}, errors.New("download client metadata not found") + } + return managedDownloadTarget{client: client, adapter: adapter}, nil +} + +func (m *DownloadManager) targets() []managedDownloadTarget { + if m == nil { + return nil + } + m.mu.RLock() + defer m.mu.RUnlock() + out := make([]managedDownloadTarget, 0, len(m.order)) + for _, id := range m.order { + adapter, ok := m.clients[id] + if !ok { + continue + } + client, ok := m.models[id] + if !ok { + continue + } + out = append(out, managedDownloadTarget{client: client, adapter: adapter}) + } + return out +} + +func (m *DownloadManager) hasClients() bool { + if m == nil { + return false + } + m.mu.RLock() + defer m.mu.RUnlock() + return len(m.clients) > 0 +} + // AddClient 动态添加并初始化一个下载客户端。 func (m *DownloadManager) AddClient(ctx context.Context, dc *model.DownloadClient) error { cfg, err := m.buildConfig(dc) @@ -149,8 +228,16 @@ func (m *DownloadManager) AddClient(ctx context.Context, dc *model.DownloadClien m.mu.Lock() defer m.mu.Unlock() + for i, current := range m.order { + if current == dc.ID { + m.order = append(m.order[:i], m.order[i+1:]...) + break + } + } m.clients[dc.ID] = adapter m.configs[dc.ID] = cfg + m.models[dc.ID] = *dc + m.order = append(m.order, dc.ID) return nil } @@ -160,6 +247,13 @@ func (m *DownloadManager) RemoveClient(id string) { defer m.mu.Unlock() delete(m.clients, id) delete(m.configs, id) + delete(m.models, id) + for i, current := range m.order { + if current == id { + m.order = append(m.order[:i], m.order[i+1:]...) + break + } + } } // UpdateClient 更新已有客户端的配置并重新初始化。 @@ -185,30 +279,22 @@ func (m *DownloadManager) TestConnection(ctx context.Context, dc *model.Download // ListAll 获取所有已加载客户端的种子列表。 func (m *DownloadManager) ListAll(ctx context.Context, filter string) (map[string][]TorrentInfo, error) { - m.mu.RLock() - ids := make([]string, 0, len(m.clients)) - for id := range m.clients { - ids = append(ids, id) - } - adapters := make([]DownloadAdapter, 0, len(m.clients)) - for _, id := range ids { - adapters = append(adapters, m.clients[id]) - } - m.mu.RUnlock() - result := make(map[string][]TorrentInfo) - for i, id := range ids { - list, err := adapters[i].List(ctx, filter) + var listErrs []error + for _, target := range m.targets() { + list, err := target.adapter.List(ctx, filter) if err != nil { m.log.Warn("failed to list torrents from client", - zap.String("id", id), + zap.String("id", target.client.ID), zap.Error(err), ) - continue + listErrs = append(listErrs, fmt.Errorf("%s (%s): %w", target.client.Name, target.client.Type, err)) + } + if len(list) > 0 || err == nil { + result[target.client.ID] = list } - result[id] = list } - return result, nil + return result, errors.Join(listErrs...) } // GetAdapterTypes 返回支持的下载客户端类型列表。 diff --git a/internal/service/download_multi_client_test.go b/internal/service/download_multi_client_test.go new file mode 100644 index 0000000..4760206 --- /dev/null +++ b/internal/service/download_multi_client_test.go @@ -0,0 +1,417 @@ +package service + +import ( + "encoding/base64" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "sort" + "sync" + "sync/atomic" + "testing" + "time" + + "go.uber.org/zap" + + "github.com/ShukeBta/MediaStationGo/internal/model" + "github.com/ShukeBta/MediaStationGo/internal/repository" +) + +func TestAddDownloadUsesDefaultTransmissionClient(t *testing.T) { + var mu sync.Mutex + var added map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + switch req.Method { + case "torrent-get": + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{Result: "success", Arguments: map[string]interface{}{"torrents": []interface{}{}}}) + case "torrent-add": + mu.Lock() + added = req.Arguments + mu.Unlock() + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{ + Result: "success", + Arguments: map[string]interface{}{ + "torrent-added": map[string]interface{}{"hashString": "transmission-hash", "name": "Movie 2026"}, + }, + }) + default: + t.Errorf("unexpected transmission method %q", req.Method) + } + })) + defer server.Close() + + db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) + repos := repository.New(db) + client := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: server.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.SetDownloadManager(manager) + + task, err := svc.AddDownloadWithMeta(t.Context(), "user-1", "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=Movie+2026", "/downloads/movies", DownloadTaskMeta{Title: "Movie 2026"}) + if err != nil { + t.Fatal(err) + } + if task.Source != "transmission" || task.DownloadClientID != client.ID || task.ExternalID != "transmission-hash" { + t.Fatalf("task downloader identity = %#v", task) + } + mu.Lock() + defer mu.Unlock() + if added["filename"] != "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=Movie+2026" { + t.Fatalf("transmission filename = %#v", added["filename"]) + } + if added["download-dir"] != "/downloads/movies" { + t.Fatalf("transmission download-dir = %#v", added["download-dir"]) + } +} + +func TestAddDownloadSendsFetchedTorrentBytesToTransmission(t *testing.T) { + torrentData := []byte("d4:infod4:name7:fixtureee") + torrentServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/x-bittorrent") + w.Header().Set("Content-Disposition", `attachment; filename="fixture.torrent"`) + _, _ = w.Write(torrentData) + })) + defer torrentServer.Close() + + var metainfo string + transmission := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + arguments := map[string]interface{}{} + switch req.Method { + case "torrent-get": + arguments["torrents"] = []interface{}{} + case "torrent-add": + metainfo, _ = req.Arguments["metainfo"].(string) + arguments["torrent-added"] = map[string]interface{}{"hashString": "torrent-file-hash"} + } + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{Result: "success", Arguments: arguments}) + })) + defer transmission.Close() + + db := newServiceTestDB(t, &model.Site{}, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) + repos := repository.New(db) + if err := repos.Site.Create(t.Context(), &model.Site{Name: "Fixture", Type: "custom_rss", URL: torrentServer.URL, AuthType: "cookie", Enabled: true}); err != nil { + t.Fatal(err) + } + client := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: transmission.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + site := NewSiteService(zap.NewNop(), repos, "") + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil, site) + svc.SetDownloadManager(manager) + task, err := svc.AddDownloadWithMeta(t.Context(), "user-1", torrentServer.URL+"/fixture.torrent", "/downloads", DownloadTaskMeta{}) + if err != nil { + t.Fatal(err) + } + if metainfo != base64.StdEncoding.EncodeToString(torrentData) { + t.Fatalf("metainfo = %q", metainfo) + } + if task.ExternalID != "torrent-file-hash" || task.Title != "fixture" { + t.Fatalf("task = %#v", task) + } +} + +func TestAddDownloadSendsPublicTorrentURLBytesToAria2(t *testing.T) { + torrentData := []byte("d4:infod4:name7:fixtureee") + torrentServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/x-bittorrent") + _, _ = w.Write(torrentData) + })) + defer torrentServer.Close() + + var addMethod string + aria2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req aria2Request + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode aria2 request: %v", err) + return + } + result := interface{}(map[string]interface{}{"version": "1.37"}) + switch req.Method { + case "aria2.tellActive", "aria2.tellWaiting", "aria2.tellStopped": + result = []interface{}{} + case "aria2.addTorrent", "aria2.addUri": + addMethod = req.Method + result = "aria2-torrent-gid" + } + _ = json.NewEncoder(w).Encode(map[string]interface{}{"jsonrpc": "2.0", "id": req.ID, "result": result}) + })) + defer aria2.Close() + + db := newServiceTestDB(t, &model.Site{}, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) + repos := repository.New(db) + client := &model.DownloadClient{Name: "aria2", Type: "aria2", Host: aria2.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + site := NewSiteService(zap.NewNop(), repos, "") + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil, site) + svc.SetDownloadManager(NewDownloadManager(zap.NewNop(), repos, nil)) + task, err := svc.AddDownloadWithMeta(t.Context(), "user-1", torrentServer.URL+"/public.torrent", "/downloads", DownloadTaskMeta{Title: "Public Torrent"}) + if err != nil { + t.Fatal(err) + } + if addMethod != "aria2.addTorrent" || task.ExternalID != "aria2-torrent-gid" { + t.Fatalf("add method = %q task = %#v", addMethod, task) + } +} + +func TestReloadConfigHotSwapsUpdatedTransmissionClient(t *testing.T) { + newServer := func(addCalls *int32, hash string) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + arguments := map[string]interface{}{} + switch req.Method { + case "torrent-get": + arguments["torrents"] = []interface{}{} + case "torrent-add": + atomic.AddInt32(addCalls, 1) + arguments["torrent-added"] = map[string]interface{}{"hashString": hash} + } + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{Result: "success", Arguments: arguments}) + })) + } + var firstCalls, secondCalls int32 + first := newServer(&firstCalls, "first-hash") + defer first.Close() + second := newServer(&secondCalls, "second-hash") + defer second.Close() + + db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) + repos := repository.New(db) + client := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: first.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.SetDownloadManager(manager) + if _, err := svc.AddDownloadWithMeta(t.Context(), "user-1", "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=First+Movie", "/downloads", DownloadTaskMeta{Title: "First Movie"}); err != nil { + t.Fatal(err) + } + + client.Host = second.URL + if err := repos.DownloadClient.Update(t.Context(), client); err != nil { + t.Fatal(err) + } + task, err := svc.AddDownloadWithMeta(t.Context(), "user-1", "magnet:?xt=urn:btih:bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb&dn=Second+Movie", "/downloads", DownloadTaskMeta{Title: "Second Movie"}) + if err != nil { + t.Fatal(err) + } + if atomic.LoadInt32(&firstCalls) != 1 || atomic.LoadInt32(&secondCalls) != 1 || task.ExternalID != "second-hash" { + t.Fatalf("hot reload calls = %d/%d task = %#v", firstCalls, secondCalls, task) + } +} + +func TestDownloadManagerPersistsOldestEnabledClientAsDefault(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + db := newServiceTestDB(t, &model.DownloadClient{}) + repos := repository.New(db) + first := &model.DownloadClient{ + Base: model.Base{CreatedAt: time.Now().Add(-time.Hour)}, + Name: "First Transmission", + Type: "transmission", + Host: server.URL, + Enabled: true, + IsDefault: false, + } + second := &model.DownloadClient{ + Base: model.Base{CreatedAt: time.Now()}, + Name: "Second Transmission", + Type: "transmission", + Host: server.URL, + Enabled: true, + IsDefault: false, + } + if err := repos.DownloadClient.Create(t.Context(), first); err != nil { + t.Fatal(err) + } + if err := repos.DownloadClient.Create(t.Context(), second); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + if err := manager.LoadAll(t.Context()); err != nil { + t.Fatal(err) + } + selected, _, err := manager.GetDefault(t.Context()) + if err != nil { + t.Fatal(err) + } + if selected.ID != first.ID { + t.Fatalf("default client = %#v", selected) + } + refreshed, err := repos.DownloadClient.FindByID(t.Context(), first.ID) + if err != nil { + t.Fatal(err) + } + if refreshed == nil || !refreshed.IsDefault { + t.Fatalf("persisted default = %#v", refreshed) + } +} + +func TestRelocateRejectsNonQBitClient(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + t.Errorf("unexpected Transmission request during unsupported relocation") + })) + defer server.Close() + + db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) + repos := repository.New(db) + client := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: server.URL, IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.SetDownloadManager(manager) + if err := svc.ReloadConfig(t.Context()); err != nil { + t.Fatal(err) + } + + err := svc.RelocateTorrent(t.Context(), "transmission-hash", "/new/location", client.ID) + if !errors.Is(err, ErrDownloadOperationUnsupported) { + t.Fatalf("err = %v", err) + } +} + +func TestListAggregatesEnabledClientsWithNormalizedProgress(t *testing.T) { + transmission := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + w.Header().Set("X-Transmission-Session-Id", "session-test") + w.WriteHeader(http.StatusConflict) + return + } + var req transmissionRPCRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode transmission request: %v", err) + return + } + arguments := map[string]interface{}{} + if req.Method == "torrent-get" { + arguments["torrents"] = []map[string]interface{}{{ + "hashString": "transmission-hash", + "name": "Transmission Movie", + "totalSize": 1000, + "percentDone": 0.5, + "rateDownload": 100, + "rateUpload": 10, + "status": 4, + "downloadDir": "/downloads/transmission", + "addedDate": 100, + "doneDate": 0, + }} + } + _ = json.NewEncoder(w).Encode(transmissionRPCResponse{Result: "success", Arguments: arguments}) + })) + defer transmission.Close() + + aria2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req aria2Request + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Errorf("decode aria2 request: %v", err) + return + } + result := interface{}(map[string]interface{}{"version": "1.37"}) + switch req.Method { + case "aria2.tellActive": + result = []map[string]interface{}{{ + "gid": "aria2-gid", + "bittorrent": map[string]interface{}{"info": map[string]interface{}{"name": "Aria Movie"}, "infoHash": "aria-info-hash"}, + "totalLength": "2000", + "completedLength": "500", + "downloadSpeed": "200", + "uploadSpeed": "20", + "status": "active", + "dir": "/downloads/aria2", + }} + case "aria2.tellWaiting", "aria2.tellStopped": + result = []interface{}{} + } + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "jsonrpc": "2.0", + "id": req.ID, + "result": result, + }) + })) + defer aria2.Close() + + db := newServiceTestDB(t, &model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}) + repos := repository.New(db) + transmissionClient := &model.DownloadClient{Name: "Transmission", Type: "transmission", Host: transmission.URL, IsDefault: true, Enabled: true} + aria2Client := &model.DownloadClient{Name: "aria2", Type: "aria2", Host: aria2.URL, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), transmissionClient); err != nil { + t.Fatal(err) + } + if err := repos.DownloadClient.Create(t.Context(), aria2Client); err != nil { + t.Fatal(err) + } + manager := NewDownloadManager(zap.NewNop(), repos, nil) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc.SetDownloadManager(manager) + if err := svc.ReloadConfig(t.Context()); err != nil { + t.Fatal(err) + } + + _, live, err := svc.List(t.Context()) + if err != nil { + t.Fatal(err) + } + if len(live) != 2 { + t.Fatalf("live torrents = %#v", live) + } + sort.Slice(live, func(i, j int) bool { return live[i].Source < live[j].Source }) + if live[0].Source != "aria2" || live[0].ClientID != aria2Client.ID || live[0].Progress != 0.25 || live[0].ContentPath != "/downloads/aria2/Aria Movie" { + t.Fatalf("aria2 live torrent = %#v", live[0]) + } + if live[1].Source != "transmission" || live[1].ClientID != transmissionClient.ID || live[1].Progress != 0.5 || live[1].ContentPath != "/downloads/transmission/Transmission Movie" { + t.Fatalf("transmission live torrent = %#v", live[1]) + } +} diff --git a/internal/service/download_polling.go b/internal/service/download_polling.go index 8f3f450..780c13f 100644 --- a/internal/service/download_polling.go +++ b/internal/service/download_polling.go @@ -13,7 +13,7 @@ const completedTorrentOrganizeQueueSize = 64 var completedTorrentOrganizeCooldown = 3 * time.Second -// poll fans out qBittorrent /torrents/info every 5 s as WS events. The +// poll aggregates every enabled downloader every 5 s as WS events. The // payload is opaque to the client; the React store merges by hash. func (d *DownloadService) poll(ctx context.Context) { t := time.NewTicker(5 * time.Second) @@ -30,8 +30,8 @@ func (d *DownloadService) poll(ctx context.Context) { return case <-t.C: } - live, err := d.qb.List(ctx, "") - if err != nil { + live, err := d.listLiveTorrents(ctx, "") + if err != nil && len(live) == 0 { continue } rows, _ := d.repo.Download.List(ctx) @@ -76,7 +76,7 @@ func (d *DownloadService) processTorrentSnapshot(ctx context.Context, torrent QB } func (d *DownloadService) downloadSnapshotTaskNeedsOrganize(ctx context.Context, torrent QBitTorrent, taskByKey map[string]model.DownloadTask) bool { - matchedTask, hasTask := findMatchingTaskByTorrentIdentity(torrent.Name, taskByKey) + matchedTask, hasTask := findMatchingTaskForTorrent(torrent, taskByKey) if !hasTask || d.completedTorrentCatchupRecorded(ctx, torrent) || !d.downloadAutoOrganizeEnabled(ctx) { return false } diff --git a/internal/service/download_runtime.go b/internal/service/download_runtime.go new file mode 100644 index 0000000..e05c64e --- /dev/null +++ b/internal/service/download_runtime.go @@ -0,0 +1,191 @@ +package service + +import ( + "context" + "encoding/base32" + "encoding/hex" + "errors" + "fmt" + "net/url" + "strings" +) + +const legacyQBitDownloadClientID = "legacy-qbittorrent" + +type downloadTarget struct { + clientID string + typ string + adapter DownloadAdapter + legacyQB bool +} + +func (d *DownloadService) listLiveTorrents(ctx context.Context, filter string) ([]QBitTorrent, error) { + if d != nil && d.manager != nil && d.manager.hasClients() { + var live []QBitTorrent + var listErrs []error + for _, target := range d.manager.targets() { + items, err := target.adapter.List(ctx, filter) + if err != nil { + listErrs = append(listErrs, fmt.Errorf("%s (%s): %w", target.client.Name, target.client.Type, err)) + } + for _, item := range items { + torrent := TorrentInfoToQBit(item) + torrent.ClientID = target.client.ID + torrent.Source = target.client.Type + live = append(live, torrent) + } + } + return live, errors.Join(listErrs...) + } + if d != nil && d.qb != nil && d.qb.IsConfigured() { + live, err := d.qb.List(ctx, filter) + for i := range live { + live[i].ClientID = legacyQBitDownloadClientID + live[i].Source = "qbittorrent" + live[i].Progress = float32(normalizedTorrentProgress(float64(live[i].Progress))) + live[i].State = canonicalTorrentState(live[i].State, float64(live[i].Progress)) + } + return live, err + } + return nil, errors.New("no download client available") +} + +func torrentURLInfoHash(raw string) string { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || !strings.EqualFold(parsed.Scheme, "magnet") { + return "" + } + for _, xt := range parsed.Query()["xt"] { + const prefix = "urn:btih:" + if strings.HasPrefix(strings.ToLower(xt), prefix) { + return normalizeTorrentInfoHash(xt[len(prefix):]) + } + } + return "" +} + +func normalizeTorrentInfoHash(value string) string { + value = strings.TrimSpace(value) + if len(value) == 40 { + if _, err := hex.DecodeString(value); err == nil { + return strings.ToLower(value) + } + } + if len(value) == 32 { + decoded, err := base32.StdEncoding.WithPadding(base32.NoPadding).DecodeString(strings.ToUpper(value)) + if err == nil && len(decoded) == 20 { + return hex.EncodeToString(decoded) + } + } + return strings.ToLower(value) +} + +func (d *DownloadService) defaultDownloadTarget(ctx context.Context) (downloadTarget, error) { + if d != nil && d.manager != nil { + client, adapter, err := d.manager.GetDefault(ctx) + if err == nil && client != nil && adapter != nil { + return downloadTarget{clientID: client.ID, typ: client.Type, adapter: adapter}, nil + } + } + if d != nil && d.qb != nil && d.qb.IsConfigured() { + return downloadTarget{clientID: legacyQBitDownloadClientID, typ: "qbittorrent", legacyQB: true}, nil + } + return downloadTarget{}, errors.New("no default downloader configured: 请在下载客户端中配置并启用默认下载器") +} + +func (d *DownloadService) downloadTargetByID(ctx context.Context, clientID string) (downloadTarget, error) { + clientID = strings.TrimSpace(clientID) + if clientID == "" { + return d.defaultDownloadTarget(ctx) + } + if clientID == legacyQBitDownloadClientID { + if d != nil && d.qb != nil && d.qb.IsConfigured() { + return downloadTarget{clientID: legacyQBitDownloadClientID, typ: "qbittorrent", legacyQB: true}, nil + } + return downloadTarget{}, errors.New("legacy qbittorrent is not configured") + } + if d != nil && d.manager != nil { + target, err := d.manager.getTarget(clientID) + if err == nil { + return downloadTarget{clientID: target.client.ID, typ: target.client.Type, adapter: target.adapter}, nil + } + } + return downloadTarget{}, errors.New("download client not found or disabled") +} + +func (d *DownloadService) resolveOperationClientID(ctx context.Context, externalID, requestedClientID string) (string, string, error) { + requestedClientID = strings.TrimSpace(requestedClientID) + live, listErr := d.listLiveTorrents(ctx, "") + if clientID, name, err, resolved := resolveLiveOperationClient(live, externalID, requestedClientID); resolved { + return clientID, name, err + } + if requestedClientID != "" { + return requestedClientID, "", nil + } + if clientID, err := d.persistedOperationClientID(ctx, externalID); clientID != "" || err != nil { + return clientID, "", err + } + if clientID, err := d.singleAvailableOperationClientID(); clientID != "" || err != nil { + return clientID, "", err + } + if d != nil && d.qb != nil && d.qb.IsConfigured() { + return legacyQBitDownloadClientID, "", nil + } + if listErr != nil { + return "", "", listErr + } + return "", "", errors.New("download client not found") +} + +func resolveLiveOperationClient(live []QBitTorrent, externalID, requestedClientID string) (string, string, error, bool) { + var matched []QBitTorrent + for _, torrent := range live { + if requestedClientID != "" && torrent.ClientID != requestedClientID { + continue + } + if strings.EqualFold(strings.TrimSpace(torrent.Hash), strings.TrimSpace(externalID)) { + matched = append(matched, torrent) + } + } + if len(matched) == 1 { + return matched[0].ClientID, matched[0].Name, nil, true + } + if len(matched) > 1 && requestedClientID == "" { + return "", "", errors.New("multiple download clients contain this task; client_id is required"), true + } + return "", "", nil, false +} + +func (d *DownloadService) persistedOperationClientID(ctx context.Context, externalID string) (string, error) { + if d != nil && d.repo != nil && d.repo.Download != nil { + rows, err := d.repo.Download.List(ctx) + if err != nil { + return "", err + } + var clientID string + for _, row := range rows { + if !strings.EqualFold(strings.TrimSpace(row.ExternalID), strings.TrimSpace(externalID)) || strings.TrimSpace(row.DownloadClientID) == "" { + continue + } + if clientID != "" && clientID != row.DownloadClientID { + return "", errors.New("multiple download clients contain this task; client_id is required") + } + clientID = row.DownloadClientID + } + return clientID, nil + } + return "", nil +} + +func (d *DownloadService) singleAvailableOperationClientID() (string, error) { + if d != nil && d.manager != nil { + targets := d.manager.targets() + if len(targets) == 1 { + return targets[0].client.ID, nil + } + if len(targets) > 1 { + return "", errors.New("client_id is required when multiple download clients are enabled") + } + } + return "", nil +} diff --git a/internal/service/download_torrent_normalize.go b/internal/service/download_torrent_normalize.go new file mode 100644 index 0000000..0b3b98d --- /dev/null +++ b/internal/service/download_torrent_normalize.go @@ -0,0 +1,80 @@ +package service + +import "strings" + +func normalizedTorrentProgress(progress float64) float64 { + if progress < 0 { + return 0 + } + if progress > 1 { + return 1 + } + return progress +} + +func canonicalTorrentState(state string, progress float64) string { + state = strings.ToLower(strings.TrimSpace(state)) + complete := normalizedTorrentProgress(progress) >= 1 + switch state { + case "completed", "complete", "seeding", "uploading", "stalledup", "pausedup", "queuedup", "forcedup": + if complete || state == "completed" || state == "complete" || state == "seeding" { + return "completed" + } + return "downloading" + case "downloading", "forceddl", "metadl", "stalleddl", "active": + if complete { + return "completed" + } + return "downloading" + case "queued", "queueddl", "download_pending", "seed_pending", "waiting": + if complete { + return "completed" + } + return "queued" + case "paused", "pauseddl", "stoppeddl", "stopped": + if complete { + return "completed" + } + return "paused" + case "checking", "checkingdl", "checkingup", "checkingresumedata", "check_pending", "moving": + return "checking" + case "error", "missingfiles": + return "error" + case "removed": + return "removed" + case "": + if complete { + return "completed" + } + return "" + default: + if complete { + return "completed" + } + return state + } +} + +func downloaderPayloadPath(dir, name string) string { + dir = strings.TrimSpace(dir) + name = strings.TrimSpace(name) + if name == "" { + return dir + } + if dir == "" { + return name + } + separator := "/" + if strings.Contains(dir, `\`) && !strings.Contains(dir, "/") { + separator = `\` + } + return strings.TrimRight(dir, `/\`) + separator + strings.TrimLeft(name, `/\`) +} + +func downloaderPathBase(value string) string { + value = strings.TrimRight(strings.TrimSpace(value), `/\`) + if i := strings.LastIndexAny(value, `/\`); i >= 0 { + return value[i+1:] + } + return value +} diff --git a/internal/service/download_views.go b/internal/service/download_views.go index 0966eb8..b0f7384 100644 --- a/internal/service/download_views.go +++ b/internal/service/download_views.go @@ -9,30 +9,34 @@ import ( ) type DownloadTaskView struct { - ID string `json:"id"` - Source string `json:"source"` - Title string `json:"title"` - PosterURL string `json:"poster_url,omitempty"` - BackdropURL string `json:"backdrop_url,omitempty"` - Overview string `json:"overview,omitempty"` - SavePath string `json:"save_path"` - MediaType string `json:"media_type,omitempty"` - MediaCategory string `json:"media_category,omitempty"` - Status string `json:"status"` - Progress float32 `json:"progress"` - State string `json:"state,omitempty"` - DLSpeed int64 `json:"dlspeed,omitempty"` - UpSpeed int64 `json:"upspeed,omitempty"` - Size int64 `json:"size,omitempty"` - Downloaded int64 `json:"downloaded,omitempty"` - NumSeeds int `json:"num_seeds,omitempty"` - NumLeechs int `json:"num_leechs,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + ID string `json:"id"` + Source string `json:"source"` + DownloadClientID string `json:"download_client_id,omitempty"` + ExternalID string `json:"external_id,omitempty"` + Title string `json:"title"` + PosterURL string `json:"poster_url,omitempty"` + BackdropURL string `json:"backdrop_url,omitempty"` + Overview string `json:"overview,omitempty"` + SavePath string `json:"save_path"` + MediaType string `json:"media_type,omitempty"` + MediaCategory string `json:"media_category,omitempty"` + Status string `json:"status"` + Progress float32 `json:"progress"` + State string `json:"state,omitempty"` + DLSpeed int64 `json:"dlspeed,omitempty"` + UpSpeed int64 `json:"upspeed,omitempty"` + Size int64 `json:"size,omitempty"` + Downloaded int64 `json:"downloaded,omitempty"` + NumSeeds int `json:"num_seeds,omitempty"` + NumLeechs int `json:"num_leechs,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` } type DownloadTorrentView struct { Hash string `json:"hash"` + ClientID string `json:"client_id"` + Source string `json:"source"` Name string `json:"name"` Title string `json:"title"` PosterURL string `json:"poster_url,omitempty"` @@ -56,21 +60,18 @@ func DownloadViews(rows []model.DownloadTask, live []QBitTorrent) ([]DownloadTas for _, torrent := range live { key := normalizeTorrentName(torrent.Name) if key != "" { - liveByKey[key] = torrent - } - } - taskByKey := map[string]model.DownloadTask{} - for _, row := range rows { - key := normalizeTorrentName(row.Title) - if key != "" { - taskByKey[key] = row + setLiveTorrentIndex(liveByKey, key, torrent) + setLiveTorrentIndex(liveByKey, downloadTaskClientTitleKey(torrent.ClientID, key), torrent) } + setLiveTorrentIndex(liveByKey, downloadTaskExternalKey(torrent.ClientID, torrent.Hash), torrent) + setLiveTorrentIndex(liveByKey, downloadTaskAnyExternalKey(torrent.Hash), torrent) } + taskByKey := tasksByTorrentIdentity(rows) taskViews := make([]DownloadTaskView, 0, len(rows)) for _, row := range rows { view := downloadTaskView(row, QBitTorrent{}) - if torrent, ok := findMatchingTorrent(row.Title, liveByKey); ok { + if torrent, ok := findMatchingTorrentForTask(row, liveByKey); ok { view = downloadTaskView(row, torrent) } taskViews = append(taskViews, view) @@ -79,7 +80,7 @@ func DownloadViews(rows []model.DownloadTask, live []QBitTorrent) ([]DownloadTas torrentViews := make([]DownloadTorrentView, 0, len(live)) for _, torrent := range live { var row model.DownloadTask - if matched, ok := findMatchingTask(torrent.Name, taskByKey); ok { + if matched, ok := findMatchingTaskForTorrent(torrent, taskByKey); ok { row = matched } torrentViews = append(torrentViews, downloadTorrentView(torrent, row)) @@ -96,26 +97,28 @@ func downloadTaskView(row model.DownloadTask, torrent QBitTorrent) DownloadTaskV } size := torrent.Size return DownloadTaskView{ - ID: row.ID, - Source: row.Source, - Title: firstNonEmpty(row.Title, "下载任务"), - PosterURL: row.PosterURL, - BackdropURL: row.BackdropURL, - Overview: row.Overview, - SavePath: row.SavePath, - MediaType: row.MediaType, - MediaCategory: row.MediaCategory, - Status: row.Status, - Progress: progress, - State: state, - DLSpeed: torrent.DLSpeed, - UpSpeed: torrent.UpSpeed, - Size: size, - Downloaded: downloadedBytes(size, progress), - NumSeeds: torrent.NumSeeds, - NumLeechs: torrent.NumLeech, - CreatedAt: row.CreatedAt, - UpdatedAt: row.UpdatedAt, + ID: row.ID, + Source: row.Source, + DownloadClientID: row.DownloadClientID, + ExternalID: row.ExternalID, + Title: firstNonEmpty(row.Title, "下载任务"), + PosterURL: row.PosterURL, + BackdropURL: row.BackdropURL, + Overview: row.Overview, + SavePath: row.SavePath, + MediaType: row.MediaType, + MediaCategory: row.MediaCategory, + Status: row.Status, + Progress: progress, + State: state, + DLSpeed: torrent.DLSpeed, + UpSpeed: torrent.UpSpeed, + Size: size, + Downloaded: downloadedBytes(size, progress), + NumSeeds: torrent.NumSeeds, + NumLeechs: torrent.NumLeech, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, } } @@ -126,6 +129,8 @@ func downloadTorrentView(torrent QBitTorrent, row model.DownloadTask) DownloadTo } return DownloadTorrentView{ Hash: torrent.Hash, + ClientID: torrent.ClientID, + Source: torrent.Source, Name: torrent.Name, Title: firstNonEmpty(title, "下载任务"), PosterURL: row.PosterURL, @@ -145,6 +150,34 @@ func downloadTorrentView(torrent QBitTorrent, row model.DownloadTask) DownloadTo } } +func setLiveTorrentIndex(index map[string]QBitTorrent, key string, torrent QBitTorrent) { + if key == "" { + return + } + if _, exists := index[key]; !exists { + index[key] = torrent + } +} + +func findMatchingTorrentForTask(row model.DownloadTask, liveByKey map[string]QBitTorrent) (QBitTorrent, bool) { + if torrent, ok := liveByKey[downloadTaskExternalKey(row.DownloadClientID, row.ExternalID)]; ok { + return torrent, true + } + if torrent, ok := liveByKey[downloadTaskAnyExternalKey(row.ExternalID)]; ok { + if strings.TrimSpace(row.DownloadClientID) == "" || row.DownloadClientID == torrent.ClientID { + return torrent, true + } + } + key := normalizeTorrentName(row.Title) + if torrent, ok := liveByKey[downloadTaskClientTitleKey(row.DownloadClientID, key)]; ok { + return torrent, true + } + if strings.TrimSpace(row.DownloadClientID) != "" { + return QBitTorrent{}, false + } + return findMatchingTorrent(row.Title, liveByKey) +} + func findMatchingTorrent(title string, liveByKey map[string]QBitTorrent) (QBitTorrent, bool) { key := normalizeTorrentName(title) if key == "" { @@ -154,6 +187,9 @@ func findMatchingTorrent(title string, liveByKey map[string]QBitTorrent) (QBitTo return torrent, true } for currentKey, torrent := range liveByKey { + if strings.HasPrefix(currentKey, "\x00") { + continue + } if strings.Contains(currentKey, key) || strings.Contains(key, currentKey) { return torrent, true } @@ -161,22 +197,6 @@ func findMatchingTorrent(title string, liveByKey map[string]QBitTorrent) (QBitTo return QBitTorrent{}, false } -func findMatchingTask(title string, taskByKey map[string]model.DownloadTask) (model.DownloadTask, bool) { - key := normalizeTorrentName(title) - if key == "" { - return model.DownloadTask{}, false - } - if row, ok := taskByKey[key]; ok { - return row, true - } - for currentKey, row := range taskByKey { - if strings.Contains(key, currentKey) || strings.Contains(currentKey, key) { - return row, true - } - } - return model.DownloadTask{}, false -} - func downloadedBytes(size int64, progress float32) int64 { if size <= 0 || progress <= 0 { return 0 diff --git a/internal/service/downloads.go b/internal/service/downloads.go index 05f0d98..a0df2d5 100644 --- a/internal/service/downloads.go +++ b/internal/service/downloads.go @@ -1,7 +1,7 @@ // Package service — download manager. // // DownloadService persists user-initiated downloads, dispatches them to -// the configured client (currently qBittorrent) and pushes live progress +// the configured qBittorrent, Transmission, or aria2 client and pushes live progress // to the WS hub so the React UI can render a live table. // // Settings consumed (system Setting table): @@ -18,6 +18,7 @@ package service import ( "context" "errors" + "fmt" "strings" "sync" "time" @@ -34,6 +35,7 @@ type DownloadService struct { repo *repository.Container hub *Hub qb *QBitClient + manager *DownloadManager organizer *OrganizerService organizePipeline *OrganizePipelineService scanner *ScannerService @@ -45,7 +47,7 @@ type DownloadService struct { stopCh chan struct{} pollOnce sync.Once organizeOnce sync.Once - prevStates map[string]bool // hash -> wasCompleted + prevStates map[string]bool // client/task identity -> wasCompleted pollInitialized bool liveTorrents []QBitTorrent liveTorrentsAt time.Time @@ -70,8 +72,12 @@ func (d *DownloadService) SetNotifyChannels(notify *NotifyChannelService) { d.notify = notify } +func (d *DownloadService) SetDownloadManager(manager *DownloadManager) { + d.manager = manager +} + // ErrDownloadAlreadyExists tells callers that the requested resource is already -// tracked locally or present in qBittorrent. Subscriptions treat this as a +// tracked locally or present in a downloader. Subscriptions treat this as a // successful dedup hit, not as a retryable enqueue failure. var ErrDownloadAlreadyExists = errors.New("download already exists") @@ -80,6 +86,8 @@ var ErrDownloadAlreadyExists = errors.New("download already exists") // downloader again. var ErrMediaAlreadyInLibrary = errors.New("media already exists in library") +var ErrDownloadOperationUnsupported = errors.New("download client operation unsupported") + func IsDownloadDedupError(err error) bool { return errors.Is(err, ErrDownloadAlreadyExists) || errors.Is(err, ErrMediaAlreadyInLibrary) } @@ -108,7 +116,9 @@ func NewDownloadService(log *zap.Logger, repo *repository.Container, hub *Hub, o // Start kicks off the background poller (idempotent). func (d *DownloadService) Start(ctx context.Context) { d.pollOnce.Do(func() { - _ = d.ReloadConfig(ctx) + if err := d.ReloadConfig(ctx); err != nil && d.log != nil { + d.log.Warn("initial download client reload failed", zap.Error(err)) + } d.startAutoOrganizeWorker(ctx) go d.poll(ctx) }) @@ -124,8 +134,8 @@ func (d *DownloadService) TorrentExistsByName(ctx context.Context, name string) if query == "" { return false } - live, err := d.qb.List(ctx, "") - if err != nil { + live, err := d.listLiveTorrents(ctx, "") + if err != nil && len(live) == 0 { return false } for _, torrent := range live { @@ -143,43 +153,52 @@ func (d *DownloadService) TorrentExistsByName(ctx context.Context, name string) return false } -// List returns every persisted download task augmented with live data -// from qBittorrent when available. +// List returns every persisted download task augmented with live downloader data. func (d *DownloadService) List(ctx context.Context) ([]model.DownloadTask, []QBitTorrent, error) { rows, err := d.repo.Download.List(ctx) if err != nil { return nil, nil, err } - live, err := d.qb.List(ctx, "") + live, err := d.listLiveTorrents(ctx, "") if err != nil { - // Network failure shouldn't break the page — return rows with no - // live data and let the UI render the persisted snapshot. - d.log.Debug("qbittorrent list failed", zap.Error(err)) - return rows, nil, nil + // A failed client must not hide healthy clients or persisted rows. + d.log.Debug("download client list failed", zap.Error(err)) + if len(live) == 0 { + return rows, nil, nil + } } return rows, live, nil } -// Delete removes a torrent (and optionally its files) from qBittorrent. -func (d *DownloadService) Delete(ctx context.Context, hash string, withFiles bool) error { +// Delete removes a task from its native downloader. clientID is optional for +// legacy callers, but disambiguates equal native IDs across multiple clients. +func (d *DownloadService) Delete(ctx context.Context, hash string, withFiles bool, clientID ...string) error { hash = strings.TrimSpace(hash) if hash == "" { return errors.New("hash is required") } - var torrentName string - if live, err := d.qb.List(ctx, ""); err == nil { - for _, torrent := range live { - if strings.EqualFold(torrent.Hash, hash) || len(live) == 1 { - torrentName = torrent.Name - break - } - } + requestedClientID := "" + if len(clientID) > 0 { + requestedClientID = clientID[0] } - if err := d.qb.Delete(ctx, hash, withFiles); err != nil { + resolvedClientID, torrentName, err := d.resolveOperationClientID(ctx, hash, requestedClientID) + if err != nil { return err } - d.markDownloadTaskDeleted(ctx, hash, torrentName) - stateKey := strings.ToLower(hash) + target, err := d.downloadTargetByID(ctx, resolvedClientID) + if err != nil { + return err + } + if target.legacyQB { + err = d.qb.Delete(ctx, hash, withFiles) + } else { + err = target.adapter.Remove(ctx, hash, withFiles) + } + if err != nil { + return err + } + d.markDownloadTaskDeleted(ctx, hash, torrentName, resolvedClientID) + stateKey := completedTorrentQueueKey(QBitTorrent{ClientID: resolvedClientID, Hash: hash}) d.mu.Lock() delete(d.prevStates, stateKey) delete(d.organizeQueued, stateKey) @@ -187,7 +206,7 @@ func (d *DownloadService) Delete(ctx context.Context, hash string, withFiles boo return nil } -func (d *DownloadService) markDownloadTaskDeleted(ctx context.Context, hash, torrentName string) { +func (d *DownloadService) markDownloadTaskDeleted(ctx context.Context, hash, torrentName string, clientID ...string) { if d == nil || d.repo == nil || d.repo.DB == nil { return } @@ -195,7 +214,7 @@ func (d *DownloadService) markDownloadTaskDeleted(ctx context.Context, hash, tor if err != nil { return } - if matched, ok := findDownloadTaskByHash(rows, hash); ok { + if matched, ok := findDownloadTaskByHash(rows, hash, clientID...); ok { _ = d.repo.DB.WithContext(ctx).Model(&model.DownloadTask{}). Where("id = ?", matched.ID). Updates(map[string]any{ @@ -208,7 +227,7 @@ func (d *DownloadService) markDownloadTaskDeleted(ctx context.Context, hash, tor return } taskByKey := tasksByTorrentIdentity(rows) - matched, ok := findMatchingTaskByTorrentIdentity(torrentName, taskByKey) + matched, ok := findMatchingTaskForTorrent(QBitTorrent{Name: torrentName, Hash: hash, ClientID: firstString(clientID)}, taskByKey) if !ok { return } @@ -220,12 +239,24 @@ func (d *DownloadService) markDownloadTaskDeleted(ctx context.Context, hash, tor }).Error } -func findDownloadTaskByHash(rows []model.DownloadTask, hash string) (model.DownloadTask, bool) { +func findDownloadTaskByHash(rows []model.DownloadTask, hash string, clientID ...string) (model.DownloadTask, bool) { hash = strings.ToLower(strings.TrimSpace(hash)) if hash == "" { return model.DownloadTask{}, false } + wantClientID := strings.TrimSpace(firstString(clientID)) for _, row := range rows { + if wantClientID != "" && strings.TrimSpace(row.DownloadClientID) != "" && row.DownloadClientID != wantClientID { + continue + } + if strings.EqualFold(strings.TrimSpace(row.ExternalID), hash) { + return row, true + } + } + for _, row := range rows { + if wantClientID != "" && strings.TrimSpace(row.DownloadClientID) != "" && row.DownloadClientID != wantClientID { + continue + } if strings.Contains(strings.ToLower(row.URL), hash) { return row, true } @@ -233,15 +264,45 @@ func findDownloadTaskByHash(rows []model.DownloadTask, hash string) (model.Downl return model.DownloadTask{}, false } +func firstString(values []string) string { + if len(values) == 0 { + return "" + } + return values[0] +} + // RelocateTorrent moves a torrent's data to a new save directory while keeping // it seeding (qBittorrent performs the physical move and resumes seeding). // 用于「移动 PT 种子文件且转移后继续做种上传」的整盘迁移场景。 -func (d *DownloadService) RelocateTorrent(ctx context.Context, hash, location string) error { +func (d *DownloadService) RelocateTorrent(ctx context.Context, hash, location string, clientID ...string) error { if strings.TrimSpace(hash) == "" { return errors.New("hash is required") } if strings.TrimSpace(location) == "" { return errors.New("location is required") } - return d.qb.SetLocation(ctx, hash, strings.TrimSpace(location)) + requestedClientID := firstString(clientID) + resolvedClientID := strings.TrimSpace(requestedClientID) + if resolvedClientID == "" { + var err error + resolvedClientID, _, err = d.resolveOperationClientID(ctx, hash, "") + if err != nil { + return err + } + } + target, err := d.downloadTargetByID(ctx, resolvedClientID) + if err != nil { + return err + } + if target.typ != "qbittorrent" { + return fmt.Errorf("%w: %s does not support torrent relocation; only qBittorrent is supported", ErrDownloadOperationUnsupported, target.typ) + } + if target.legacyQB { + return d.qb.SetLocation(ctx, hash, strings.TrimSpace(location)) + } + relocator, ok := target.adapter.(TorrentRelocateAdapter) + if !ok { + return fmt.Errorf("%w: configured qBittorrent adapter cannot relocate torrents", ErrDownloadOperationUnsupported) + } + return relocator.Relocate(ctx, hash, strings.TrimSpace(location)) } diff --git a/internal/service/downloads_config_runtime.go b/internal/service/downloads_config_runtime.go index 9506045..a228df7 100644 --- a/internal/service/downloads_config_runtime.go +++ b/internal/service/downloads_config_runtime.go @@ -13,39 +13,59 @@ import ( const settingDownloadClientsManaged = "download_clients.managed" -// ReloadConfig rebuilds the qBittorrent client from the configured -// download clients (preferred) or the legacy Setting table (fallback). +// ReloadConfig reloads managed adapters and preserves the legacy qBittorrent +// settings fallback for deployments that never used download_clients. // // 配置来源优先级: // -// 1. download_clients 表中 type=qbittorrent 且 is_default=true 且 enabled=true -// 的行(侧边栏「下载器」页面写入的数据)。 -// 2. system Setting 表中的 qbittorrent.url / username / password +// 1. download_clients 表中已启用的显式默认客户端;没有默认时持久化最早 +// 创建的已启用客户端。 +// 2. 从未使用多客户端下载器配置的旧部署,读取 system Setting 表中的 +// qbittorrent.url / username / password // (旧版「系统设置」表单写入的数据;保留作向后兼容)。 // // 这避免了两套配置各跑各的:之前操作员明明已经在「下载器」页面填好 // 默认 qb,但实际下载链路读的还是 Setting 表,导致一直连不上。 func (d *DownloadService) ReloadConfig(ctx context.Context) error { + if d.manager != nil { + if err := d.manager.LoadAll(ctx); err != nil { + return err + } + if d.manager.hasClients() { + // Managed clients use their native adapters. Keep the legacy qB + // client blank so a Transmission/aria2 default cannot be silently + // overridden by an unrelated qB row. + d.qb.Configure(QBitConfig{}) + return nil + } + } cfg := QBitConfig{} hasConfiguredClients := false managedByDownloadClients := false - // Path 1: download_clients 表 + // Path 1: download_clients 表。此分支主要服务未注入 DownloadManager 的 + // 单元/兼容调用;生产容器已在上方通过原生适配器返回。 if d.repo.DownloadClient != nil { hasConfiguredClients, _ = d.repo.DownloadClient.HasAnyIncludingDeleted(ctx) - if c, err := d.repo.DownloadClient.FindDefault(ctx); err == nil && c != nil && c.Type == "qbittorrent" { - cfg.BaseURL = strings.TrimRight(c.Host, "/") - cfg.Username = c.Username - cfg.Password = c.Password - } else if c, err := d.preferredEnabledQBitClient(ctx); err == nil && c != nil { - cfg.BaseURL = strings.TrimRight(c.Host, "/") - cfg.Username = c.Username - cfg.Password = c.Password - _ = d.repo.DownloadClient.SetDefault(ctx, c.ID) + selected, _ := d.repo.DownloadClient.FindDefault(ctx) + if selected == nil { + selected, _ = d.preferredEnabledClient(ctx) + if selected != nil { + _ = d.repo.DownloadClient.SetDefault(ctx, selected.ID) + selected.IsDefault = true + } + } + if selected != nil { if d.log != nil { - d.log.Warn("default downloader missing; selected first enabled qbittorrent client", - zap.String("client_id", c.ID), - zap.String("client", c.Name)) + d.log.Debug("selected managed default downloader", + zap.String("client_id", selected.ID), + zap.String("client", selected.Name), + zap.String("type", selected.Type)) + } + if selected.Type == "qbittorrent" { + cfg.BaseURL = strings.TrimRight(selected.Host, "/") + cfg.Username = selected.Username + cfg.Password = selected.Password } } } @@ -72,7 +92,7 @@ func (d *DownloadService) ReloadConfig(ctx context.Context) error { return nil } -func (d *DownloadService) preferredEnabledQBitClient(ctx context.Context) (*model.DownloadClient, error) { +func (d *DownloadService) preferredEnabledClient(ctx context.Context) (*model.DownloadClient, error) { if d == nil || d.repo == nil || d.repo.DownloadClient == nil { return nil, nil } @@ -80,33 +100,27 @@ func (d *DownloadService) preferredEnabledQBitClient(ctx context.Context) (*mode if err != nil { return nil, err } - var selected *model.DownloadClient - for i := range rows { - if rows[i].Type != "qbittorrent" { - continue - } - row := rows[i] - selected = &row - break + if len(rows) == 0 { + return nil, nil } - return selected, nil + selected := rows[0] + return &selected, nil } func (d *DownloadService) defaultDownloaderNotConfiguredError(ctx context.Context) error { const prefix = "no default downloader configured" if d == nil || d.repo == nil || d.repo.DownloadClient == nil { - return errors.New(prefix + ": 请在下载客户端中配置并启用 qBittorrent") + return errors.New(prefix + ": 请在下载客户端中配置并启用下载器") } rows, err := d.repo.DownloadClient.ListEnabled(ctx) if err != nil { return fmt.Errorf("%s: 读取下载客户端配置失败: %w", prefix, err) } if len(rows) == 0 { - return errors.New(prefix + ": 请在下载客户端中启用 qBittorrent 并设为默认;当前没有已启用的下载器") + return errors.New(prefix + ": 请在下载客户端中启用下载器;当前没有已启用的下载器") } var enabled []string - var hasQBit bool for _, row := range rows { label := strings.TrimSpace(row.Name) if label == "" { @@ -115,12 +129,6 @@ func (d *DownloadService) defaultDownloaderNotConfiguredError(ctx context.Contex label += "(" + row.Type + ")" } enabled = append(enabled, label) - if strings.EqualFold(strings.TrimSpace(row.Type), "qbittorrent") { - hasQBit = true - } } - if !hasQBit { - return fmt.Errorf("%s: 订阅投递目前需要 qBittorrent;当前启用的下载器为 %s", prefix, strings.Join(enabled, ", ")) - } - return errors.New(prefix + ": 请在下载客户端中选择一个启用的 qBittorrent 作为默认下载器") + return fmt.Errorf("%s: 请检查已启用下载器的连接和默认设置;当前启用的下载器为 %s", prefix, strings.Join(enabled, ", ")) } diff --git a/internal/service/downloads_progress_test.go b/internal/service/downloads_progress_test.go index aa67f08..d204318 100644 --- a/internal/service/downloads_progress_test.go +++ b/internal/service/downloads_progress_test.go @@ -75,6 +75,92 @@ func TestSyncDownloadTaskProgressMatchesSeasonFolderTorrentName(t *testing.T) { } } +func TestSyncDownloadTaskProgressBackfillsDownloaderIdentity(t *testing.T) { + repos := newOrganizerTestRepo(t) + if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil { + t.Fatal(err) + } + task := &model.DownloadTask{ + UserID: "u1", + Source: "qbittorrent", + URL: "https://pt.example/download?id=legacy", + Title: "Legacy Identity Movie 2026", + Status: "queued", + Progress: 0, + } + if err := repos.Download.Create(t.Context(), task); err != nil { + t.Fatal(err) + } + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + torrent := QBitTorrent{ + Hash: "transmission-hash", + ClientID: "transmission-client", + Source: "transmission", + Name: "Legacy Identity Movie 2026", + State: "downloading", + Progress: 0.4, + } + svc.syncDownloadTaskProgress(t.Context(), torrent, tasksByTorrentIdentity([]model.DownloadTask{*task})) + + var updated model.DownloadTask + if err := repos.DB.Where("id = ?", task.ID).First(&updated).Error; err != nil { + t.Fatal(err) + } + if updated.DownloadClientID != torrent.ClientID || updated.ExternalID != torrent.Hash || updated.Source != torrent.Source { + t.Fatalf("downloader identity = %#v", updated) + } +} + +func TestProcessDownloadSnapshotCompletesTransmissionTask(t *testing.T) { + repos := newOrganizerTestRepo(t) + if err := repos.DB.AutoMigrate(&model.DownloadTask{}); err != nil { + t.Fatal(err) + } + if err := repos.Setting.Set(t.Context(), "organizer.auto_after_download", "true"); err != nil { + t.Fatal(err) + } + task := &model.DownloadTask{ + UserID: "u1", + Source: "transmission", + DownloadClientID: "transmission-client", + ExternalID: "shared-native-id", + URL: "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + Title: "Transmission Complete Movie", + Status: "downloading", + Progress: 0.8, + } + if err := repos.Download.Create(t.Context(), task); err != nil { + t.Fatal(err) + } + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + torrent := QBitTorrent{ + Hash: task.ExternalID, + ClientID: task.DownloadClientID, + Source: task.Source, + Name: task.Title, + State: "completed", + Progress: 1, + ContentPath: "/downloads/Transmission Complete Movie", + CompletionOn: time.Now().Unix(), + } + svc.processDownloadSnapshot(t.Context(), []QBitTorrent{torrent}, tasksByTorrentIdentity([]model.DownloadTask{*task})) + if got := len(svc.organizeQueue); got != 1 { + t.Fatalf("queued completed organize jobs = %d", got) + } + var updated model.DownloadTask + if err := repos.DB.Where("id = ?", task.ID).First(&updated).Error; err != nil { + t.Fatal(err) + } + if updated.Status != "completed" || updated.Progress != 1 { + t.Fatalf("task completion = %q/%v", updated.Status, updated.Progress) + } + other := torrent + other.ClientID = "another-client" + if completedTorrentQueueKey(torrent) == completedTorrentQueueKey(other) { + t.Fatal("completion queue keys collided across download clients") + } +} + func TestProcessDownloadSnapshotQueuesCompletedPendingTaskOnFirstSnapshot(t *testing.T) { db := newServiceTestDB(t, &model.DownloadTask{}, &model.Setting{}) repos := repository.New(db) diff --git a/internal/service/qbittorrent.go b/internal/service/qbittorrent.go index 226fbfb..0a1d89e 100644 --- a/internal/service/qbittorrent.go +++ b/internal/service/qbittorrent.go @@ -40,6 +40,8 @@ type QBitConfig struct { // QBitTorrent is the subset of /torrents/info we surface to the API. type QBitTorrent struct { Hash string `json:"hash"` + ClientID string `json:"client_id,omitempty"` + Source string `json:"source,omitempty"` Name string `json:"name"` State string `json:"state"` Progress float32 `json:"progress"` @@ -196,6 +198,49 @@ func (q *QBitClient) Delete(ctx context.Context, hash string, deleteFiles bool) return nil } +func (q *QBitClient) Pause(ctx context.Context, hash string) error { + return q.torrentAction(ctx, hash, "pause", "stop") +} + +func (q *QBitClient) Resume(ctx context.Context, hash string) error { + return q.torrentAction(ctx, hash, "resume", "start") +} + +func (q *QBitClient) torrentAction(ctx context.Context, hash string, actions ...string) error { + q.mu.Lock() + defer q.mu.Unlock() + if err := q.ensureAuth(ctx); err != nil { + return err + } + baseURL := strings.TrimRight(q.cfg.BaseURL, "/") + form := url.Values{"hashes": []string{hash}} + var lastErr error + for _, action := range actions { + req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, + baseURL+"/api/v2/torrents/"+action, strings.NewReader(form.Encode())) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.Header.Set("Referer", baseURL) + req.Header.Set("Origin", baseURL) + resp, err := q.client.Do(req) + if err != nil { + return err + } + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode < 400 { + return nil + } + lastErr = fmt.Errorf("qbittorrent %s: %d: %s", action, resp.StatusCode, strings.TrimSpace(string(body))) + if resp.StatusCode != http.StatusNotFound && resp.StatusCode != http.StatusMethodNotAllowed { + break + } + } + return lastErr +} + // SetLocation moves a torrent's data to a new save directory via // POST /api/v2/torrents/setLocation. qBittorrent performs the physical move // itself and keeps seeding from the new location — this is the seeding-safe diff --git a/internal/service/qbittorrent_adp.go b/internal/service/qbittorrent_adp.go index 1d783b1..bd541cd 100644 --- a/internal/service/qbittorrent_adp.go +++ b/internal/service/qbittorrent_adp.go @@ -7,6 +7,7 @@ package service import ( "bytes" "context" + "errors" "fmt" "io" "mime/multipart" @@ -61,6 +62,12 @@ func (a *QBitAdapter) Ping(ctx context.Context) error { // AddTorrent 通过 URL 添加种子。 func (a *QBitAdapter) AddTorrent(ctx context.Context, torrentURL, savePath string) (string, error) { + return a.AddTorrentWithCategory(ctx, torrentURL, savePath, "") +} + +// AddTorrentWithCategory submits a URL or magnet while preserving the +// qBittorrent category selected by MediaStationGo's classification rules. +func (a *QBitAdapter) AddTorrentWithCategory(ctx context.Context, torrentURL, savePath, category string) (string, error) { a.mu.Lock() defer a.mu.Unlock() if err := a.ensureAuthLocked(ctx); err != nil { @@ -70,10 +77,49 @@ func (a *QBitAdapter) AddTorrent(ctx context.Context, torrentURL, savePath strin body := &bytes.Buffer{} w := multipart.NewWriter(body) _ = w.WriteField("urls", torrentURL) + return a.addTorrentMultipartLocked(ctx, body, w, savePath, category, torrentURLInfoHash(torrentURL)) +} + +// AddMagnet 通过磁力链接添加种子。 +func (a *QBitAdapter) AddMagnet(ctx context.Context, magnet, savePath string) (string, error) { + return a.AddTorrent(ctx, magnet, savePath) +} + +// AddTorrentFile uploads .torrent bytes as multipart/form-data. +func (a *QBitAdapter) AddTorrentFile(ctx context.Context, data []byte, name, savePath string) (string, error) { + return a.AddTorrentFileWithCategory(ctx, data, name, savePath, "") +} + +// AddTorrentFileWithCategory uploads .torrent bytes and preserves the +// qBittorrent category selected by the caller. +func (a *QBitAdapter) AddTorrentFileWithCategory(ctx context.Context, data []byte, name, savePath, category string) (string, error) { + a.mu.Lock() + defer a.mu.Unlock() + if err := a.ensureAuthLocked(ctx); err != nil { + return "", err + } + body := &bytes.Buffer{} + w := multipart.NewWriter(body) + part, err := w.CreateFormFile("torrents", name) + if err != nil { + return "", err + } + if _, err := part.Write(data); err != nil { + return "", err + } + return a.addTorrentMultipartLocked(ctx, body, w, savePath, category, torrentInfoHash(data)) +} + +func (a *QBitAdapter) addTorrentMultipartLocked(ctx context.Context, body *bytes.Buffer, w *multipart.Writer, savePath, category, externalID string) (string, error) { if savePath != "" { _ = w.WriteField("savepath", savePath) } - _ = w.Close() + if category != "" { + _ = w.WriteField("category", category) + } + if err := w.Close(); err != nil { + return "", err + } baseURL := strings.TrimRight(a.cfg.Host, "/") req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, @@ -95,14 +141,9 @@ func (a *QBitAdapter) AddTorrent(ctx context.Context, torrentURL, savePath strin return "", fmt.Errorf("qbittorrent add torrent: %d: %s", resp.StatusCode, strings.TrimSpace(string(raw))) } if strings.EqualFold(strings.TrimSpace(string(raw)), "Fails.") { - return "", fmt.Errorf("qbittorrent add torrent: rejected by downloader") + return externalID, errors.New("qbittorrent add torrent: rejected by downloader") } - return "", nil -} - -// AddMagnet 通过磁力链接添加种子。 -func (a *QBitAdapter) AddMagnet(ctx context.Context, magnet, savePath string) (string, error) { - return a.AddTorrent(ctx, magnet, savePath) + return externalID, nil } // Pause 暂停种子。 @@ -193,6 +234,36 @@ func (a *QBitAdapter) Remove(ctx context.Context, hash string, deleteFiles bool) return nil } +func (a *QBitAdapter) Relocate(ctx context.Context, hash, location string) error { + a.mu.Lock() + defer a.mu.Unlock() + if err := a.ensureAuthLocked(ctx); err != nil { + return err + } + baseURL := strings.TrimRight(a.cfg.Host, "/") + form := url.Values{} + form.Set("hashes", hash) + form.Set("location", location) + req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, + baseURL+"/api/v2/torrents/setLocation", strings.NewReader(form.Encode())) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.Header.Set("Referer", baseURL) + req.Header.Set("Origin", baseURL) + resp, err := a.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + body, _ := io.ReadAll(resp.Body) + if resp.StatusCode >= 400 { + return fmt.Errorf("qbittorrent setLocation: %d: %s", resp.StatusCode, strings.TrimSpace(string(body))) + } + return nil +} + // loginLocked 执行登录(调用者必须持有锁)。 func (a *QBitAdapter) loginLocked(ctx context.Context) error { if a.cfg.Host == "" { diff --git a/internal/service/qbittorrent_adp_list.go b/internal/service/qbittorrent_adp_list.go index 8f882be..d75251a 100644 --- a/internal/service/qbittorrent_adp_list.go +++ b/internal/service/qbittorrent_adp_list.go @@ -59,38 +59,42 @@ func (a *QBitAdapter) GetInfo(ctx context.Context, hash string) (*TorrentInfo, e } type qbitTorrentListItem struct { - Hash string `json:"hash"` - Name string `json:"name"` - State string `json:"state"` - Progress float32 `json:"progress"` - DLSpeed int64 `json:"dlspeed"` - UPSpeed int64 `json:"upspeed"` - NumSeeds int `json:"num_seeds"` - NumLeechs int `json:"num_leechs"` - Size int64 `json:"size"` - SavePath string `json:"save_path"` - AddedOn int64 `json:"added_on"` - Category string `json:"category"` - Tags string `json:"tags"` + Hash string `json:"hash"` + Name string `json:"name"` + State string `json:"state"` + Progress float32 `json:"progress"` + DLSpeed int64 `json:"dlspeed"` + UPSpeed int64 `json:"upspeed"` + NumSeeds int `json:"num_seeds"` + NumLeechs int `json:"num_leechs"` + Size int64 `json:"size"` + SavePath string `json:"save_path"` + AddedOn int64 `json:"added_on"` + Category string `json:"category"` + Tags string `json:"tags"` + ContentPath string `json:"content_path"` + CompletionOn int64 `json:"completion_on"` } func qbitTorrentListToInfo(items []qbitTorrentListItem) []TorrentInfo { result := make([]TorrentInfo, 0, len(items)) for _, item := range items { result = append(result, TorrentInfo{ - Hash: item.Hash, - Name: item.Name, - Size: item.Size, - Progress: float64(item.Progress), - DLSpeed: item.DLSpeed, - UPSpeed: item.UPSpeed, - State: item.State, - SavePath: item.SavePath, - NumSeeds: item.NumSeeds, - NumLeechs: item.NumLeechs, - AddedOn: time.Unix(item.AddedOn, 0), - Category: item.Category, - Tags: item.Tags, + Hash: item.Hash, + Name: item.Name, + Size: item.Size, + Progress: normalizedTorrentProgress(float64(item.Progress)), + DLSpeed: item.DLSpeed, + UPSpeed: item.UPSpeed, + State: canonicalTorrentState(item.State, float64(item.Progress)), + SavePath: item.SavePath, + NumSeeds: item.NumSeeds, + NumLeechs: item.NumLeechs, + AddedOn: time.Unix(item.AddedOn, 0), + Category: item.Category, + Tags: item.Tags, + ContentPath: item.ContentPath, + CompletionOn: item.CompletionOn, }) } return result diff --git a/internal/service/qbittorrent_convert.go b/internal/service/qbittorrent_convert.go index b5c07ca..e65e827 100644 --- a/internal/service/qbittorrent_convert.go +++ b/internal/service/qbittorrent_convert.go @@ -3,31 +3,38 @@ package service // QBitTorrentToInfo 将旧的 QBitTorrent 转换为新的 TorrentInfo。 func QBitTorrentToInfo(q QBitTorrent) TorrentInfo { return TorrentInfo{ - Hash: q.Hash, - Name: q.Name, - Size: q.Size, - Progress: float64(q.Progress), - DLSpeed: q.DLSpeed, - UPSpeed: q.UpSpeed, - State: q.State, - SavePath: q.SavePath, - NumSeeds: q.NumSeeds, - NumLeechs: q.NumLeech, + Hash: q.Hash, + Name: q.Name, + Size: q.Size, + Progress: float64(q.Progress), + DLSpeed: q.DLSpeed, + UPSpeed: q.UpSpeed, + State: q.State, + SavePath: q.SavePath, + NumSeeds: q.NumSeeds, + NumLeechs: q.NumLeech, + ContentPath: q.ContentPath, + CompletionOn: q.CompletionOn, } } // TorrentInfoToQBit 将 TorrentInfo 转换回旧的 QBitTorrent 格式(兼容性)。 func TorrentInfoToQBit(t TorrentInfo) QBitTorrent { return QBitTorrent{ - Hash: t.Hash, - Name: t.Name, - State: t.State, - Progress: float32(t.Progress), - DLSpeed: t.DLSpeed, - UpSpeed: t.UPSpeed, - NumSeeds: t.NumSeeds, - NumLeech: t.NumLeechs, - Size: t.Size, - SavePath: t.SavePath, + Hash: t.Hash, + ClientID: "", + Source: "", + Name: t.Name, + State: canonicalTorrentState(t.State, t.Progress), + Progress: float32(normalizedTorrentProgress(t.Progress)), + DLSpeed: t.DLSpeed, + UpSpeed: t.UPSpeed, + NumSeeds: t.NumSeeds, + NumLeech: t.NumLeechs, + Size: t.Size, + SavePath: t.SavePath, + Category: t.Category, + ContentPath: t.ContentPath, + CompletionOn: t.CompletionOn, } } diff --git a/internal/service/service.go b/internal/service/service.go index ab2288e..9a298f2 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -109,11 +109,6 @@ func (c *Container) Boot() { } go c.warmMediaSearchIndex(c.stopCtx) - // 加载所有已配置的下载客户端 - if err := c.DownloadMgr.LoadAll(c.stopCtx); err != nil { - c.Log.Warn("failed to load download clients", zap.Error(err)) - } - // 启动调度器定时任务 c.Scheduler.Start(c.stopCtx) diff --git a/internal/service/service_builder.go b/internal/service/service_builder.go index b7606c6..00bd233 100644 --- a/internal/service/service_builder.go +++ b/internal/service/service_builder.go @@ -158,6 +158,7 @@ func (b *serviceContainerBuilder) initIdentityServices() { func (b *serviceContainerBuilder) initSiteDownloadServices() { b.c.Site = NewSiteService(b.log, b.repos, b.flareSolverrURL()) b.c.Downloads = NewDownloadService(b.log, b.repos, b.c.WSHub, b.c.Organizer, b.c.Site) + b.c.Downloads.SetDownloadManager(b.c.DownloadMgr) b.c.Organizer.SetActiveDownloadPathProvider(b.c.Downloads.ActiveDownloadPaths) b.c.Downloads.SetScanner(b.c.Scan) b.c.Downloads.SetTaskTracker(b.c.Tasks) diff --git a/internal/service/site_download.go b/internal/service/site_download.go index da483fe..4e3bc74 100644 --- a/internal/service/site_download.go +++ b/internal/service/site_download.go @@ -81,11 +81,15 @@ func redactSensitiveDownloadURL(raw string) string { } func (s *SiteService) FetchTorrentFile(ctx context.Context, raw string) ([]byte, string, error) { - matched := s.matchSiteForURL(ctx, raw) - if matched == nil { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Host == "" { return nil, "", errors.New("no matching PT site for torrent URL") } - cfg := s.siteModelToConfig(matched) + matched := s.matchSiteForURL(ctx, raw) + cfg := SiteConfig{Timeout: 30 * time.Second} + if matched != nil { + cfg = s.siteModelToConfig(matched) + } timeout := cfg.Timeout if timeout <= 0 { timeout = 30 * time.Second diff --git a/internal/service/subscription_availability.go b/internal/service/subscription_availability.go index 6462aff..90d2e52 100644 --- a/internal/service/subscription_availability.go +++ b/internal/service/subscription_availability.go @@ -137,11 +137,11 @@ func (s *SubscriptionService) downloadTaskCountsAsPending(ctx context.Context, r } func (s *SubscriptionService) addLiveTorrentAvailability(ctx context.Context, queries []string, out *LocalAvailability) { - if s == nil || s.downloads == nil || s.downloads.qb == nil || out == nil { + if s == nil || s.downloads == nil || out == nil { return } - live, err := s.downloads.qb.List(ctx, "") - if err != nil { + live, err := s.downloads.listLiveTorrents(ctx, "") + if err != nil && len(live) == 0 { return } for _, torrent := range live { diff --git a/internal/service/subscription_delete.go b/internal/service/subscription_delete.go index 22af697..115f2fd 100644 --- a/internal/service/subscription_delete.go +++ b/internal/service/subscription_delete.go @@ -3,7 +3,6 @@ package service import ( "context" "fmt" - "net/url" "strings" "github.com/ShukeBta/MediaStationGo/internal/model" @@ -29,19 +28,21 @@ func (s *SubscriptionService) deleteSubscriptionDownloads(ctx context.Context, s } var live []QBitTorrent - if s.downloads != nil && s.downloads.qb != nil && s.downloads.qb.IsConfigured() { - live, _ = s.downloads.qb.List(ctx, "") + if s.downloads != nil { + live, _ = s.downloads.listLiveTorrents(ctx, "") } deletedHashes := map[string]struct{}{} for _, task := range candidates { - hash := downloadTaskInfoHash(task) - if hash == "" { - hash = matchingLiveTorrentHash(task, live) + hash := firstNonEmpty(task.ExternalID, downloadTaskInfoHash(task)) + clientID := strings.TrimSpace(task.DownloadClientID) + if matched, ok := matchingLiveTorrent(task, live); ok { + hash = firstNonEmpty(hash, matched.Hash) + clientID = firstNonEmpty(clientID, matched.ClientID) } - if hash != "" && s.downloads != nil && s.downloads.qb != nil && s.downloads.qb.IsConfigured() { - key := strings.ToLower(hash) + if hash != "" && s.downloads != nil { + key := strings.ToLower(clientID + ":" + hash) if _, ok := deletedHashes[key]; !ok { - if err := s.downloads.Delete(ctx, hash, false); err != nil { + if err := s.downloads.Delete(ctx, hash, false, clientID); err != nil { return fmt.Errorf("删除订阅关联下载任务 %q 失败: %w", task.Title, err) } deletedHashes[key] = struct{}{} @@ -79,32 +80,31 @@ func subscriptionDeleteMatchesTask(ctx context.Context, s *SubscriptionService, } func downloadTaskInfoHash(task model.DownloadTask) string { - raw := strings.TrimSpace(task.URL) - if raw == "" { - return "" - } - parsed, err := url.Parse(raw) - if err != nil { - return "" - } - if strings.EqualFold(parsed.Scheme, "magnet") { - for _, xt := range parsed.Query()["xt"] { - const prefix = "urn:btih:" - if strings.HasPrefix(strings.ToLower(xt), prefix) { - return strings.TrimSpace(xt[len(prefix):]) - } - } + return torrentURLInfoHash(task.URL) +} + +func matchingLiveTorrentHash(task model.DownloadTask, live []QBitTorrent) string { + if torrent, ok := matchingLiveTorrent(task, live); ok { + return strings.TrimSpace(torrent.Hash) } return "" } -func matchingLiveTorrentHash(task model.DownloadTask, live []QBitTorrent) string { +func matchingLiveTorrent(task model.DownloadTask, live []QBitTorrent) (QBitTorrent, bool) { + for _, torrent := range live { + if strings.TrimSpace(task.DownloadClientID) != "" && task.DownloadClientID != torrent.ClientID { + continue + } + if strings.TrimSpace(task.ExternalID) != "" && strings.EqualFold(task.ExternalID, torrent.Hash) { + return torrent, true + } + } key := downloadTaskIdentityKey(task.Title) if key == "" { key = downloadTaskIdentityKey(publicDownloadTitle(task.URL)) } if key == "" { - return "" + return QBitTorrent{}, false } for _, torrent := range live { current := downloadTaskIdentityKey(torrent.Name) @@ -112,10 +112,10 @@ func matchingLiveTorrentHash(task model.DownloadTask, live []QBitTorrent) string continue } if current == key || strings.Contains(current, key) || strings.Contains(key, current) { - return strings.TrimSpace(torrent.Hash) + return torrent, true } } - return "" + return QBitTorrent{}, false } func markDownloadTaskDeletedByID(ctx context.Context, db *gorm.DB, task model.DownloadTask) { diff --git a/internal/service/transmission_adp.go b/internal/service/transmission_adp.go index 1fe3ade..b609d5b 100644 --- a/internal/service/transmission_adp.go +++ b/internal/service/transmission_adp.go @@ -6,6 +6,7 @@ package service import ( "context" + "encoding/base64" "fmt" "net/http" "strings" @@ -34,6 +35,20 @@ func (a *TransmissionAdapter) AddTorrent(ctx context.Context, torrentURL, savePa a.mu.Lock() defer a.mu.Unlock() args := map[string]interface{}{"filename": torrentURL} + return a.addTorrentLocked(ctx, args, savePath) +} + +// AddTorrentFile submits application-fetched .torrent bytes through +// Transmission's base64 metainfo field. This keeps private tracker cookies and +// signed URLs inside MediaStationGo instead of asking Transmission to refetch. +func (a *TransmissionAdapter) AddTorrentFile(ctx context.Context, data []byte, _ string, savePath string) (string, error) { + a.mu.Lock() + defer a.mu.Unlock() + args := map[string]interface{}{"metainfo": base64.StdEncoding.EncodeToString(data)} + return a.addTorrentLocked(ctx, args, savePath) +} + +func (a *TransmissionAdapter) addTorrentLocked(ctx context.Context, args map[string]interface{}, savePath string) (string, error) { if savePath != "" { args["download-dir"] = savePath } @@ -99,7 +114,7 @@ func (a *TransmissionAdapter) List(ctx context.Context, filter string) ([]Torren "hashString", "name", "totalSize", "percentDone", "rateDownload", "rateUpload", "status", "downloadDir", "peersSendingToUs", "peersGettingFromUs", "addedDate", - "labels", "isStalled", + "doneDate", "labels", "isStalled", }, } resp, err := a.rpcLocked(ctx, "torrent-get", args) @@ -132,7 +147,7 @@ func (a *TransmissionAdapter) List(ctx context.Context, filter string) ([]Torren // Transmission 状态码转字符串 status := int(toFloat64(t["status"])) - state := transmissionStateStr(status) + state := canonicalTorrentState(transmissionStateStr(status), progress) // 过滤 if filter != "" && !strings.EqualFold(state, filter) { @@ -140,18 +155,20 @@ func (a *TransmissionAdapter) List(ctx context.Context, filter string) ([]Torren } result = append(result, TorrentInfo{ - Hash: hash, - Name: name, - Size: size, - Progress: progress * 100, - DLSpeed: dlSpeed, - UPSpeed: upSpeed, - State: state, - SavePath: savePath, - NumSeeds: numSeeds, - NumLeechs: numLeechs, - AddedOn: time.Unix(addedOn, 0), - Tags: toJSONLabels(t["labels"]), + Hash: hash, + Name: name, + Size: size, + Progress: normalizedTorrentProgress(progress), + DLSpeed: dlSpeed, + UPSpeed: upSpeed, + State: state, + SavePath: savePath, + NumSeeds: numSeeds, + NumLeechs: numLeechs, + AddedOn: time.Unix(addedOn, 0), + Tags: toJSONLabels(t["labels"]), + ContentPath: downloaderPayloadPath(savePath, name), + CompletionOn: toInt64(t["doneDate"]), }) } return result, nil @@ -166,7 +183,7 @@ func (a *TransmissionAdapter) GetInfo(ctx context.Context, hash string) (*Torren "fields": []string{ "hashString", "name", "totalSize", "percentDone", "rateDownload", "rateUpload", "status", "downloadDir", - "peersSendingToUs", "peersGettingFromUs", "addedDate", "labels", + "peersSendingToUs", "peersGettingFromUs", "addedDate", "doneDate", "labels", }, } resp, err := a.rpcLocked(ctx, "torrent-get", args) @@ -184,18 +201,20 @@ func (a *TransmissionAdapter) GetInfo(ctx context.Context, hash string) (*Torren status := int(toFloat64(t["status"])) info := &TorrentInfo{ - Hash: hash, - Name: strVal(t["name"]), - Size: toInt64(t["totalSize"]), - Progress: toFloat64(t["percentDone"]) * 100, - DLSpeed: toInt64(t["rateDownload"]), - UPSpeed: toInt64(t["rateUpload"]), - State: transmissionStateStr(status), - SavePath: strVal(t["downloadDir"]), - NumSeeds: int(toInt64(t["peersSendingToUs"])), - NumLeechs: int(toInt64(t["peersGettingFromUs"])), - AddedOn: time.Unix(int64(toFloat64(t["addedDate"])), 0), - Tags: toJSONLabels(t["labels"]), + Hash: hash, + Name: strVal(t["name"]), + Size: toInt64(t["totalSize"]), + Progress: normalizedTorrentProgress(toFloat64(t["percentDone"])), + DLSpeed: toInt64(t["rateDownload"]), + UPSpeed: toInt64(t["rateUpload"]), + State: canonicalTorrentState(transmissionStateStr(status), toFloat64(t["percentDone"])), + SavePath: strVal(t["downloadDir"]), + NumSeeds: int(toInt64(t["peersSendingToUs"])), + NumLeechs: int(toInt64(t["peersGettingFromUs"])), + AddedOn: time.Unix(int64(toFloat64(t["addedDate"])), 0), + Tags: toJSONLabels(t["labels"]), + ContentPath: downloaderPayloadPath(strVal(t["downloadDir"]), strVal(t["name"])), + CompletionOn: toInt64(t["doneDate"]), } return info, nil } diff --git a/web/src/api/downloads.ts b/web/src/api/downloads.ts index dea11b2..c630ceb 100644 --- a/web/src/api/downloads.ts +++ b/web/src/api/downloads.ts @@ -26,9 +26,9 @@ export const downloadsAPI = { .post('/downloads', { url, save_path: savePath, ...meta }) .then((r) => r.data), - remove: (hash: string, deleteFiles = false) => + remove: (hash: string, clientID: string, deleteFiles = false) => api - .delete(`/downloads/${hash}?delete_files=${deleteFiles ? 'true' : 'false'}`) + .delete(`/downloads/${hash}`, { params: { client_id: clientID, delete_files: deleteFiles } }) .then((r) => r.data), reload: () => api.post('/downloads/reload').then((r) => r.data), diff --git a/web/src/pages/DownloadsPage.tsx b/web/src/pages/DownloadsPage.tsx index 9713231..961b3aa 100644 --- a/web/src/pages/DownloadsPage.tsx +++ b/web/src/pages/DownloadsPage.tsx @@ -102,12 +102,12 @@ export function DownloadsPage() {
{torrents.map((torrent) => ( { if (!(await confirmAction({ title: '删除下载任务', message: `删除「${torrent.title || torrent.name}」?`, confirmText: '删除' }))) return - await downloadsAPI.remove(torrent.hash, false) + await downloadsAPI.remove(torrent.hash, torrent.client_id, false) toast.success('已删除任务') await refresh() }} diff --git a/web/src/types/downloads.ts b/web/src/types/downloads.ts index 5b0a1e1..974447d 100644 --- a/web/src/types/downloads.ts +++ b/web/src/types/downloads.ts @@ -1,6 +1,8 @@ export interface DownloadTask { id: string source: string + download_client_id?: string + external_id?: string title: string poster_url?: string backdrop_url?: string @@ -21,6 +23,8 @@ export interface DownloadTask { export interface QBitTorrent { hash: string + client_id: string + source: 'qbittorrent' | 'transmission' | 'aria2' name: string title: string poster_url?: string